249 lines
9.8 KiBLFS
Python
249 lines
9.8 KiBLFS
Python
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()
|