Files
SkillCompiler/data/skills-bench/docs/paper-figures/scripts/04_task_heterogeneity.py
T
2026-09-04 14:58:42 +08:00

249 lines
9.8 KiBLFS
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
r"""04 Task Heterogeneity (before vs. after).
Square scatter: x = baseline pass-rate (no skills), y = pass-rate WITH curated
skills. The dashed y = x diagonal is the "skills did nothing" reference; the
median-baseline vertical line splits low- vs. high-baseline tasks. Together
they yield four named quadrants:
top-left low baseline, skills helped ("skill-rescued")
bottom-left low baseline, skills didn't lift ("stuck-low")
top-right high baseline, skills helped ("skill-amplified")
bottom-right high baseline, skills hurt ("context-burden")
================================================================================
FAKE DATA FORMAT
================================================================================
Module-level constants:
DOMAIN_PALETTE: dict[str, hex_color] — 10 domains and their colors
DOMAINS: list[str] — keys of DOMAIN_PALETTE (point colors)
TIER_SIZE: dict[str, int] — Core/Extended/Extreme → marker size
REGION_SPECS: list[dict] — quadrant labels + colors
`_fake_tasks(n=84)` generates a list of 84 task dicts:
domain: str — sampled from DOMAINS with the per-domain probability
array hard-coded inside the function (sums to 1)
tier: str — "Core" (0.5) / "Extended" (0.35) / "Extreme" (0.15)
baseline: float — beta(2.5, 2.5) * 0.95, clipped to [0.02, 0.95]
delta_pp: float — drawn from a tier-conditional Normal so that the resulting
(baseline, baseline+delta) point lands in a plausible
quadrant (most points lift; some stuck-low or context-burden)
name: str — "{domain[:3]}-task-{i:02d}"
The last 4 entries are overwritten by `extremes`, four hand-picked extreme
points used as visual anchors (e.g., a +85.7 pp rescue and a -39.3 pp burden).
================================================================================
"""
from __future__ import annotations
from pathlib import Path
import matplotlib.pyplot as plt
import matplotlib.ticker as mticker
from matplotlib.lines import Line2D
from matplotlib.patches import Patch, Polygon
import numpy as np
from utils import apply_style
OUTPUT_PATH = Path(__file__).resolve().parent.parent / "figures" / "04_task_heterogeneity.pdf"
DOMAIN_PALETTE = {
"Office": "#2563eb",
"SoftwareEng": "#059669",
"Web/Vis": "#d97706",
"Security": "#dc2626",
"Manufacturing": "#7c3aed",
"BizAuto": "#0891b2",
"Graph/Dialog": "#64748b",
"ScientificComp": "#db2777",
"Optimization": "#ea580c",
"Multimodal": "#16a34a",
}
DOMAINS = list(DOMAIN_PALETTE.keys())
TIER_SIZE = {"Core": 36, "Extended": 76, "Extreme": 130}
REGION_SPECS = [
{
"key": "skill-rescued",
"label": "Skill-rescued",
"description": "low baseline, helped",
"color": "#dcfce7",
"text_color": "#166534",
"alpha": 0.60,
},
{
"key": "stuck-low",
"label": "Stuck-low",
"description": "low baseline, little lift",
"color": "#fef3c7",
"text_color": "#a16207",
"alpha": 0.65,
},
{
"key": "skill-amplified",
"label": "Skill-amplified",
"description": "high baseline, helped",
"color": "#dbeafe",
"text_color": "#1e40af",
"alpha": 0.55,
},
{
"key": "context-burden",
"label": "Context-burden",
"description": "high baseline, hurt",
"color": "#fee2e2",
"text_color": "#991b1b",
"alpha": 0.65,
},
]
def _fake_tasks(n: int = 84, seed: int = 11) -> list[dict]:
rng = np.random.default_rng(seed)
domain_dist = rng.choice(len(DOMAINS), size=n,
p=np.array([10, 11, 9, 8, 9, 8, 6, 10, 7, 6]) / 84)
tier_dist = rng.choice(["Core", "Extended", "Extreme"], size=n, p=[0.5, 0.35, 0.15])
tasks = []
for i in range(n):
d = DOMAINS[domain_dist[i]]
tier = tier_dist[i]
baseline = float(np.clip(rng.beta(2.5, 2.5) * 0.95, 0.02, 0.95))
delta_mean_pp = 24 * (1 - abs(baseline - 0.4)) - 4
delta_pp = float(rng.normal(delta_mean_pp, 16))
if i % 84 in (3, 9, 14, 21, 28, 33, 41, 49, 55, 60, 65, 70, 74, 77, 80, 82):
delta_pp = float(rng.uniform(-40, -2))
tasks.append(dict(domain=d, tier=tier, baseline=baseline,
delta_pp=delta_pp,
name=f"{d.lower()[:3]}-task-{i:02d}"))
extremes = [
dict(domain="Office", tier="Core", baseline=0.00, delta_pp=85.7, name="mario-coin-counting"),
dict(domain="BizAuto", tier="Core", baseline=0.00, delta_pp=85.7, name="sales-pivot-analysis"),
dict(domain="Graph/Dialog", tier="Extended", baseline=0.46, delta_pp=-39.3, name="taxonomy-tree-merge"),
dict(domain="Optimization", tier="Extreme", baseline=0.32, delta_pp=-14.3, name="energy-ac-opf"),
]
tasks[-4:] = extremes
return tasks[:84]
def _with_skill(task: dict) -> float:
"""Map (baseline, delta_pp) -> with-skill pass rate clipped to [0, 1]."""
return float(np.clip(task["baseline"] + task["delta_pp"] / 100.0, 0.0, 1.0))
def main() -> None:
apply_style()
tasks = _fake_tasks()
median_baseline = float(np.median([t["baseline"] for t in tasks]))
# Half-page wrapfigure layout: legends stacked above the square scatter,
# no marginal histograms (user removed them to keep the figure compact).
# Region legend goes inside the main scatter (lower-right corner);
# domain legend stays below the scatter, multi-row, full figure width.
fig = plt.figure(figsize=(4.4, 5.5))
gs = fig.add_gridspec(
2, 1,
height_ratios=[5.5, 0.95],
hspace=0.05,
)
ax = fig.add_subplot(gs[0, 0])
ax_leg_domain = fig.add_subplot(gs[1, 0]); ax_leg_domain.axis("off")
# 4-region shading (median baseline split × diagonal y=x).
m = median_baseline
region_polygons = {
# top-left: x<m AND y>x
"skill-rescued": [(0, 0), (m, m), (m, 1), (0, 1)],
# bottom-left: x<m AND y<=x
"stuck-low": [(0, 0), (m, 0), (m, m)],
# top-right: x>=m AND y>x
"skill-amplified": [(m, m), (1, 1), (m, 1)],
# bottom-right: x>=m AND y<=x
"context-burden": [(m, 0), (1, 0), (1, 1), (m, m)],
}
for spec in REGION_SPECS:
ax.add_patch(Polygon(region_polygons[spec["key"]], closed=True,
facecolor=spec["color"], edgecolor="none",
alpha=spec["alpha"], zorder=0))
# Diagonal "skills did nothing" reference + median-baseline vertical.
ax.plot([0, 1], [0, 1], color="#1f2937", linewidth=0.9,
linestyle=(0, (5, 3)), alpha=0.65, zorder=1)
ax.axvline(m, color="#1f2937", linewidth=0.6, alpha=0.45,
linestyle=(0, (2, 2)), zorder=1)
# Plot points: each task at (baseline, with-skill).
for t in tasks:
ax.scatter(
t["baseline"], _with_skill(t),
s=TIER_SIZE[t["tier"]], marker="o",
facecolor=DOMAIN_PALETTE[t["domain"]],
edgecolor="white", linewidth=0.5,
alpha=0.85, zorder=3,
)
ax.set_xlim(-0.035, 1.035)
ax.set_ylim(-0.035, 1.035)
ax.set_aspect("equal")
ax.set_xlabel("Baseline pass rate (no skills)")
ax.set_ylabel("Pass rate with curated skills")
ticks = np.linspace(0, 1, 6)
ax.set_xticks(ticks)
ax.set_yticks(ticks)
ax.xaxis.set_major_formatter(mticker.PercentFormatter(xmax=1.0, decimals=0))
ax.yaxis.set_major_formatter(mticker.PercentFormatter(xmax=1.0, decimals=0))
ax.grid(True, alpha=0.30)
ax.set_axisbelow(True)
# Marginal histograms removed (user request: no top/right hist bars).
region_handles = [
Patch(facecolor=spec["color"], edgecolor="#d1d5db",
alpha=spec["alpha"], label=spec["label"])
for spec in REGION_SPECS
]
reference_handles = [
Line2D([0], [0], color="#1f2937", linewidth=0.9,
linestyle=(0, (5, 3)), alpha=0.65, label="y = x: no skill effect"),
Line2D([0], [0], color="#1f2937", linewidth=0.8,
linestyle=(0, (2, 2)), alpha=0.45,
label=f"Median baseline: {m * 100:.0f}%"),
]
domain_handles = [
Line2D([0], [0], marker="o", linestyle="none", markersize=7,
markerfacecolor=DOMAIN_PALETTE[d], markeredgecolor="white",
markeredgewidth=0.6, label=d)
for d in DOMAINS
]
tier_handles = [
Line2D([0], [0], marker="o", linestyle="none",
markersize=np.sqrt(TIER_SIZE[t]),
markerfacecolor="#475569", markeredgecolor="white",
markeredgewidth=0.6, label=t)
for t in ["Core", "Extended", "Extreme"]
]
# Region legend inside the scatter, lower-right corner (Context-burden
# quadrant — typically sparse). White background to stay readable on top
# of any markers in the area.
ax.legend(handles=region_handles,
loc="lower right", ncol=1,
fontsize=6.8, frameon=True, framealpha=0.95,
handlelength=1.2, handletextpad=0.4,
labelspacing=0.30, borderpad=0.4).get_frame().set_edgecolor("#d1d5db")
# Domain legend below the scatter, ncol=6 → 2 rows, width matches figure.
ax_leg_domain.legend(handles=domain_handles,
loc="center", ncol=5,
fontsize=6.8, frameon=False,
handletextpad=0.35, columnspacing=1.0,
labelspacing=0.30)
fig.subplots_adjust(top=0.99, bottom=0.07, left=0.12, right=0.98)
OUTPUT_PATH.parent.mkdir(parents=True, exist_ok=True)
fig.savefig(OUTPUT_PATH, bbox_inches="tight")
plt.close(fig)
print(f"[04] wrote {OUTPUT_PATH}")
if __name__ == "__main__":
main()