Files
llm-atlas/experiments/deepseek/compare_routing_length_control.py

376 lines
13 KiB
Python

#!/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()