feat: add paired DeepSeek routing length control
This commit is contained in:
@@ -0,0 +1,375 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Compare two routing traces that use the same prompts at different lengths.
|
||||
|
||||
The comparison is paired at prompt level: every bootstrap replicate draws the
|
||||
same source-prompt indices for the short and long trace. This isolates prefix
|
||||
length within the fixed matched cohort more cleanly than comparing independent
|
||||
confidence intervals.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
from datetime import datetime, timezone
|
||||
from itertools import combinations
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--short", type=Path, required=True)
|
||||
parser.add_argument("--long", type=Path, required=True)
|
||||
parser.add_argument("--output", type=Path, required=True)
|
||||
parser.add_argument("--bootstrap", type=int, default=2000)
|
||||
parser.add_argument("--seed", type=int, default=20260729)
|
||||
parser.add_argument("--captured-at", default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def sha256(path: Path) -> str:
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as handle:
|
||||
for block in iter(lambda: handle.read(1024 * 1024), b""):
|
||||
digest.update(block)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def scoped_seed(seed: int, scope: str) -> int:
|
||||
payload = f"{seed}:{scope}".encode()
|
||||
return int.from_bytes(hashlib.sha256(payload).digest()[:8], "big")
|
||||
|
||||
|
||||
def interval(values: np.ndarray) -> list[float]:
|
||||
low, high = np.quantile(values, [0.025, 0.975], axis=0)
|
||||
return [float(low), float(high)]
|
||||
|
||||
|
||||
def distributions(loads: np.ndarray, mode: str, sampled: np.ndarray | None = None) -> np.ndarray:
|
||||
values = loads if sampled is None else loads[sampled]
|
||||
if values.ndim == 2:
|
||||
values = values[None, ...]
|
||||
if mode == "token_weighted":
|
||||
result = values.sum(axis=1, dtype=np.float64)
|
||||
elif mode == "prompt_balanced":
|
||||
normalized = values / values.sum(axis=2, keepdims=True)
|
||||
result = normalized.mean(axis=1, dtype=np.float64)
|
||||
else:
|
||||
raise ValueError(mode)
|
||||
return result / result.sum(axis=1, keepdims=True)
|
||||
|
||||
|
||||
def metric_vector(values: np.ndarray) -> dict[str, np.ndarray]:
|
||||
values = np.atleast_2d(values).astype(np.float64, copy=False)
|
||||
expert_count = values.shape[1]
|
||||
mean = values.mean(axis=1)
|
||||
ordered = np.sort(values, axis=1)
|
||||
indices = np.arange(1, expert_count + 1, dtype=np.float64)
|
||||
gini = (
|
||||
((2 * indices - expert_count - 1) * ordered).sum(axis=1)
|
||||
/ (expert_count * ordered.sum(axis=1))
|
||||
)
|
||||
logs = np.zeros_like(values)
|
||||
np.log(values, out=logs, where=values > 0)
|
||||
entropy = -(values * logs).sum(axis=1)
|
||||
return {
|
||||
"cv": values.std(axis=1) / mean,
|
||||
"gini": gini,
|
||||
"effective_experts": np.exp(entropy),
|
||||
"top_expert_share": values.max(axis=1),
|
||||
"used_experts": (values > 0).sum(axis=1).astype(np.float64),
|
||||
}
|
||||
|
||||
|
||||
def js_divergence(left: np.ndarray, right: np.ndarray) -> np.ndarray:
|
||||
left = np.atleast_2d(left).astype(np.float64, copy=False)
|
||||
right = np.atleast_2d(right).astype(np.float64, copy=False)
|
||||
midpoint = (left + right) / 2
|
||||
left_ratio = np.ones_like(left)
|
||||
right_ratio = np.ones_like(right)
|
||||
np.divide(left, midpoint, out=left_ratio, where=left > 0)
|
||||
np.divide(right, midpoint, out=right_ratio, where=right > 0)
|
||||
left_log = np.zeros_like(left)
|
||||
right_log = np.zeros_like(right)
|
||||
np.log(left_ratio, out=left_log, where=left > 0)
|
||||
np.log(right_ratio, out=right_log, where=right > 0)
|
||||
return 0.5 * (
|
||||
(left * left_log).sum(axis=1)
|
||||
+ (right * right_log).sum(axis=1)
|
||||
)
|
||||
|
||||
|
||||
def prompt_identity(data: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
return [
|
||||
{
|
||||
key: row[key]
|
||||
for key in (
|
||||
"id",
|
||||
"domain",
|
||||
"within_domain_index",
|
||||
"selection_rank",
|
||||
"text_sha256",
|
||||
"characters",
|
||||
"source_tokens",
|
||||
)
|
||||
}
|
||||
for row in data["corpus_contract"]["selected"]
|
||||
]
|
||||
|
||||
|
||||
def load_by_domain(
|
||||
layer: dict[str, Any],
|
||||
domains: list[str],
|
||||
) -> tuple[dict[str, np.ndarray], dict[str, list[str]]]:
|
||||
loads: dict[str, np.ndarray] = {}
|
||||
ids: dict[str, list[str]] = {}
|
||||
for domain in domains:
|
||||
rows = sorted(
|
||||
(row for row in layer["prompts"] if row["domain"] == domain),
|
||||
key=lambda row: row["id"],
|
||||
)
|
||||
ids[domain] = [row["id"] for row in rows]
|
||||
loads[domain] = np.asarray([row["load"] for row in rows], dtype=np.int64)
|
||||
return loads, ids
|
||||
|
||||
|
||||
def compare_domain(
|
||||
short: np.ndarray,
|
||||
long: np.ndarray,
|
||||
mode: str,
|
||||
bootstrap: int,
|
||||
seed: int,
|
||||
scope: str,
|
||||
) -> dict[str, Any]:
|
||||
if short.shape != long.shape:
|
||||
raise ValueError(f"paired loads differ in shape: {short.shape} vs {long.shape}")
|
||||
prompt_count = short.shape[0]
|
||||
rng = np.random.default_rng(scoped_seed(seed, scope))
|
||||
sampled = rng.integers(
|
||||
0,
|
||||
prompt_count,
|
||||
size=(bootstrap, prompt_count),
|
||||
endpoint=False,
|
||||
)
|
||||
short_point = distributions(short, mode)[0]
|
||||
long_point = distributions(long, mode)[0]
|
||||
short_boot = distributions(short, mode, sampled)
|
||||
long_boot = distributions(long, mode, sampled)
|
||||
short_metrics = metric_vector(short_point)
|
||||
long_metrics = metric_vector(long_point)
|
||||
short_boot_metrics = metric_vector(short_boot)
|
||||
long_boot_metrics = metric_vector(long_boot)
|
||||
metrics = {}
|
||||
for name in short_metrics:
|
||||
delta = long_boot_metrics[name] - short_boot_metrics[name]
|
||||
metrics[name] = {
|
||||
"short": float(short_metrics[name][0]),
|
||||
"long": float(long_metrics[name][0]),
|
||||
"delta_long_minus_short": float(
|
||||
long_metrics[name][0] - short_metrics[name][0]
|
||||
),
|
||||
"delta_ci95": interval(delta),
|
||||
}
|
||||
total_variation_boot = 0.5 * np.abs(long_boot - short_boot).sum(axis=1)
|
||||
return {
|
||||
"metrics": metrics,
|
||||
"total_variation": {
|
||||
"point": float(0.5 * np.abs(long_point - short_point).sum()),
|
||||
"ci95": interval(total_variation_boot),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def compare_pair(
|
||||
short_left: np.ndarray,
|
||||
short_right: np.ndarray,
|
||||
long_left: np.ndarray,
|
||||
long_right: np.ndarray,
|
||||
mode: str,
|
||||
bootstrap: int,
|
||||
seed: int,
|
||||
scope: str,
|
||||
) -> dict[str, Any]:
|
||||
rng = np.random.default_rng(scoped_seed(seed, scope))
|
||||
left_indices = rng.integers(
|
||||
0,
|
||||
short_left.shape[0],
|
||||
size=(bootstrap, short_left.shape[0]),
|
||||
endpoint=False,
|
||||
)
|
||||
right_indices = rng.integers(
|
||||
0,
|
||||
short_right.shape[0],
|
||||
size=(bootstrap, short_right.shape[0]),
|
||||
endpoint=False,
|
||||
)
|
||||
short_point = js_divergence(
|
||||
distributions(short_left, mode),
|
||||
distributions(short_right, mode),
|
||||
)[0]
|
||||
long_point = js_divergence(
|
||||
distributions(long_left, mode),
|
||||
distributions(long_right, mode),
|
||||
)[0]
|
||||
short_boot = js_divergence(
|
||||
distributions(short_left, mode, left_indices),
|
||||
distributions(short_right, mode, right_indices),
|
||||
)
|
||||
long_boot = js_divergence(
|
||||
distributions(long_left, mode, left_indices),
|
||||
distributions(long_right, mode, right_indices),
|
||||
)
|
||||
return {
|
||||
"short": float(short_point),
|
||||
"long": float(long_point),
|
||||
"delta_long_minus_short": float(long_point - short_point),
|
||||
"delta_ci95": interval(long_boot - short_boot),
|
||||
"unit": "nats",
|
||||
"upper_bound": math.log(2),
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
if args.bootstrap < 100:
|
||||
raise ValueError("bootstrap replicates must be at least 100")
|
||||
short = json.loads(args.short.read_text())
|
||||
long = json.loads(args.long.read_text())
|
||||
domains = short["corpus_contract"]["domains"]
|
||||
if domains != long["corpus_contract"]["domains"]:
|
||||
raise ValueError("domain order differs")
|
||||
if prompt_identity(short) != prompt_identity(long):
|
||||
raise ValueError("short and long traces do not use the same source prompts")
|
||||
for key in ("model", "corpora"):
|
||||
if short["provenance"][key] != long["provenance"][key]:
|
||||
raise ValueError(f"provenance differs: {key}")
|
||||
if short["configuration"] != long["configuration"]:
|
||||
raise ValueError("model configuration differs")
|
||||
if short["corpus_contract"]["sample_salt"] != long["corpus_contract"]["sample_salt"]:
|
||||
raise ValueError("sample salt differs")
|
||||
|
||||
short_tokens = short["corpus_contract"]["max_tokens"]
|
||||
long_tokens = long["corpus_contract"]["max_tokens"]
|
||||
if short_tokens >= long_tokens:
|
||||
raise ValueError("short max tokens must be less than long max tokens")
|
||||
expected_short = len(domains) * short["corpus_contract"]["per_domain"] * short_tokens
|
||||
expected_long = len(domains) * long["corpus_contract"]["per_domain"] * long_tokens
|
||||
if short["inference_contract"]["valid_tokens"] != expected_short:
|
||||
raise ValueError("short trace is not exactly length-controlled")
|
||||
if long["inference_contract"]["valid_tokens"] != expected_long:
|
||||
raise ValueError("long trace is not exactly length-controlled")
|
||||
|
||||
layers = []
|
||||
for short_layer, long_layer in zip(
|
||||
short["layers"][1:],
|
||||
long["layers"][1:],
|
||||
strict=True,
|
||||
):
|
||||
if short_layer["layer"] != long_layer["layer"]:
|
||||
raise ValueError("layer order differs")
|
||||
layer_index = short_layer["layer"]
|
||||
short_loads, short_ids = load_by_domain(short_layer, domains)
|
||||
long_loads, long_ids = load_by_domain(long_layer, domains)
|
||||
if short_ids != long_ids:
|
||||
raise ValueError(f"prompt IDs differ at layer {layer_index}")
|
||||
modes = {}
|
||||
for mode in ("token_weighted", "prompt_balanced"):
|
||||
domain_results = {
|
||||
domain: compare_domain(
|
||||
short_loads[domain],
|
||||
long_loads[domain],
|
||||
mode,
|
||||
args.bootstrap,
|
||||
args.seed,
|
||||
f"layer={layer_index}|mode={mode}|domain={domain}",
|
||||
)
|
||||
for domain in domains
|
||||
}
|
||||
pairs = []
|
||||
for left, right in combinations(domains, 2):
|
||||
pairs.append(
|
||||
{
|
||||
"left": left,
|
||||
"right": right,
|
||||
"js_divergence": compare_pair(
|
||||
short_loads[left],
|
||||
short_loads[right],
|
||||
long_loads[left],
|
||||
long_loads[right],
|
||||
mode,
|
||||
args.bootstrap,
|
||||
args.seed,
|
||||
f"layer={layer_index}|mode={mode}|pair={left}:{right}",
|
||||
),
|
||||
}
|
||||
)
|
||||
modes[mode] = {"domains": domain_results, "pairs": pairs}
|
||||
layers.append({"layer": layer_index, "modes": modes})
|
||||
|
||||
captured_at = args.captured_at or datetime.now(timezone.utc).isoformat()
|
||||
result = {
|
||||
"schema_version": 1,
|
||||
"captured_at": captured_at,
|
||||
"evidence_identity": "S / paired prompt-level bootstrap over two real routing traces",
|
||||
"boundary": {
|
||||
"same_source_prompts": True,
|
||||
"nested_prefixes": True,
|
||||
"causal_claim_beyond_fixed_cohort": False,
|
||||
"expert_semantics_inferred": False,
|
||||
"hypothesis_test": False,
|
||||
},
|
||||
"inputs": {
|
||||
"short": {
|
||||
"path": args.short.name,
|
||||
"sha256": sha256(args.short),
|
||||
"tokens_per_prompt": short_tokens,
|
||||
"valid_tokens": short["inference_contract"]["valid_tokens"],
|
||||
"routes": short["inference_contract"]["routes_per_moe_layer"] * 6,
|
||||
},
|
||||
"long": {
|
||||
"path": args.long.name,
|
||||
"sha256": sha256(args.long),
|
||||
"tokens_per_prompt": long_tokens,
|
||||
"valid_tokens": long["inference_contract"]["valid_tokens"],
|
||||
"routes": long["inference_contract"]["routes_per_moe_layer"] * 6,
|
||||
},
|
||||
},
|
||||
"paired_contract": {
|
||||
"domains": domains,
|
||||
"prompts_per_domain": short["corpus_contract"]["per_domain"],
|
||||
"sample_salt": short["corpus_contract"]["sample_salt"],
|
||||
"eligibility_min_tokens": short["corpus_contract"][
|
||||
"uniform_eligibility_min_tokens"
|
||||
],
|
||||
"resampling_unit": "source prompt",
|
||||
"same_indices_for_short_and_long": True,
|
||||
"replicates": args.bootstrap,
|
||||
"seed": args.seed,
|
||||
"interval": "95% percentile interval of paired long-minus-short deltas",
|
||||
"modes": ["token_weighted", "prompt_balanced"],
|
||||
},
|
||||
"layers": layers,
|
||||
}
|
||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.output.write_text(json.dumps(result, ensure_ascii=False, indent=2) + "\n")
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"output": str(args.output),
|
||||
"layers": len(layers),
|
||||
"short_tokens": short_tokens,
|
||||
"long_tokens": long_tokens,
|
||||
"bootstrap": args.bootstrap,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user