Include mixed steps in Qwen235 state analysis

This commit is contained in:
2026-07-19 18:39:36 +08:00
parent b6dfbfcad7
commit 670eca176d

View File

@@ -183,7 +183,7 @@ def categorized_components(components: dict[str, Any]) -> dict[str, float]:
}
def load_decode_rows(state_root: Path, config: str) -> tuple[list[dict[str, Any]], dict]:
def load_stage_rows(state_root: Path, config: str) -> tuple[list[dict[str, Any]], dict]:
result_path = state_root / config / "result.json"
result = json.loads(result_path.read_text())
if result.get("status") != "PASS":
@@ -194,8 +194,8 @@ def load_decode_rows(state_root: Path, config: str) -> tuple[list[dict[str, Any]
rows = []
for row in load_jsonl(ledger):
token_counts = [int(value) for value in row["request_num_tokens"]]
if not token_counts or any(value != 1 for value in token_counts):
continue
if not token_counts:
raise ValueError(f"empty stage batch in {config}")
components = categorized_components(
row["execution_time"]["component_ledger_ms"]
)
@@ -205,6 +205,9 @@ def load_decode_rows(state_root: Path, config: str) -> tuple[list[dict[str, Any]
rows.append(
{
"batch_size": len(row["request_ids"]),
"decode_requests": sum(value == 1 for value in token_counts),
"prefill_requests": sum(value > 1 for value in token_counts),
"prefill_tokens": sum(value for value in token_counts if value > 1),
"total_time_ms": total,
"categories_ms": components,
"per_expert_tokens": row.get("per_expert_tokens"),
@@ -215,11 +218,15 @@ def load_decode_rows(state_root: Path, config: str) -> tuple[list[dict[str, Any]
return rows, result
def weighted_component_mean(rows: Iterable[dict[str, Any]]) -> dict[str, float]:
def weighted_component_mean(
rows: Iterable[dict[str, Any]], *, weight_name: str = "batch_size"
) -> dict[str, float]:
selected = list(rows)
if not selected:
raise ValueError("component mean needs rows")
weights = [row["batch_size"] for row in selected]
weights = [int(row[weight_name]) for row in selected]
if any(weight <= 0 for weight in weights):
raise ValueError(f"component weights must be positive: {weight_name}")
denominator = sum(weights)
values = {
category: sum(
@@ -288,13 +295,32 @@ def analyze(real_root: Path, state_root: Path) -> dict[str, Any]:
sim = {}
by_batch = {}
for config in CONFIGS:
rows, replay = load_decode_rows(state_root, config)
grouped = means_by_batch(rows)
rows, replay = load_stage_rows(state_root, config)
decode_bearing = [row for row in rows if row["decode_requests"] > 0]
pure_decode = [row for row in decode_bearing if row["prefill_requests"] == 0]
mixed = [row for row in decode_bearing if row["prefill_requests"] > 0]
if not pure_decode:
raise ValueError(f"no pure-decode ledger rows for {config}")
grouped = means_by_batch(pure_decode)
by_batch[config] = grouped
sim[config] = {
"decode_only_rows": len(rows),
"decode_batch_size": numeric(row["batch_size"] for row in rows),
"token_weighted_components_ms": weighted_component_mean(rows),
"stage_rows": len(rows),
"decode_bearing_rows": len(decode_bearing),
"pure_decode_rows": len(pure_decode),
"mixed_prefill_decode_rows": len(mixed),
"decode_batch_size": numeric(
row["decode_requests"] for row in decode_bearing
),
"decode_token_weighted_all_step_components_ms": weighted_component_mean(
decode_bearing, weight_name="decode_requests"
),
"decode_token_weighted_mixed_step_share": sum(
row["decode_requests"] for row in mixed
)
/ sum(row["decode_requests"] for row in decode_bearing),
"pure_decode_token_weighted_components_ms": weighted_component_mean(
pure_decode, weight_name="decode_requests"
),
"batch_support": {
str(batch): value for batch, value in grouped.items()
},
@@ -307,6 +333,10 @@ def analyze(real_root: Path, state_root: Path) -> dict[str, Any]:
)
for config in CONFIGS
}
simulator_internal_contrast = subtract(
sim[CONFIGS[1]]["decode_token_weighted_all_step_components_ms"],
sim[CONFIGS[0]]["decode_token_weighted_all_step_components_ms"],
)
coverages = [real_reweighted[config]["coverage"] for config in CONFIGS]
proxy_contrast = None
if min(coverages) >= 0.8:
@@ -357,6 +387,7 @@ def analyze(real_root: Path, state_root: Path) -> dict[str, Any]:
},
"real": real,
"simulator": sim,
"simulator_internal_all_step_tp8_minus_tp4_ms": simulator_internal_contrast,
"proxy_matched": {
"method": (
"exact batch-size lookup; each simulator config is reweighted to its own "
@@ -390,6 +421,22 @@ def markdown(result: dict[str, Any]) -> str:
f"| {result['proxy_matched']['configs'][config]['coverage']:.1%} |"
)
contrast = result["proxy_matched"]["tp8_minus_tp4_ms"]
internal = result["simulator_internal_all_step_tp8_minus_tp4_ms"]
lines.extend(
[
"",
"## Frontier internal component contrast over its own executed composition",
"",
"Decode-token-weighted over both pure-decode and mixed prefill/decode steps.",
"",
"| Component | TP8 - TP4 (ms/decoded token step) |",
"|---|---:|",
]
)
for name, value in sorted(
internal.items(), key=lambda item: abs(item[1]), reverse=True
):
lines.append(f"| {name} | {value:+.4f} |")
if contrast is not None:
lines.extend(
[