376 lines
13 KiB
Python
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()
|