Spaces:
Running
Running
| from __future__ import annotations | |
| import math | |
| import random | |
| from copy import deepcopy | |
| from statistics import mean, median | |
| from typing import Any | |
| from .execution import ExecutionLearningConfig, generate_workflows, run_execution_learning | |
| POLICIES = ( | |
| ("No prefetch", "none"), | |
| ("Top-1 decayed", "decayed"), | |
| ("Multi-step top-k", "multistep"), | |
| ("Utility-aware multi-step", "utility"), | |
| ) | |
| def _percentile(values: list[float], q: float) -> float: | |
| if not values: | |
| return 0.0 | |
| ordered = sorted(float(value) for value in values) | |
| if len(ordered) == 1: | |
| return ordered[0] | |
| pos = min(max(float(q), 0.0), 1.0) * (len(ordered) - 1) | |
| lo = int(math.floor(pos)) | |
| hi = int(math.ceil(pos)) | |
| if lo == hi: | |
| return ordered[lo] | |
| frac = pos - lo | |
| return ordered[lo] * (1.0 - frac) + ordered[hi] * frac | |
| def _bootstrap_mean_ci(values: list[float], samples: int, seed: int) -> tuple[float, float]: | |
| if not values: | |
| return (0.0, 0.0) | |
| if len(values) == 1: | |
| return (values[0], values[0]) | |
| rng = random.Random(seed) | |
| n = len(values) | |
| draws = [] | |
| for _ in range(max(100, int(samples))): | |
| draws.append(mean(values[rng.randrange(n)] for _ in range(n))) | |
| return _percentile(draws, 0.025), _percentile(draws, 0.975) | |
| def _summary_row(label: str, result: dict[str, Any]) -> dict[str, Any]: | |
| return { | |
| "label": label, | |
| "p95_ttft_ms": float(result["latency"]["step_ttft_ms"]["p95"]), | |
| "p95_workflow_e2e_ms": float(result["latency"]["workflow_e2e_ms"]["p95"]), | |
| "workflow_throughput_rps": float(result["summary"]["workflow_throughput_rps"]), | |
| "completion_rate": float(result["summary"]["workflow_completion_rate"]), | |
| "prefix_hit_rate": float(result["resource"]["prefix_hit_rate"]), | |
| "prefetch_utilization": float(result["resource"].get("prefetch_utilization", 0.0)), | |
| "future_role_recall_at_k": float(result["resource"].get("forecast_recall", 0.0)), | |
| "mean_hbm_gb": float(result["resource"]["mean_prefix_hbm_gb"]), | |
| "unused_prefetch_gb": float(result["resource"].get("unused_prefetch_gb", 0.0)), | |
| "pressure_evictions": int(result["resource"].get("pressure_evictions", 0)), | |
| "saved_prefill_tokens": int(result["resource"].get("prefill_tokens_saved", 0)), | |
| } | |
| def _dominates(a: dict[str, Any], b: dict[str, Any]) -> bool: | |
| # The robust frontier deliberately treats speculative traffic and HBM as | |
| # first-class resources rather than ranking on latency alone. | |
| a_obj = (float(a["p95_ttft_ms"]), float(a["unused_prefetch_gb"]), float(a["mean_hbm_gb"])) | |
| b_obj = (float(b["p95_ttft_ms"]), float(b["unused_prefetch_gb"]), float(b["mean_hbm_gb"])) | |
| return all(x <= y + 1e-12 for x, y in zip(a_obj, b_obj, strict=True)) and any( | |
| x < y - 1e-12 for x, y in zip(a_obj, b_obj, strict=True) | |
| ) | |
| def _pareto_labels(rows: list[dict[str, Any]]) -> set[str]: | |
| labels: set[str] = set() | |
| for candidate in rows: | |
| if not any(_dominates(other, candidate) for other in rows if other is not candidate): | |
| labels.add(str(candidate["label"])) | |
| return labels | |
| def _offline_constrained_oracle( | |
| base: ExecutionLearningConfig, | |
| workflows: list[Any], | |
| deployable_results: dict[str, dict[str, Any]], | |
| ) -> dict[str, Any]: | |
| """Return a bounded full-trace information upper bound. | |
| This is intentionally *not* described as a globally optimal cache controller. | |
| It is an exhaustive oracle over a declared candidate family: all deployable | |
| policies already evaluated plus clairvoyant future-set plans for horizons | |
| 1..5 and top-k 1..3. Every candidate uses the same cache budget, transfer | |
| bandwidth, model, device profile, and exact realized workflow trace. | |
| """ | |
| pool: list[dict[str, Any]] = [] | |
| for label, result in deployable_results.items(): | |
| pool.append( | |
| { | |
| "label": label, | |
| "kind": "deployable", | |
| "config": { | |
| "policy": result["config"]["prefetch_policy"], | |
| "horizon": result["config"].get("forecast_horizon", base.forecast_horizon), | |
| "top_k": result["config"].get("prefetch_top_k", base.prefetch_top_k), | |
| }, | |
| "result": result, | |
| "row": _summary_row(label, result), | |
| } | |
| ) | |
| for horizon in range(1, 6): | |
| for top_k in range(1, 4): | |
| cfg = ExecutionLearningConfig.from_dict(base.to_dict()) | |
| cfg.prefetch_policy = "oracle_horizon" | |
| cfg.forecast_horizon = horizon | |
| cfg.prefetch_top_k = top_k | |
| result = run_execution_learning(cfg.to_dict(), workflows) | |
| label = f"clairvoyant H{horizon}/K{top_k}" | |
| pool.append( | |
| { | |
| "label": label, | |
| "kind": "clairvoyant", | |
| "config": {"policy": "oracle_horizon", "horizon": horizon, "top_k": top_k}, | |
| "result": result, | |
| "row": _summary_row(label, result), | |
| } | |
| ) | |
| feasible = [item for item in pool if item["row"]["completion_rate"] >= 1.0 - 1e-12] | |
| if not feasible: | |
| feasible = pool | |
| winner = min( | |
| feasible, | |
| key=lambda item: ( | |
| item["row"]["p95_ttft_ms"], | |
| item["row"]["p95_workflow_e2e_ms"], | |
| item["row"]["unused_prefetch_gb"], | |
| item["row"]["mean_hbm_gb"], | |
| ), | |
| ) | |
| return { | |
| "label": winner["label"], | |
| "kind": winner["kind"], | |
| "config": winner["config"], | |
| "metrics": winner["row"], | |
| "candidate_count": len(pool), | |
| "definition": "bounded-full-trace-serving-oracle", | |
| "note": ( | |
| "Exhaustive upper bound over the declared candidate family, including clairvoyant future-set plans. " | |
| "It uses the complete trace for policy selection and future-role actions, but is not a proof of global optimality." | |
| ), | |
| } | |
| def repeated_seed_policy_study( | |
| config: dict[str, Any], | |
| repetitions: int = 12, | |
| bootstrap_samples: int = 600, | |
| ) -> dict[str, Any]: | |
| base = ExecutionLearningConfig.from_dict(config) | |
| repetitions = max(4, min(int(repetitions), 24)) | |
| bootstrap_samples = max(100, min(int(bootstrap_samples), 4000)) | |
| seed_runs: list[dict[str, Any]] = [] | |
| per_policy: dict[str, list[dict[str, Any]]] = {label: [] for label, _ in POLICIES} | |
| for rep in range(repetitions): | |
| seed = base.seed + rep * 1009 | |
| seeded = ExecutionLearningConfig.from_dict(base.to_dict()) | |
| seeded.seed = seed | |
| workflows = generate_workflows(seeded) | |
| deployable_results: dict[str, dict[str, Any]] = {} | |
| rows: list[dict[str, Any]] = [] | |
| for label, policy in POLICIES: | |
| cfg = ExecutionLearningConfig.from_dict(seeded.to_dict()) | |
| cfg.prefetch_policy = policy | |
| result = run_execution_learning(cfg.to_dict(), workflows) | |
| deployable_results[label] = result | |
| row = _summary_row(label, result) | |
| rows.append(row) | |
| per_policy[label].append(row) | |
| pareto = _pareto_labels(rows) | |
| nominal_winner = min(rows, key=lambda row: (row["p95_ttft_ms"], row["unused_prefetch_gb"]))["label"] | |
| oracle = _offline_constrained_oracle(seeded, workflows, deployable_results) | |
| oracle_ttft = float(oracle["metrics"]["p95_ttft_ms"]) | |
| for row in rows: | |
| row["pareto"] = row["label"] in pareto | |
| row["ttft_winner"] = row["label"] == nominal_winner | |
| row["oracle_regret_ms"] = max(0.0, float(row["p95_ttft_ms"]) - oracle_ttft) | |
| row["oracle_regret_pct"] = ( | |
| row["oracle_regret_ms"] / oracle_ttft * 100.0 if oracle_ttft > 1e-12 else 0.0 | |
| ) | |
| seed_runs.append({"rep": rep + 1, "seed": seed, "rows": rows, "oracle": oracle}) | |
| baseline_label = "Top-1 decayed" | |
| baseline_ttfts = [float(row["p95_ttft_ms"]) for row in per_policy[baseline_label]] | |
| summaries: list[dict[str, Any]] = [] | |
| for index, (label, _) in enumerate(POLICIES): | |
| rows = per_policy[label] | |
| ttfts = [float(row["p95_ttft_ms"]) for row in rows] | |
| ci_low, ci_high = _bootstrap_mean_ci(ttfts, bootstrap_samples, base.seed ^ (index + 1) * 7919) | |
| paired_deltas = [ttft - base_ttft for ttft, base_ttft in zip(ttfts, baseline_ttfts, strict=True)] | |
| delta_low, delta_high = _bootstrap_mean_ci( | |
| paired_deltas, bootstrap_samples, base.seed ^ (index + 1) * 104729 | |
| ) | |
| seed_rows = [ | |
| next(row for row in seed_run["rows"] if row["label"] == label) | |
| for seed_run in seed_runs | |
| ] | |
| regrets = [float(row["oracle_regret_ms"]) for row in seed_rows] | |
| regret_low, regret_high = _bootstrap_mean_ci( | |
| regrets, bootstrap_samples, base.seed ^ (index + 1) * 15485863 | |
| ) | |
| summaries.append( | |
| { | |
| "label": label, | |
| "mean_ttft_ms": mean(ttfts), | |
| "median_ttft_ms": median(ttfts), | |
| "ttft_ci95_low_ms": ci_low, | |
| "ttft_ci95_high_ms": ci_high, | |
| "paired_delta_vs_top1_mean_ms": mean(paired_deltas), | |
| "paired_delta_ci95_low_ms": delta_low, | |
| "paired_delta_ci95_high_ms": delta_high, | |
| "ttft_win_rate": mean(1.0 if row["ttft_winner"] else 0.0 for row in seed_rows), | |
| "pareto_stability": mean(1.0 if row["pareto"] else 0.0 for row in seed_rows), | |
| "median_oracle_regret_ms": median(regrets), | |
| "mean_oracle_regret_ms": mean(regrets), | |
| "oracle_regret_ci95_low_ms": regret_low, | |
| "oracle_regret_ci95_high_ms": regret_high, | |
| "worst_seed_ttft_ms": max(ttfts), | |
| "mean_unused_prefetch_gb": mean(float(row["unused_prefetch_gb"]) for row in rows), | |
| "mean_hbm_gb": mean(float(row["mean_hbm_gb"]) for row in rows), | |
| "mean_completion_rate": mean(float(row["completion_rate"]) for row in rows), | |
| } | |
| ) | |
| ranked = sorted( | |
| summaries, | |
| key=lambda row: ( | |
| -float(row["ttft_win_rate"]), | |
| float(row["median_ttft_ms"]), | |
| float(row["median_oracle_regret_ms"]), | |
| -float(row["pareto_stability"]), | |
| float(row["mean_unused_prefetch_gb"]), | |
| ), | |
| ) | |
| rank_map = {row["label"]: rank + 1 for rank, row in enumerate(ranked)} | |
| for row in summaries: | |
| row["robust_rank"] = rank_map[row["label"]] | |
| nominal_first_seed = min( | |
| seed_runs[0]["rows"], key=lambda row: (row["p95_ttft_ms"], row["unused_prefetch_gb"]) | |
| )["label"] | |
| oracle_ttfts = [float(seed_run["oracle"]["metrics"]["p95_ttft_ms"]) for seed_run in seed_runs] | |
| oracle_labels: dict[str, int] = {} | |
| for seed_run in seed_runs: | |
| label = str(seed_run["oracle"]["label"]) | |
| oracle_labels[label] = oracle_labels.get(label, 0) + 1 | |
| return { | |
| "study": "repeated-seed-policy-consolidation", | |
| "protocol": "matched-seeds-bootstrap-and-bounded-offline-oracle", | |
| "repetitions": repetitions, | |
| "bootstrap_samples": bootstrap_samples, | |
| "config": base.to_dict(), | |
| "policies": summaries, | |
| "seed_runs": seed_runs, | |
| "nominal_winner_first_seed": nominal_first_seed, | |
| "robust_winner": ranked[0]["label"] if ranked else None, | |
| "oracle": { | |
| "definition": "bounded-full-trace-serving-oracle", | |
| "candidate_count_per_seed": seed_runs[0]["oracle"]["candidate_count"] if seed_runs else 0, | |
| "median_ttft_ms": median(oracle_ttfts) if oracle_ttfts else 0.0, | |
| "mean_ttft_ms": mean(oracle_ttfts) if oracle_ttfts else 0.0, | |
| "selected_plan_frequency": oracle_labels, | |
| "note": ( | |
| "The oracle exhaustively selects among the evaluated deployable policies plus clairvoyant future-set " | |
| "plans over H=1..5 and K=1..3 on each complete trace, under the same cache/bandwidth constraints. " | |
| "It is a bounded information upper bound, not a proof of globally optimal action scheduling." | |
| ), | |
| }, | |
| "pareto_objectives": ["p95_step_ttft_ms", "unused_prefetch_gb", "mean_hbm_gb"], | |
| "note": ( | |
| "Robust rank prioritizes how often a policy wins TTFT across matched seeds, then median TTFT and regret to " | |
| "the bounded offline oracle. Pareto stability is reported separately rather than silently overriding latency. " | |
| "Bootstrap intervals quantify seed uncertainty, not real-hardware error." | |
| ), | |
| } | |