145 lines
5.8 KiBLFS
Python
145 lines
5.8 KiBLFS
Python
"""03 Per-domain TSAPR attribution profile (stacked bars).
|
||
|
||
Each domain has one bar stacked by TSAPR function attribution count. Because
|
||
attributions are multi-label, total height can exceed n_success. Domains
|
||
sorted left-to-right by total successful trajectories.
|
||
|
||
================================================================================
|
||
FAKE DATA FORMAT
|
||
================================================================================
|
||
This script is fully synthetic. The data source is two module-level constants:
|
||
|
||
DOMAINS: list[str] — 11 domain names (x-axis order)
|
||
TSAPR: list[str] — 5 function letters: T, S, A, P, R
|
||
|
||
`_demo()` returns a dict {domain: ndarray of shape (n_success, 5)} where:
|
||
- n_success per domain is hard-coded inside `_demo()` (proportional to
|
||
task count; 20–62 successful trajectories per domain)
|
||
- each row is one successful trajectory, each column is one TSAPR function
|
||
- cell (i, j) = 1 if function j was attributed to trajectory i, else 0
|
||
- rows are drawn from Bernoulli(p_d_j) where p_d_j is a per-(domain,
|
||
function) base probability stored in `_demo()`'s `domain_bias` dict
|
||
The plot then sums each column to get attribution counts per (domain, function).
|
||
================================================================================
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
from pathlib import Path
|
||
import numpy as np
|
||
import matplotlib.pyplot as plt
|
||
from matplotlib.patches import Patch
|
||
|
||
from utils import apply_style, save_figure
|
||
|
||
OUTPUT_PATH = Path(__file__).resolve().parent.parent / "figures" / "03_tsapr_attribution_domain.pdf"
|
||
|
||
TSAPR = ["T", "S", "A", "P", "R"]
|
||
LABELS = {
|
||
"T": "Trigger",
|
||
"S": "State",
|
||
"A": "Action",
|
||
"P": "Procedure",
|
||
"R": "Runtime",
|
||
}
|
||
COLORS = {
|
||
"T": "#0891b2", "S": "#7c3aed", "A": "#059669",
|
||
"P": "#d97706", "R": "#dc2626",
|
||
}
|
||
DOMAINS = [
|
||
"Office", "SoftwareEng", "Web/Vis", "Security",
|
||
"Manufacturing", "BizAuto", "Graph/Dialog",
|
||
"ScientificComp", "Optimization", "Multimodal", "Retrieval",
|
||
]
|
||
|
||
|
||
def _demo(seed: int = 7) -> dict[str, np.ndarray]:
|
||
"""Returns {domain: (n_success, 5)} — each row is one successful trajectory."""
|
||
rng = np.random.default_rng(seed)
|
||
# n_success per domain (proportional to task count)
|
||
n_success = {
|
||
"Office": 55, "SoftwareEng": 42, "Web/Vis": 38,
|
||
"Security": 30, "Manufacturing": 28, "BizAuto": 45,
|
||
"Graph/Dialog": 62, "ScientificComp": 20, "Optimization": 24,
|
||
"Multimodal": 32, "Retrieval": 58,
|
||
}
|
||
# Domain-specific TSAPR bias (base probability per function)
|
||
domain_bias = {
|
||
"Office": [0.25, 0.55, 0.40, 0.80, 0.20],
|
||
"SoftwareEng": [0.30, 0.35, 0.75, 0.65, 0.15],
|
||
"Web/Vis": [0.20, 0.30, 0.70, 0.60, 0.10],
|
||
"Security": [0.40, 0.50, 0.45, 0.70, 0.55],
|
||
"Manufacturing": [0.35, 0.60, 0.50, 0.85, 0.30],
|
||
"BizAuto": [0.45, 0.40, 0.55, 0.75, 0.25],
|
||
"Graph/Dialog": [0.50, 0.65, 0.30, 0.60, 0.20],
|
||
"ScientificComp": [0.25, 0.70, 0.60, 0.55, 0.35],
|
||
"Optimization": [0.30, 0.50, 0.65, 0.80, 0.28],
|
||
"Multimodal": [0.35, 0.55, 0.45, 0.65, 0.40],
|
||
"Retrieval": [0.40, 0.60, 0.35, 0.70, 0.18],
|
||
}
|
||
out = {}
|
||
for d in DOMAINS:
|
||
n = n_success[d]
|
||
p = np.array(domain_bias[d])
|
||
mat = (rng.random((n, 5)) < p).astype(int)
|
||
out[d] = mat
|
||
return out
|
||
|
||
|
||
def main() -> None:
|
||
apply_style()
|
||
data = _demo()
|
||
|
||
# Sort domains by total successful trajectories descending
|
||
domains_sorted = sorted(DOMAINS, key=lambda d: -len(data[d]))
|
||
|
||
fig, ax_main = plt.subplots(figsize=(11.0, 3.6))
|
||
|
||
# ── Per-domain stacked bars ───────────────────────────────────────────────
|
||
xs = np.arange(len(domains_sorted))
|
||
bar_w = 0.62
|
||
|
||
for xi, dom in enumerate(domains_sorted):
|
||
mat = data[dom]
|
||
bottom = 0
|
||
for fi, func in enumerate(TSAPR):
|
||
count = int(mat[:, fi].sum())
|
||
if count == 0:
|
||
continue
|
||
ax_main.bar(xi, count, bottom=bottom, width=bar_w,
|
||
color=COLORS[func], edgecolor="white",
|
||
linewidth=0.5, zorder=3, alpha=0.90)
|
||
if count >= 6:
|
||
ax_main.text(xi, bottom + count / 2, func,
|
||
ha="center", va="center",
|
||
fontsize=7.5, fontweight="bold",
|
||
color="white", zorder=4)
|
||
bottom += count
|
||
|
||
ax_main.set_xticks(xs)
|
||
ax_main.set_xticklabels(domains_sorted, rotation=38, ha="right",
|
||
rotation_mode="anchor", fontsize=9.8)
|
||
ax_main.set_ylabel("Attribution count\n(multi-label: sums exceed n_success)", fontsize=10.5)
|
||
ax_main.set_title(
|
||
"TSAPR attribution per domain — successful skill trajectories (demo data)",
|
||
fontsize=11.0, pad=8, loc="left",
|
||
)
|
||
ax_main.grid(axis="y", alpha=0.35, zorder=0)
|
||
ax_main.set_axisbelow(True)
|
||
|
||
# ── Legend ────────────────────────────────────────────────────────────────
|
||
handles = [Patch(facecolor=COLORS[f], edgecolor="white",
|
||
label=f"{f} {LABELS[f]}")
|
||
for f in TSAPR]
|
||
ax_main.legend(handles=handles, loc="upper right",
|
||
ncol=1, fontsize=9.8, frameon=True, framealpha=0.93,
|
||
handlelength=1.2, columnspacing=1.0,
|
||
title="TSAPR function", title_fontsize=9.8)
|
||
|
||
fig.subplots_adjust(top=0.88, bottom=0.24, left=0.09, right=0.97)
|
||
save_figure(fig, OUTPUT_PATH)
|
||
print(f"[03] wrote {OUTPUT_PATH}")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|