feat: add paired DeepSeek routing length control

This commit is contained in:
wuyang
2026-07-29 16:17:18 +08:00
parent 059458f8e3
commit 9ca08501a8
15 changed files with 318874 additions and 69 deletions
+53
View File
@@ -100,3 +100,56 @@ The committed independent rerun is byte-exact. Both JSON files have SHA-256:
See `research/DEEPSEEK_ROUTING_CORPUS_AUDIT.md` for corpus revisions and hashes,
metric definitions, interval semantics, results, and claim boundaries.
## Paired length-control cohort
The corpus runner can also select one fixed cohort by untruncated source length
and execute nested prefixes. The committed 16-token and 24-token traces use the
same 128 source prompts, all selected from records with at least 24 DeepSeek
tokens:
```bash
common_args=(
--per-domain 32
--batch-size 16
--bootstrap 2000
--sample-salt llm-atlas-deepseek-routing-length-control-v1
--eligibility-min-tokens 24
)
python -B experiments/deepseek/v2_lite_routing_corpus.py \
...source arguments... \
"${common_args[@]}" \
--max-tokens 16 \
--output src/data/deepseek-v2-lite-routing-matched16.json
python -B experiments/deepseek/v2_lite_routing_corpus.py \
...source arguments... \
"${common_args[@]}" \
--max-tokens 24 \
--output src/data/deepseek-v2-lite-routing-matched24.json
```
The two real traces add 184,320 top-6 route selections. Their independent
reruns are byte-exact:
```text
matched-16 f8d437d5379ffb41ac8dca5a8e97c0f44ba10ce7b63f95d7be0b7c88ac0baebd
matched-24 bed54835ad243ca2ab46bf9574e53137e6c2c0e19267f581719b5f2f65546436
```
Use the paired comparison runner to resample identical prompt indices in the
short and long traces:
```bash
python -B experiments/deepseek/compare_routing_length_control.py \
--short src/data/deepseek-v2-lite-routing-matched16.json \
--long src/data/deepseek-v2-lite-routing-matched24.json \
--output src/data/deepseek-v2-lite-routing-length-sensitivity.json \
--bootstrap 2000 \
--seed 20260729
```
See `research/DEEPSEEK_ROUTING_LENGTH_CONTROL_AUDIT.md` for the sampling bias
audit, paired CV/JSD deltas, total-variation accounting, and interpretation
boundaries.
@@ -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()
+53 -11
View File
@@ -69,6 +69,8 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--layers", type=int, default=7)
parser.add_argument("--bootstrap", type=int, default=2000)
parser.add_argument("--seed", type=int, default=20260729)
parser.add_argument("--sample-salt", default=SAMPLE_SALT)
parser.add_argument("--eligibility-min-tokens", type=int, default=None)
parser.add_argument("--device", default="cuda")
parser.add_argument("--captured-at", default=None)
return parser.parse_args()
@@ -175,6 +177,8 @@ def select_corpus(
candidates: dict[str, list[dict[str, str]]],
per_domain: int,
max_tokens: int,
sample_salt: str,
eligibility_min_tokens: int | None,
) -> tuple[list[dict[str, Any]], dict[str, dict[str, int]]]:
selected: list[dict[str, Any]] = []
counts: dict[str, dict[str, int]] = {}
@@ -183,21 +187,31 @@ def select_corpus(
texts = [row["text"] for row in rows]
encoded: list[list[int]] = []
for start in range(0, len(texts), 512):
result = tokenizer(
texts[start : start + 512],
add_special_tokens=True,
truncation=True,
max_length=max_tokens,
padding=False,
)
tokenizer_kwargs: dict[str, Any] = {
"add_special_tokens": True,
"padding": False,
}
if eligibility_min_tokens is None:
tokenizer_kwargs.update(
{
"truncation": True,
"max_length": max_tokens,
}
)
else:
tokenizer_kwargs["truncation"] = False
result = tokenizer(texts[start : start + 512], **tokenizer_kwargs)
encoded.extend(result.input_ids)
eligible = []
for row, token_ids in zip(rows, encoded, strict=True):
if len(token_ids) < MIN_TOKENS[domain]:
minimum = eligibility_min_tokens or MIN_TOKENS[domain]
if len(token_ids) < minimum:
continue
source_tokens = len(token_ids)
token_ids = token_ids[:max_tokens]
rank = hashlib.sha256(
f"{SAMPLE_SALT}|{domain}|{row['id']}".encode()
f"{sample_salt}|{domain}|{row['id']}".encode()
).hexdigest()
eligible.append(
{
@@ -209,6 +223,11 @@ def select_corpus(
"characters": len(row["text"]),
"token_ids": token_ids,
"tokens": len(token_ids),
**(
{"source_tokens": source_tokens}
if eligibility_min_tokens is not None
else {}
),
"selection_rank": rank,
}
)
@@ -509,6 +528,10 @@ def main() -> None:
raise ValueError("per-domain sample must be at least two")
if args.bootstrap < 100:
raise ValueError("bootstrap replicates must be at least 100")
if args.eligibility_min_tokens is not None and args.eligibility_min_tokens < 2:
raise ValueError("eligibility minimum must be at least two tokens")
if not args.sample_salt.strip():
raise ValueError("sample salt must not be empty")
torch.manual_seed(args.seed)
torch.cuda.manual_seed_all(args.seed)
@@ -532,6 +555,8 @@ def main() -> None:
candidates,
args.per_domain,
args.max_tokens,
args.sample_salt,
args.eligibility_min_tokens,
)
batches = make_batches(samples, args.batch_size, tokenizer.pad_token_id)
device = torch.device(args.device)
@@ -658,6 +683,11 @@ def main() -> None:
"text_sha256": sample["text_sha256"],
"characters": sample["characters"],
"tokens": sample["tokens"],
**(
{"source_tokens": sample["source_tokens"]}
if "source_tokens" in sample
else {}
),
}
for sample in samples
]
@@ -754,11 +784,23 @@ def main() -> None:
"corpus_contract": {
"domains": list(DOMAIN_ORDER),
"domain_labels": DOMAIN_LABELS,
"sample_salt": SAMPLE_SALT,
"sample_salt": args.sample_salt,
"selection": "ascending SHA256(salt|domain|source_id), then source_id",
"per_domain": args.per_domain,
"max_tokens": args.max_tokens,
"minimum_tokens": MIN_TOKENS,
"minimum_tokens": (
{domain: args.eligibility_min_tokens for domain in DOMAIN_ORDER}
if args.eligibility_min_tokens is not None
else MIN_TOKENS
),
**(
{
"uniform_eligibility_min_tokens": args.eligibility_min_tokens,
"matched_length_control": True,
}
if args.eligibility_min_tokens is not None
else {}
),
"special_tokens": True,
"truncation": "right",
"counts": corpus_counts,