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

862 lines
31 KiB
Python

#!/usr/bin/env python3
"""Run a fixed, multi-domain DeepSeek-V2-Lite routing corpus.
The probe executes the official embedding and decoder layers 0 through 6 from
the first official BF16 safetensors shard. Layers 1 through 6 are MoE layers.
Their real top-6 routed-expert IDs are aggregated at prompt level before any
statistics are computed.
Sampling is deterministic and source-addressable. Confidence intervals use a
prompt-level, within-domain bootstrap; tokens from the same prompt are never
treated as independent observations.
"""
from __future__ import annotations
import argparse
import gc
import gzip
import hashlib
import importlib.util
import json
import math
import platform
import re
import subprocess
import sys
import types
from datetime import datetime, timezone
from itertools import combinations
from pathlib import Path
from typing import Any
import numpy as np
import pyarrow.parquet as pq
import torch
import torch.nn.functional as F
from safetensors import safe_open
from transformers import AutoTokenizer
DOMAIN_ORDER = ("english", "chinese", "code", "math")
DOMAIN_LABELS = {
"english": "English encyclopedia",
"chinese": "中文新闻",
"code": "Python code",
"math": "Grade-school math",
}
MIN_TOKENS = {
"english": 24,
"chinese": 8,
"code": 24,
"math": 16,
}
SAMPLE_SALT = "llm-atlas-deepseek-routing-v1"
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--artifact-dir", type=Path, required=True)
parser.add_argument("--human-eval", type=Path, required=True)
parser.add_argument("--gsm8k", type=Path, required=True)
parser.add_argument("--tnews", type=Path, required=True)
parser.add_argument("--tnews-archive", type=Path, required=True)
parser.add_argument("--wikitext", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--per-domain", type=int, default=32)
parser.add_argument("--max-tokens", type=int, default=96)
parser.add_argument("--batch-size", type=int, default=16)
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()
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 text_sha256(value: str) -> str:
return hashlib.sha256(value.encode("utf-8")).hexdigest()
def git_revision(path: Path) -> str | None:
current = path.resolve()
for candidate in (current, *current.parents):
if (candidate / ".git").exists():
return subprocess.check_output(
["git", "-C", str(candidate), "rev-parse", "HEAD"],
text=True,
).strip()
return None
def scoped_seed(seed: int, scope: str) -> int:
payload = f"{seed}:{scope}".encode()
return int.from_bytes(hashlib.sha256(payload).digest()[:8], "big")
def load_official_modules(root: Path) -> tuple[Any, Any]:
package_name = "deepseek_v2_lite_official_routing_corpus"
package = types.ModuleType(package_name)
package.__path__ = [str(root)]
sys.modules[package_name] = package
loaded = {}
for leaf in ("configuration_deepseek", "modeling_deepseek"):
name = f"{package_name}.{leaf}"
spec = importlib.util.spec_from_file_location(name, root / f"{leaf}.py")
if spec is None or spec.loader is None:
raise RuntimeError(f"cannot load official module: {leaf}")
module = importlib.util.module_from_spec(spec)
sys.modules[name] = module
spec.loader.exec_module(module)
loaded[leaf] = module
return loaded["configuration_deepseek"], loaded["modeling_deepseek"]
def read_jsonl(path: Path) -> list[dict[str, Any]]:
with path.open(encoding="utf-8") as handle:
return [json.loads(line) for line in handle if line.strip()]
def load_candidates(args: argparse.Namespace) -> dict[str, list[dict[str, str]]]:
with gzip.open(args.human_eval, "rt", encoding="utf-8") as handle:
human_eval = [json.loads(line) for line in handle if line.strip()]
gsm8k = read_jsonl(args.gsm8k)
tnews = read_jsonl(args.tnews)
wikitext = pq.read_table(args.wikitext, columns=["text"]).column("text").to_pylist()
return {
"english": [
{
"id": f"wikitext2/raw-validation/{index:04d}",
"text": str(text).strip(),
}
for index, text in enumerate(wikitext)
if str(text).strip()
and not re.fullmatch(r"=+\s*.*?\s*=+", str(text).strip())
],
"chinese": [
{
"id": f"tnews/test/{row['id']}",
"text": str(row["sentence"]).strip(),
}
for row in tnews
if str(row.get("sentence", "")).strip()
],
"code": [
{
"id": str(row["task_id"]),
"text": str(row["prompt"]).rstrip(),
}
for row in human_eval
if str(row.get("prompt", "")).strip()
],
"math": [
{
"id": f"gsm8k/test/{index:04d}",
"text": str(row["question"]).strip(),
}
for index, row in enumerate(gsm8k)
if str(row.get("question", "")).strip()
],
}
def select_corpus(
tokenizer: Any,
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]] = {}
for domain in DOMAIN_ORDER:
rows = candidates[domain]
texts = [row["text"] for row in rows]
encoded: list[list[int]] = []
for start in range(0, len(texts), 512):
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):
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()
).hexdigest()
eligible.append(
{
"id": row["id"],
"domain": domain,
"label": DOMAIN_LABELS[domain],
"text": row["text"],
"text_sha256": text_sha256(row["text"]),
"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,
}
)
eligible.sort(key=lambda row: (row["selection_rank"], row["id"]))
if len(eligible) < per_domain:
raise RuntimeError(
f"{domain} has only {len(eligible)} eligible prompts; need {per_domain}"
)
domain_selection = eligible[:per_domain]
for within_domain_index, row in enumerate(domain_selection):
row["within_domain_index"] = within_domain_index
selected.extend(domain_selection)
counts[domain] = {
"candidate_records_after_text_filter": len(rows),
"eligible_records": len(eligible),
"selected_records": len(domain_selection),
"valid_tokens": sum(row["tokens"] for row in domain_selection),
}
return selected, counts
def make_batches(
samples: list[dict[str, Any]],
batch_size: int,
pad_token_id: int,
) -> list[dict[str, Any]]:
ordered = sorted(
samples,
key=lambda row: (row["tokens"], row["domain"], row["id"]),
)
batches = []
for start in range(0, len(ordered), batch_size):
rows = ordered[start : start + batch_size]
sequence = max(row["tokens"] for row in rows)
input_ids = torch.full(
(len(rows), sequence),
pad_token_id,
dtype=torch.long,
)
attention_mask = torch.zeros((len(rows), sequence), dtype=torch.long)
for index, row in enumerate(rows):
length = row["tokens"]
input_ids[index, :length] = torch.tensor(row["token_ids"])
attention_mask[index, :length] = 1
batches.append(
{
"samples": rows,
"input_ids": input_ids,
"attention_mask": attention_mask,
"padded_sequence": sequence,
}
)
return batches
def distribution(loads: np.ndarray, mode: str) -> np.ndarray:
if mode == "token_weighted":
values = loads.sum(axis=0, dtype=np.float64)
return values / values.sum()
if mode == "prompt_balanced":
prompt_distributions = loads / loads.sum(axis=1, keepdims=True)
values = prompt_distributions.mean(axis=0, dtype=np.float64)
return values / values.sum()
raise ValueError(mode)
def metric_vector(distributions: np.ndarray) -> dict[str, np.ndarray]:
values = np.atleast_2d(distributions).astype(np.float64, copy=False)
expert_count = values.shape[1]
means = values.mean(axis=1)
cv = values.std(axis=1) / means
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))
)
log_values = np.zeros_like(values)
np.log(values, out=log_values, where=values > 0)
entropy = -(values * log_values).sum(axis=1)
return {
"cv": cv,
"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 interval(values: np.ndarray) -> list[float]:
low, high = np.quantile(values, [0.025, 0.975], axis=0)
if np.ndim(low) == 0:
return [float(low), float(high)]
return np.stack([low, high], axis=-1).tolist()
def bootstrap_domain(
loads: np.ndarray,
mode: str,
replicates: int,
seed: int,
scope: str,
) -> dict[str, Any]:
point_distribution = distribution(loads, mode)
prompt_count = loads.shape[0]
rng = np.random.default_rng(scoped_seed(seed, scope))
sampled = rng.integers(
0,
prompt_count,
size=(replicates, prompt_count),
endpoint=False,
)
if mode == "token_weighted":
bootstrap_loads = loads[sampled].sum(axis=1, dtype=np.float64)
bootstrap_distributions = bootstrap_loads / bootstrap_loads.sum(
axis=1, keepdims=True
)
else:
prompt_distributions = loads / loads.sum(axis=1, keepdims=True)
bootstrap_distributions = prompt_distributions[sampled].mean(axis=1)
bootstrap_distributions /= bootstrap_distributions.sum(axis=1, keepdims=True)
point_metrics = metric_vector(point_distribution)
bootstrap_metrics = metric_vector(bootstrap_distributions)
metrics = {
name: {
"point": float(point_metrics[name][0]),
"ci95": interval(bootstrap_metrics[name]),
}
for name in point_metrics
}
order = np.argsort(point_distribution)[::-1][:8]
return {
"distribution": point_distribution.tolist(),
"expert_share_ci95": interval(bootstrap_distributions),
"metrics": metrics,
"top_experts": [
{
"expert": int(expert),
"share": float(point_distribution[expert]),
"ci95": interval(bootstrap_distributions[:, expert]),
}
for expert in order
],
}
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_log_ratio = np.zeros_like(left)
right_log_ratio = np.zeros_like(right)
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)
np.log(left_ratio, out=left_log_ratio, where=left > 0)
np.log(right_ratio, out=right_log_ratio, where=right > 0)
return 0.5 * (
(left * left_log_ratio).sum(axis=1)
+ (right * right_log_ratio).sum(axis=1)
)
def bootstrap_pair(
left: np.ndarray,
right: np.ndarray,
mode: str,
replicates: int,
seed: int,
scope: str,
) -> dict[str, Any]:
left_point = distribution(left, mode)
right_point = distribution(right, mode)
rng = np.random.default_rng(scoped_seed(seed, scope))
left_indices = rng.integers(
0, left.shape[0], size=(replicates, left.shape[0]), endpoint=False
)
right_indices = rng.integers(
0, right.shape[0], size=(replicates, right.shape[0]), endpoint=False
)
if mode == "token_weighted":
left_boot = left[left_indices].sum(axis=1, dtype=np.float64)
right_boot = right[right_indices].sum(axis=1, dtype=np.float64)
left_boot /= left_boot.sum(axis=1, keepdims=True)
right_boot /= right_boot.sum(axis=1, keepdims=True)
else:
left_prompt = left / left.sum(axis=1, keepdims=True)
right_prompt = right / right.sum(axis=1, keepdims=True)
left_boot = left_prompt[left_indices].mean(axis=1)
right_boot = right_prompt[right_indices].mean(axis=1)
left_boot /= left_boot.sum(axis=1, keepdims=True)
right_boot /= right_boot.sum(axis=1, keepdims=True)
bootstrap_jsd = js_divergence(left_boot, right_boot)
return {
"point": float(js_divergence(left_point, right_point)[0]),
"ci95": interval(bootstrap_jsd),
"unit": "nats",
"upper_bound": math.log(2),
}
def layer_statistics(
prompt_rows: list[dict[str, Any]],
replicates: int,
seed: int,
layer_index: int,
) -> dict[str, Any]:
loads_by_domain = {
domain: np.asarray(
[row["load"] for row in prompt_rows if row["domain"] == domain],
dtype=np.int64,
)
for domain in DOMAIN_ORDER
}
modes: dict[str, Any] = {}
for mode in ("token_weighted", "prompt_balanced"):
domains = {
domain: bootstrap_domain(
loads_by_domain[domain],
mode,
replicates,
seed,
f"layer={layer_index}|mode={mode}|domain={domain}",
)
for domain in DOMAIN_ORDER
}
pairs = []
for left, right in combinations(DOMAIN_ORDER, 2):
pairs.append(
{
"left": left,
"right": right,
"js_divergence": bootstrap_pair(
loads_by_domain[left],
loads_by_domain[right],
mode,
replicates,
seed,
f"layer={layer_index}|mode={mode}|pair={left}:{right}",
),
}
)
modes[mode] = {"domains": domains, "pairs": pairs}
return modes
def gpu_identity(device: torch.device) -> dict[str, Any] | None:
if device.type != "cuda":
return None
index = device.index or 0
properties = torch.cuda.get_device_properties(index)
driver = subprocess.check_output(
[
"nvidia-smi",
"--query-gpu=driver_version",
"--format=csv,noheader,nounits",
f"--id={index}",
],
text=True,
).strip()
return {
"name": properties.name,
"total_memory_bytes": properties.total_memory,
"compute_capability": list(torch.cuda.get_device_capability(index)),
"driver": driver,
}
def main() -> None:
args = parse_args()
root = args.artifact_dir.resolve()
shard = root / "model-00001-of-000004.safetensors"
required = [
root / "config.json",
root / "configuration_deepseek.py",
root / "modeling_deepseek.py",
root / "model.safetensors.index.json",
root / "tokenizer.json",
shard,
args.human_eval,
args.gsm8k,
args.tnews,
args.tnews_archive,
args.wikitext,
]
missing = [str(path) for path in required if not path.exists()]
if missing:
raise FileNotFoundError(f"missing artifacts: {missing}")
if args.device.startswith("cuda") and not torch.cuda.is_available():
raise RuntimeError("CUDA requested but unavailable")
if not 2 <= args.layers <= 7:
raise ValueError("need layer 0 plus at least one MoE layer; shard ends at layer 6")
if args.per_domain < 2:
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)
torch.backends.cuda.matmul.allow_tf32 = False
configuration, modeling = load_official_modules(root)
config = configuration.DeepseekV2Config.from_pretrained(root)
config._attn_implementation = "eager"
tokenizer = AutoTokenizer.from_pretrained(
root,
trust_remote_code=True,
local_files_only=True,
)
if tokenizer.pad_token_id is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "right"
candidates = load_candidates(args)
samples, corpus_counts = select_corpus(
tokenizer,
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)
with safe_open(shard, framework="pt", device="cpu") as handle:
embedding = handle.get_tensor("model.embed_tokens.weight")
for batch in batches:
batch["hidden"] = F.embedding(batch["input_ids"], embedding)
del embedding
layer_results = []
for layer_index in range(args.layers):
prefix = f"model.layers.{layer_index}."
with safe_open(shard, framework="pt", device="cpu") as handle:
state = {
key[len(prefix) :]: handle.get_tensor(key)
for key in handle.keys()
if key.startswith(prefix)
}
state_numel = sum(value.numel() for value in state.values())
state_bytes = sum(value.numel() * value.element_size() for value in state.values())
with torch.device("meta"):
layer = modeling.DeepseekV2DecoderLayer(config, layer_index)
layer.to_empty(device="cpu")
layer.load_state_dict(state, strict=True, assign=True)
del state
layer = layer.to(device=device, dtype=torch.bfloat16).eval()
prompt_rows = []
next_hidden = []
for batch in batches:
hidden = batch["hidden"].to(device)
attention_mask = batch["attention_mask"].to(device)
sequence = batch["padded_sequence"]
position_ids = torch.arange(sequence, device=device).unsqueeze(0)
causal_mask = modeling._prepare_4d_causal_attention_mask(
attention_mask,
batch["input_ids"].shape,
hidden,
0,
)
captures: dict[str, torch.Tensor] = {}
hook = None
if layer_index > 0:
def capture_gate(_module: Any, _inputs: Any, output: Any) -> None:
captures["topk_ids"] = output[0].detach().cpu()
hook = layer.mlp.gate.register_forward_hook(capture_gate)
with torch.inference_mode():
outputs = layer(
hidden,
attention_mask=causal_mask,
position_ids=position_ids,
past_key_value=None,
use_cache=False,
)
next_hidden.append(outputs[0].cpu())
if hook is not None:
hook.remove()
topk_ids = captures["topk_ids"].view(
len(batch["samples"]),
sequence,
config.num_experts_per_tok,
)
for sample_index, sample in enumerate(batch["samples"]):
length = sample["tokens"]
load = torch.bincount(
topk_ids[sample_index, :length].flatten(),
minlength=config.n_routed_experts,
)
prompt_rows.append(
{
"id": sample["id"],
"domain": sample["domain"],
"tokens": length,
"routes": int(load.sum()),
"load": load.tolist(),
}
)
del hidden, attention_mask, causal_mask, outputs
for batch, hidden in zip(batches, next_hidden, strict=True):
batch["hidden"] = hidden
result: dict[str, Any] = {
"layer": layer_index,
"ffn": "dense" if layer_index == 0 else "moe",
"state_numel": state_numel,
"state_bytes": state_bytes,
}
if layer_index > 0:
prompt_rows.sort(
key=lambda row: (
DOMAIN_ORDER.index(row["domain"]),
row["id"],
)
)
result["valid_tokens"] = sum(row["tokens"] for row in prompt_rows)
result["routes"] = sum(row["routes"] for row in prompt_rows)
result["prompts"] = prompt_rows
result["statistics"] = layer_statistics(
prompt_rows,
args.bootstrap,
args.seed,
layer_index,
)
layer_results.append(result)
del layer, next_hidden
gc.collect()
if device.type == "cuda":
torch.cuda.empty_cache()
captured_at = args.captured_at or datetime.now(timezone.utc).isoformat()
index = json.loads((root / "model.safetensors.index.json").read_text())
selected_identity = [
{
"id": sample["id"],
"domain": sample["domain"],
"within_domain_index": sample["within_domain_index"],
"selection_rank": sample["selection_rank"],
"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
]
result = {
"schema_version": 1,
"captured_at": captured_at,
"evidence_identity": "X / official BF16 weights, fixed public corpus, local truncated forward",
"boundary": {
"model": "DeepSeek-V2-Lite base",
"executed_layers": list(range(args.layers)),
"measured_moe_layers": list(range(1, args.layers)),
"total_model_layers": config.num_hidden_layers,
"full_model_generation": False,
"training_or_online_load": False,
"expert_semantics_inferred": False,
"population": (
f"{len(samples)} fixed public prompts across four domains; "
"not training data, online traffic, or a task-performance benchmark"
),
"code_execution": False,
"answers_used": False,
},
"provenance": {
"model": {
"huggingface_model": "deepseek-ai/DeepSeek-V2-Lite",
"huggingface_revision": "604d5664dddd88a0433dbae533b7fe9472482de0",
"sha256": {
"config": sha256(root / "config.json"),
"modeling_code": sha256(root / "modeling_deepseek.py"),
"tokenizer": sha256(root / "tokenizer.json"),
"index": sha256(root / "model.safetensors.index.json"),
"shard_1": sha256(shard),
},
"checkpoint_tensor_bytes": index["metadata"]["total_size"],
"shard_1_bytes": shard.stat().st_size,
},
"corpora": {
"english": {
"name": "WikiText-2 raw validation",
"url": "https://huggingface.co/datasets/Salesforce/wikitext",
"revision": "b08601e04326c79dfdd32d625aee71d232d685c3",
"file_sha256": sha256(args.wikitext),
"field_used": "text",
},
"chinese": {
"name": "CLUE TNEWS public test",
"url": "https://github.com/CLUEbenchmark/CLUE",
"download_url": "https://storage.googleapis.com/cluebenchmark/tasks/tnews_public.zip",
"revision": "9e61ddd3659ddb57ed82b4d0ba0a8613dfb55a2e",
"archive_sha256": sha256(args.tnews_archive),
"file_sha256": sha256(args.tnews),
"field_used": "sentence",
},
"code": {
"name": "OpenAI HumanEval",
"url": "https://github.com/openai/human-eval",
"revision": git_revision(args.human_eval),
"file_sha256": sha256(args.human_eval),
"field_used": "prompt",
},
"math": {
"name": "OpenAI GSM8K test",
"url": "https://github.com/openai/grade-school-math",
"revision": git_revision(args.gsm8k),
"file_sha256": sha256(args.gsm8k),
"field_used": "question",
},
},
},
"environment": {
"python": platform.python_version(),
"platform": platform.platform(),
"torch": torch.__version__,
"torch_cuda": torch.version.cuda,
"transformers": __import__("transformers").__version__,
"safetensors": __import__("safetensors").__version__,
"numpy": np.__version__,
"pyarrow": __import__("pyarrow").__version__,
"device": str(device),
"gpu": gpu_identity(device),
"matmul_allow_tf32": torch.backends.cuda.matmul.allow_tf32,
},
"configuration": {
"layers": config.num_hidden_layers,
"hidden": config.hidden_size,
"routed_experts": config.n_routed_experts,
"active_routed_experts": config.num_experts_per_tok,
"shared_experts": config.n_shared_experts,
"first_dense_layers": config.first_k_dense_replace,
"router_scoring": config.scoring_func,
"router_topk_method": config.topk_method,
"normalize_selected_weights": config.norm_topk_prob,
},
"corpus_contract": {
"domains": list(DOMAIN_ORDER),
"domain_labels": DOMAIN_LABELS,
"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": (
{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,
"selected": selected_identity,
},
"inference_contract": {
"batch_size": args.batch_size,
"batches": len(batches),
"batch_order": "ascending token count, then domain, then source id",
"attention": "official eager",
"dtype": "BF16",
"total_prompts": len(samples),
"valid_tokens": sum(sample["tokens"] for sample in samples),
"routes_per_moe_layer": sum(sample["tokens"] for sample in samples)
* config.num_experts_per_tok,
},
"statistical_contract": {
"resampling_unit": "prompt",
"strata": "domain",
"replicates": args.bootstrap,
"seed": args.seed,
"interval": "95% percentile bootstrap",
"modes": {
"token_weighted": "sum route counts, so longer prompts contribute more",
"prompt_balanced": "normalize each prompt first, then give every prompt equal weight",
},
"metrics": {
"cv": "population standard deviation across 64 expert shares divided by their mean",
"gini": "Gini coefficient across 64 expert shares",
"effective_experts": "exp(Shannon entropy) in nats",
"top_expert_share": "largest routed-expert share",
"used_experts": "experts with nonzero share in the resample",
"js_divergence": "Jensen-Shannon divergence in nats; bounded by ln(2)",
},
"multiple_comparison_correction": False,
},
"layers": layer_results,
}
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),
"prompts": len(samples),
"tokens": result["inference_contract"]["valid_tokens"],
"routes_per_moe_layer": result["inference_contract"][
"routes_per_moe_layer"
],
"measured_moe_layers": args.layers - 1,
},
ensure_ascii=False,
)
)
if __name__ == "__main__":
main()