From 9c9479c313923c89d81f90bc3d8b2f8bef250337 Mon Sep 17 00:00:00 2001 From: Gahow Wang Date: Sun, 19 Jul 2026 19:00:18 +0800 Subject: [PATCH] Match Frontier to exact iteration composition --- .../analyze_qwen235_fixed_pd_state.py | 202 +++++++++++++++++- .../test_fidelity_envelope.py | 15 ++ 2 files changed, 210 insertions(+), 7 deletions(-) diff --git a/runs/frontier-fidelity-envelope-v1/analyze_qwen235_fixed_pd_state.py b/runs/frontier-fidelity-envelope-v1/analyze_qwen235_fixed_pd_state.py index 03d021d..e812b95 100644 --- a/runs/frontier-fidelity-envelope-v1/analyze_qwen235_fixed_pd_state.py +++ b/runs/frontier-fidelity-envelope-v1/analyze_qwen235_fixed_pd_state.py @@ -31,6 +31,13 @@ LOG_PATTERN = re.compile( r"Running: (?P[0-9]+) reqs, Waiting: (?P[0-9]+) reqs, " r"GPU KV cache usage: (?P[0-9.]+)%" ) +ITERATION_PATTERN = re.compile( + r"Iteration\((?P[0-9]+)\): (?P[0-9]+) context requests, " + r"(?P[0-9]+) context tokens, " + r"(?P[0-9]+) generation requests, " + r"(?P[0-9]+) generation tokens, iteration elapsed time: " + r"(?P[0-9.]+) ms" +) CATEGORIES = { @@ -229,6 +236,84 @@ def summarize_real(real_root: Path) -> dict[str, Any]: return result +def parse_iteration_log(path: Path) -> list[dict[str, float | int]]: + rows = [] + for line in path.read_text(errors="replace").splitlines(): + match = ITERATION_PATTERN.search(line) + if match is None: + continue + rows.append( + { + "index": int(match.group("index")), + "context_requests": int(match.group("context_requests")), + "context_tokens": int(match.group("context_tokens")), + "generation_requests": int(match.group("generation_requests")), + "generation_tokens": int(match.group("generation_tokens")), + "elapsed_ms": float(match.group("elapsed_ms")), + } + ) + if not rows: + raise ValueError(f"no iteration details found: {path}") + start = next( + (index for index, row in enumerate(rows) if row["context_tokens"] >= 4096), + None, + ) + if start is None: + raise ValueError(f"measured Fixed-PD interval not found: {path}") + measured = rows[start:] + if any(row["generation_tokens"] != row["generation_requests"] for row in measured): + raise ValueError("Fixed-PD decode must schedule one token per generation request") + return measured + + +def summarize_iteration_real(iteration_root: Path) -> dict[str, Any]: + result = {} + for config in CONFIGS: + path = iteration_root / config / "logs/server.log" + rows = parse_iteration_log(path) + decode = [row for row in rows if row["generation_requests"] > 0] + pure = [row for row in decode if row["context_tokens"] == 0] + mixed = [row for row in decode if row["context_tokens"] > 0] + if not pure or not mixed: + raise ValueError(f"iteration state lacks pure or mixed decode rows: {path}") + joint = Counter( + (int(row["context_tokens"]), int(row["generation_requests"])) + for row in decode + ) + token_joint = { + f"{context}:{generation}": count * generation + for (context, generation), count in sorted(joint.items()) + } + pure_hist = Counter(int(row["generation_requests"]) for row in pure) + weights = [int(row["generation_requests"]) for row in decode] + result[config] = { + "path": str(path.resolve()), + "measured_rows": len(rows), + "decode_bearing_rows": len(decode), + "pure_decode_rows": len(pure), + "mixed_prefill_decode_rows": len(mixed), + "decode_batch_size": numeric( + row["generation_requests"] for row in decode + ), + "pure_decode_batch_histogram": { + str(key): value for key, value in sorted(pure_hist.items()) + }, + "decode_token_weighted_joint_state_histogram": token_joint, + "decode_token_weighted_iteration_elapsed_ms": sum( + float(row["elapsed_ms"]) * weight + for row, weight in zip(decode, weights, strict=True) + ) + / sum(weights), + "pure_decode_iteration_elapsed_ms": numeric( + row["elapsed_ms"] for row in pure + ), + "mixed_iteration_elapsed_ms": numeric( + row["elapsed_ms"] for row in mixed + ), + } + return result + + def categorized_components(components: dict[str, Any]) -> dict[str, float]: covered = set().union(*CATEGORIES.values()) unknown = set(components) - covered @@ -396,6 +481,24 @@ def means_by_batch(rows: list[dict[str, Any]]) -> dict[int, dict[str, Any]]: } +def means_by_joint_state(rows: list[dict[str, Any]]) -> dict[str, dict[str, Any]]: + groups: dict[str, list[dict[str, Any]]] = defaultdict(list) + for row in rows: + if row["decode_requests"] <= 0: + continue + key = f"{row['prefill_tokens']}:{row['decode_requests']}" + groups[key].append(row) + return { + key: { + "n": len(group), + "components_ms": weighted_component_mean( + group, weight_name="decode_requests" + ), + } + for key, group in sorted(groups.items()) + } + + def reweight( by_batch: dict[int, dict[str, Any]], histogram: dict[str, int] ) -> dict[str, Any]: @@ -405,6 +508,42 @@ def reweight( for batch_size, count in histogram.items() if int(batch_size) in by_batch } + + +def reweight_joint( + by_state: dict[str, dict[str, Any]], token_histogram: dict[str, int] +) -> dict[str, Any]: + total_tokens = sum(token_histogram.values()) + supported = { + state: count + for state, count in token_histogram.items() + if state in by_state + } + covered = sum(supported.values()) + if covered == 0: + return { + "coverage": 0.0, + "components_ms": None, + "unsupported": token_histogram, + } + return { + "coverage": covered / total_tokens, + "supported_decode_tokens": covered, + "total_decode_tokens": total_tokens, + "components_ms": { + category: sum( + by_state[state]["components_ms"][category] * count + for state, count in supported.items() + ) + / covered + for category in (*CATEGORIES, "total") + }, + "unsupported": { + state: count + for state, count in token_histogram.items() + if state not in by_state + }, + } covered = sum(supported.values()) if covered == 0: return {"coverage": 0.0, "components_ms": None, "unsupported": histogram} @@ -434,11 +573,18 @@ def subtract(right: dict[str, float], left: dict[str, float]) -> dict[str, float def analyze( - real_root: Path, state_root: Path, op_trace_root: Path | None = None + real_root: Path, + state_root: Path, + op_trace_root: Path | None = None, + iteration_root: Path | None = None, ) -> dict[str, Any]: real = summarize_real(real_root) + iteration_real = ( + summarize_iteration_real(iteration_root) if iteration_root is not None else None + ) sim = {} by_batch = {} + by_joint_state = {} for config in CONFIGS: component_source = "stage_batch_ledger" trace_validation = None @@ -454,6 +600,7 @@ def analyze( raise ValueError(f"no pure-decode ledger rows for {config}") grouped = means_by_batch(pure_decode) by_batch[config] = grouped + by_joint_state[config] = means_by_joint_state(decode_bearing) sim[config] = { "stage_rows": len(rows), "decode_bearing_rows": len(decode_bearing), @@ -486,6 +633,23 @@ def analyze( ) for config in CONFIGS } + exact_reweighted = None + exact_contrast = None + if iteration_real is not None: + exact_reweighted = { + config: reweight_joint( + by_joint_state[config], + iteration_real[config][ + "decode_token_weighted_joint_state_histogram" + ], + ) + for config in CONFIGS + } + if min(value["coverage"] for value in exact_reweighted.values()) >= 0.8: + exact_contrast = subtract( + exact_reweighted[CONFIGS[1]]["components_ms"], + exact_reweighted[CONFIGS[0]]["components_ms"], + ) simulator_internal_contrast = subtract( sim[CONFIGS[1]]["decode_token_weighted_all_step_components_ms"], sim[CONFIGS[0]]["decode_token_weighted_all_step_components_ms"], @@ -512,12 +676,13 @@ def analyze( if min(by_batch[config][batch]["n"] for config in CONFIGS) >= 10 } observed_real_contrast = REAL_TPOT_MS[CONFIGS[1]] - REAL_TPOT_MS[CONFIGS[0]] - if proxy_contrast is None: + decision_contrast = exact_contrast if exact_contrast is not None else proxy_contrast + if decision_contrast is None: verdict = "STOP: real Running proxy has insufficient exact simulator support" - elif proxy_contrast["total"] < 0: + elif decision_contrast["total"] < 0: verdict = ( - "Active-batch-count mismatch alone is insufficient: after exact reweighting " - "to each config's real Running histogram, Frontier still predicts TP8 faster." + "State-composition mismatch is insufficient: after reweighting Frontier " + "to measured real decode composition, it still predicts TP8 faster." ) else: verdict = ( @@ -539,6 +704,7 @@ def analyze( ), }, "real": real, + "real_iteration_state": iteration_real, "simulator": sim, "simulator_internal_all_step_tp8_minus_tp4_ms": simulator_internal_contrast, "proxy_matched": { @@ -549,6 +715,16 @@ def analyze( "configs": real_reweighted, "tp8_minus_tp4_ms": proxy_contrast, }, + "exact_state_matched": { + "method": ( + "decode-token-weighted exact (context_tokens, generation_requests) " + "composition from vLLM iteration details; no interpolation" + ), + "configs": exact_reweighted, + "tp8_minus_tp4_ms": exact_contrast, + } + if iteration_real is not None + else None, "same_batch_contrasts": shared_contrasts, "reference": { "real_tpot_ms": REAL_TPOT_MS, @@ -573,7 +749,12 @@ def markdown(result: dict[str, Any]) -> str: f"| {result['simulator'][config]['decode_batch_size']['mean']:.3f} " f"| {result['proxy_matched']['configs'][config]['coverage']:.1%} |" ) - contrast = result["proxy_matched"]["tp8_minus_tp4_ms"] + exact = result["exact_state_matched"] + contrast = ( + exact["tp8_minus_tp4_ms"] + if exact is not None and exact["tp8_minus_tp4_ms"] is not None + else result["proxy_matched"]["tp8_minus_tp4_ms"] + ) internal = result["simulator_internal_all_step_tp8_minus_tp4_ms"] lines.extend( [ @@ -591,10 +772,15 @@ def markdown(result: dict[str, Any]) -> str: ): lines.append(f"| {name} | {value:+.4f} |") if contrast is not None: + heading = ( + "Frontier component contrast at exact real token composition" + if exact is not None and exact["tp8_minus_tp4_ms"] is not None + else "Frontier internal component contrast at real Running proxy" + ) lines.extend( [ "", - "## Frontier internal component contrast at real Running proxy", + f"## {heading}", "", "Positive means TP8 slower; negative means Frontier gives TP8 an advantage.", "", @@ -626,6 +812,7 @@ def parse_args() -> argparse.Namespace: parser.add_argument("--real-root", type=Path, required=True) parser.add_argument("--state-root", type=Path, required=True) parser.add_argument("--op-trace-root", type=Path) + parser.add_argument("--iteration-root", type=Path) parser.add_argument("--json-output", type=Path, required=True) parser.add_argument("--markdown-output", type=Path, required=True) return parser.parse_args() @@ -637,6 +824,7 @@ def main() -> None: args.real_root.resolve(), args.state_root.resolve(), args.op_trace_root.resolve() if args.op_trace_root else None, + args.iteration_root.resolve() if args.iteration_root else None, ) atomic_write(args.json_output, json.dumps(result, indent=2, sort_keys=True) + "\n") atomic_write(args.markdown_output, markdown(result)) diff --git a/runs/frontier-fidelity-envelope-v1/test_fidelity_envelope.py b/runs/frontier-fidelity-envelope-v1/test_fidelity_envelope.py index 7abf5d7..b75e17d 100644 --- a/runs/frontier-fidelity-envelope-v1/test_fidelity_envelope.py +++ b/runs/frontier-fidelity-envelope-v1/test_fidelity_envelope.py @@ -218,6 +218,21 @@ class FidelityEnvelopeTest(unittest.TestCase): ], ) + def test_qwen235_iteration_state_parser_excludes_warmup(self) -> None: + module = load("analyze_qwen235_fixed_pd_state.py") + rows = [ + "Iteration(0): 1 context requests, 512 context tokens, 0 generation requests, 0 generation tokens, iteration elapsed time: 10.00 ms", + "Iteration(1): 1 context requests, 4096 context tokens, 4 generation requests, 4 generation tokens, iteration elapsed time: 30.00 ms", + "Iteration(2): 0 context requests, 0 context tokens, 5 generation requests, 5 generation tokens, iteration elapsed time: 20.00 ms", + ] + with tempfile.TemporaryDirectory() as temporary: + path = Path(temporary) / "server.log" + path.write_text("\n".join(rows) + "\n") + parsed = module.parse_iteration_log(path) + self.assertEqual([row["index"] for row in parsed], [1, 2]) + self.assertEqual(parsed[0]["context_tokens"], 4096) + self.assertEqual(parsed[1]["generation_requests"], 5) + def test_qwen235_component_categories_are_additive(self) -> None: module = load("analyze_qwen235_fixed_pd_state.py") components = {