1121 lines
43 KiB
Python
1121 lines
43 KiB
Python
#!/usr/bin/env python3
|
|
"""Probe DeepSeek-V2-Lite routing sensitivity to its official chat template.
|
|
|
|
The experiment keeps one source-addressable public cohort fixed and renders
|
|
three inputs per prompt:
|
|
|
|
1. regular tokenizer input (BOS + content);
|
|
2. the official one-user chat template;
|
|
3. the same official template with the generation prompt appended.
|
|
|
|
All three variants execute together through official BF16 layers 0--6. The
|
|
output separates full-input loads from content-only loads, pairs bootstrap
|
|
indices by source prompt, aligns content tokens by character span, and checks
|
|
the causal invariant that appending a suffix cannot change a shared prefix.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import gc
|
|
import hashlib
|
|
import json
|
|
import math
|
|
import platform
|
|
from datetime import datetime, timezone
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import numpy as np
|
|
import torch
|
|
import torch.nn.functional as F
|
|
from safetensors import safe_open
|
|
from transformers import AutoTokenizer
|
|
|
|
from v2_lite_routing_corpus import (
|
|
DOMAIN_LABELS,
|
|
DOMAIN_ORDER,
|
|
bootstrap_domain,
|
|
distribution,
|
|
git_revision,
|
|
gpu_identity,
|
|
interval,
|
|
js_divergence,
|
|
load_candidates,
|
|
load_official_modules,
|
|
metric_vector,
|
|
scoped_seed,
|
|
sha256,
|
|
text_sha256,
|
|
)
|
|
|
|
|
|
CONDITIONS = ("raw", "user", "generation")
|
|
CONDITION_LABELS = {
|
|
"raw": "BOS + content",
|
|
"user": "official user template",
|
|
"generation": "official user template + Assistant:",
|
|
}
|
|
COMPARISONS = (
|
|
("raw_to_user", "raw", "user"),
|
|
("user_to_generation", "user", "generation"),
|
|
)
|
|
SAMPLE_SALT = "llm-atlas-deepseek-routing-template-control-v1"
|
|
CHAT_TEMPLATE_REVISION = "604d5664dddd88a0433dbae533b7fe9472482de0"
|
|
|
|
|
|
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("--content-tokens", type=int, default=23)
|
|
parser.add_argument("--batch-prompts", type=int, default=8)
|
|
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("--device", default="cuda")
|
|
parser.add_argument("--captured-at", default=None)
|
|
return parser.parse_args()
|
|
|
|
|
|
def canonical_hash(value: Any) -> str:
|
|
payload = json.dumps(
|
|
value,
|
|
ensure_ascii=False,
|
|
sort_keys=True,
|
|
separators=(",", ":"),
|
|
).encode()
|
|
return hashlib.sha256(payload).hexdigest()
|
|
|
|
|
|
def tokenize_with_offsets(
|
|
tokenizer: Any,
|
|
text: str,
|
|
*,
|
|
add_special_tokens: bool,
|
|
) -> tuple[list[int], list[tuple[int, int]]]:
|
|
encoded = tokenizer(
|
|
text,
|
|
add_special_tokens=add_special_tokens,
|
|
return_offsets_mapping=True,
|
|
padding=False,
|
|
truncation=False,
|
|
)
|
|
return list(encoded.input_ids), [tuple(pair) for pair in encoded.offset_mapping]
|
|
|
|
|
|
def content_positions(
|
|
token_ids: list[int],
|
|
offsets: list[tuple[int, int]],
|
|
content_start: int,
|
|
content_end: int,
|
|
) -> tuple[list[int], list[dict[str, int]], int]:
|
|
positions = []
|
|
records = []
|
|
crossing = 0
|
|
for index, (token_id, (start, end)) in enumerate(zip(token_ids, offsets, strict=True)):
|
|
if end <= start:
|
|
continue
|
|
overlaps = start < content_end and end > content_start
|
|
inside = start >= content_start and end <= content_end
|
|
if inside:
|
|
positions.append(index)
|
|
records.append(
|
|
{
|
|
"position": index,
|
|
"token_id": token_id,
|
|
"start": start - content_start,
|
|
"end": end - content_start,
|
|
}
|
|
)
|
|
elif overlaps:
|
|
crossing += 1
|
|
return positions, records, crossing
|
|
|
|
|
|
def common_prefix_length(left: list[int], right: list[int]) -> int:
|
|
length = 0
|
|
for left_id, right_id in zip(left, right):
|
|
if left_id != right_id:
|
|
break
|
|
length += 1
|
|
return length
|
|
|
|
|
|
def render_variant(tokenizer: Any, content: str, condition: str) -> dict[str, Any]:
|
|
if condition == "raw":
|
|
rendered = content
|
|
token_ids, offsets = tokenize_with_offsets(
|
|
tokenizer,
|
|
rendered,
|
|
add_special_tokens=True,
|
|
)
|
|
content_start = 0
|
|
else:
|
|
messages = [{"role": "user", "content": content}]
|
|
add_generation_prompt = condition == "generation"
|
|
rendered = tokenizer.apply_chat_template(
|
|
messages,
|
|
tokenize=False,
|
|
add_generation_prompt=add_generation_prompt,
|
|
)
|
|
official_ids = list(
|
|
tokenizer.apply_chat_template(
|
|
messages,
|
|
tokenize=True,
|
|
add_generation_prompt=add_generation_prompt,
|
|
)
|
|
)
|
|
token_ids, offsets = tokenize_with_offsets(
|
|
tokenizer,
|
|
rendered,
|
|
add_special_tokens=False,
|
|
)
|
|
if token_ids != official_ids:
|
|
raise RuntimeError(f"{condition} rendered token IDs differ from apply_chat_template")
|
|
content_start = rendered.index(content)
|
|
|
|
content_end = content_start + len(content)
|
|
positions, records, crossing = content_positions(
|
|
token_ids,
|
|
offsets,
|
|
content_start,
|
|
content_end,
|
|
)
|
|
if not positions:
|
|
raise RuntimeError(f"{condition} has no content tokens")
|
|
return {
|
|
"condition": condition,
|
|
"rendered_sha256": text_sha256(rendered),
|
|
"token_ids": token_ids,
|
|
"tokens": len(token_ids),
|
|
"token_ids_sha256": canonical_hash(token_ids),
|
|
"content_positions": positions,
|
|
"content_records": records,
|
|
"content_tokens": len(positions),
|
|
"wrapper_tokens": len(token_ids) - len(positions),
|
|
"boundary_crossing_tokens": crossing,
|
|
}
|
|
|
|
|
|
def select_samples(
|
|
tokenizer: Any,
|
|
candidates: dict[str, list[dict[str, str]]],
|
|
per_domain: int,
|
|
content_tokens: int,
|
|
sample_salt: str,
|
|
) -> tuple[list[dict[str, Any]], dict[str, dict[str, int]]]:
|
|
def fixed_token_prefix(
|
|
text: str,
|
|
full_offsets: list[tuple[int, int]],
|
|
) -> tuple[str, list[int]]:
|
|
lower_index = max(0, content_tokens - 5)
|
|
upper_index = min(len(full_offsets) - 1, content_tokens + 7)
|
|
lower = max(1, full_offsets[lower_index][0])
|
|
upper = max(lower, full_offsets[upper_index][1])
|
|
matches = []
|
|
for end in range(lower, upper + 1):
|
|
prefix = text[:end]
|
|
ids = list(tokenizer(prefix, add_special_tokens=False).input_ids)
|
|
if len(ids) == content_tokens:
|
|
matches.append((end, ids))
|
|
if not matches:
|
|
raise RuntimeError("could not construct a fixed-token character prefix")
|
|
end, ids = matches[-1]
|
|
return text[:end], ids
|
|
|
|
selected = []
|
|
counts = {}
|
|
for domain in DOMAIN_ORDER:
|
|
ranked = []
|
|
for row in candidates[domain]:
|
|
full_ids, full_offsets = tokenize_with_offsets(
|
|
tokenizer,
|
|
row["text"],
|
|
add_special_tokens=False,
|
|
)
|
|
if len(full_ids) < content_tokens:
|
|
continue
|
|
rank = hashlib.sha256(
|
|
f"{sample_salt}|{domain}|{row['id']}".encode()
|
|
).hexdigest()
|
|
ranked.append((rank, row, len(full_ids), full_offsets))
|
|
ranked.sort(key=lambda item: (item[0], item[1]["id"]))
|
|
|
|
eligible = []
|
|
canonicalization_failures = 0
|
|
for rank, row, source_tokens, full_offsets in ranked:
|
|
try:
|
|
content, prefix_ids = fixed_token_prefix(row["text"], full_offsets)
|
|
except RuntimeError:
|
|
canonicalization_failures += 1
|
|
continue
|
|
variants = {
|
|
condition: render_variant(tokenizer, content, condition)
|
|
for condition in CONDITIONS
|
|
}
|
|
shared_content_keys = set.intersection(
|
|
*(
|
|
{
|
|
(record["start"], record["end"], record["token_id"])
|
|
for record in variants[condition]["content_records"]
|
|
}
|
|
for condition in CONDITIONS
|
|
)
|
|
)
|
|
for condition in CONDITIONS:
|
|
variants[condition]["aligned_content_positions"] = [
|
|
record["position"]
|
|
for record in variants[condition]["content_records"]
|
|
if (record["start"], record["end"], record["token_id"])
|
|
in shared_content_keys
|
|
]
|
|
variants[condition]["aligned_content_tokens"] = len(
|
|
variants[condition]["aligned_content_positions"]
|
|
)
|
|
eligible.append(
|
|
{
|
|
"id": row["id"],
|
|
"domain": domain,
|
|
"label": DOMAIN_LABELS[domain],
|
|
"text_sha256": text_sha256(row["text"]),
|
|
"source_characters": len(row["text"]),
|
|
"source_tokens": source_tokens,
|
|
"content": content,
|
|
"content_sha256": text_sha256(content),
|
|
"content_characters": len(content),
|
|
"canonical_content_token_ids_sha256": canonical_hash(prefix_ids),
|
|
"selection_rank": rank,
|
|
"variants": variants,
|
|
}
|
|
)
|
|
if len(eligible) == per_domain:
|
|
break
|
|
if len(eligible) < per_domain:
|
|
raise RuntimeError(
|
|
f"{domain} has only {len(eligible)} eligible prompts; need {per_domain}"
|
|
)
|
|
domain_rows = eligible[:per_domain]
|
|
for index, row in enumerate(domain_rows):
|
|
row["within_domain_index"] = index
|
|
selected.extend(domain_rows)
|
|
counts[domain] = {
|
|
"candidate_records_after_text_filter": len(candidates[domain]),
|
|
"eligible_records": len(ranked),
|
|
"canonicalization_failures_before_selection_complete": (
|
|
canonicalization_failures
|
|
),
|
|
"selected_records": len(domain_rows),
|
|
"source_tokens_min": min(row["source_tokens"] for row in domain_rows),
|
|
"source_tokens_mean": float(
|
|
np.mean([row["source_tokens"] for row in domain_rows])
|
|
),
|
|
"source_tokens_max": max(row["source_tokens"] for row in domain_rows),
|
|
}
|
|
return selected, counts
|
|
|
|
|
|
def make_batches(
|
|
samples: list[dict[str, Any]],
|
|
batch_prompts: int,
|
|
pad_token_id: int,
|
|
) -> list[dict[str, Any]]:
|
|
ordered = sorted(
|
|
samples,
|
|
key=lambda row: (
|
|
max(row["variants"][condition]["tokens"] for condition in CONDITIONS),
|
|
DOMAIN_ORDER.index(row["domain"]),
|
|
row["id"],
|
|
),
|
|
)
|
|
batches = []
|
|
for start in range(0, len(ordered), batch_prompts):
|
|
prompt_rows = ordered[start : start + batch_prompts]
|
|
variants = [
|
|
{
|
|
"sample": sample,
|
|
"condition": condition,
|
|
**sample["variants"][condition],
|
|
}
|
|
for sample in prompt_rows
|
|
for condition in CONDITIONS
|
|
]
|
|
sequence = max(row["tokens"] for row in variants)
|
|
input_ids = torch.full(
|
|
(len(variants), sequence),
|
|
pad_token_id,
|
|
dtype=torch.long,
|
|
)
|
|
attention_mask = torch.zeros((len(variants), sequence), dtype=torch.long)
|
|
for index, row in enumerate(variants):
|
|
length = row["tokens"]
|
|
input_ids[index, :length] = torch.tensor(row["token_ids"])
|
|
attention_mask[index, :length] = 1
|
|
batches.append(
|
|
{
|
|
"prompt_rows": prompt_rows,
|
|
"variants": variants,
|
|
"input_ids": input_ids,
|
|
"attention_mask": attention_mask,
|
|
"padded_sequence": sequence,
|
|
}
|
|
)
|
|
return batches
|
|
|
|
|
|
def bootstrap_distributions(
|
|
loads: np.ndarray,
|
|
mode: str,
|
|
sampled: np.ndarray,
|
|
) -> np.ndarray:
|
|
if mode == "token_weighted":
|
|
values = loads[sampled].sum(axis=1, dtype=np.float64)
|
|
return values / values.sum(axis=1, keepdims=True)
|
|
prompt_distributions = loads / loads.sum(axis=1, keepdims=True)
|
|
values = prompt_distributions[sampled].mean(axis=1)
|
|
return values / values.sum(axis=1, keepdims=True)
|
|
|
|
|
|
def paired_domain(
|
|
before: np.ndarray,
|
|
after: np.ndarray,
|
|
mode: str,
|
|
replicates: int,
|
|
seed: int,
|
|
scope: str,
|
|
) -> dict[str, Any]:
|
|
if before.shape != after.shape:
|
|
raise ValueError(f"paired shape mismatch: {before.shape} != {after.shape}")
|
|
rng = np.random.default_rng(scoped_seed(seed, scope))
|
|
sampled = rng.integers(
|
|
0,
|
|
before.shape[0],
|
|
size=(replicates, before.shape[0]),
|
|
endpoint=False,
|
|
)
|
|
before_point = distribution(before, mode)
|
|
after_point = distribution(after, mode)
|
|
before_boot = bootstrap_distributions(before, mode, sampled)
|
|
after_boot = bootstrap_distributions(after, mode, sampled)
|
|
before_metrics = metric_vector(before_point)
|
|
after_metrics = metric_vector(after_point)
|
|
before_boot_metrics = metric_vector(before_boot)
|
|
after_boot_metrics = metric_vector(after_boot)
|
|
metrics = {}
|
|
for name in before_metrics:
|
|
delta_boot = after_boot_metrics[name] - before_boot_metrics[name]
|
|
metrics[name] = {
|
|
"before": float(before_metrics[name][0]),
|
|
"after": float(after_metrics[name][0]),
|
|
"delta_after_minus_before": float(
|
|
after_metrics[name][0] - before_metrics[name][0]
|
|
),
|
|
"delta_ci95": interval(delta_boot),
|
|
}
|
|
tv_boot = 0.5 * np.abs(after_boot - before_boot).sum(axis=1)
|
|
jsd_boot = js_divergence(before_boot, after_boot)
|
|
share_delta = after_point - before_point
|
|
return {
|
|
"metrics": metrics,
|
|
"total_variation": {
|
|
"point": float(0.5 * np.abs(share_delta).sum()),
|
|
"ci95": interval(tv_boot),
|
|
},
|
|
"js_divergence": {
|
|
"point": float(js_divergence(before_point, after_point)[0]),
|
|
"ci95": interval(jsd_boot),
|
|
"unit": "nats",
|
|
"upper_bound": math.log(2),
|
|
},
|
|
"expert_share_delta": share_delta.tolist(),
|
|
"expert_share_delta_ci95": interval(after_boot - before_boot),
|
|
}
|
|
|
|
|
|
def aligned_pairs(
|
|
left: dict[str, Any],
|
|
right: dict[str, Any],
|
|
) -> list[tuple[int, int]]:
|
|
right_by_key = {
|
|
(row["start"], row["end"], row["token_id"]): row["position"]
|
|
for row in right["content_records"]
|
|
}
|
|
return [
|
|
(
|
|
row["position"],
|
|
right_by_key[(row["start"], row["end"], row["token_id"])],
|
|
)
|
|
for row in left["content_records"]
|
|
if (row["start"], row["end"], row["token_id"]) in right_by_key
|
|
]
|
|
|
|
|
|
def route_alignment(
|
|
left_routes: torch.Tensor,
|
|
right_routes: torch.Tensor,
|
|
pairs: list[tuple[int, int]],
|
|
) -> dict[str, Any]:
|
|
if not pairs:
|
|
return {
|
|
"aligned_tokens": 0,
|
|
"ordered_topk_exact": 0,
|
|
"set_topk_exact": 0,
|
|
"mean_topk_overlap": None,
|
|
"mean_jaccard": None,
|
|
}
|
|
ordered_exact = 0
|
|
set_exact = 0
|
|
overlaps = []
|
|
jaccards = []
|
|
for left_position, right_position in pairs:
|
|
left = left_routes[left_position].tolist()
|
|
right = right_routes[right_position].tolist()
|
|
ordered_exact += int(left == right)
|
|
left_set = set(left)
|
|
right_set = set(right)
|
|
intersection = len(left_set & right_set)
|
|
union = len(left_set | right_set)
|
|
set_exact += int(left_set == right_set)
|
|
overlaps.append(intersection)
|
|
jaccards.append(intersection / union)
|
|
return {
|
|
"aligned_tokens": len(pairs),
|
|
"ordered_topk_exact": ordered_exact,
|
|
"set_topk_exact": set_exact,
|
|
"ordered_topk_exact_rate": ordered_exact / len(pairs),
|
|
"set_topk_exact_rate": set_exact / len(pairs),
|
|
"mean_topk_overlap": float(np.mean(overlaps)),
|
|
"mean_jaccard": float(np.mean(jaccards)),
|
|
}
|
|
|
|
|
|
def layer_statistics(
|
|
prompt_rows: list[dict[str, Any]],
|
|
replicates: int,
|
|
seed: int,
|
|
layer_index: int,
|
|
) -> dict[str, Any]:
|
|
scopes = {}
|
|
for load_scope, load_key in (
|
|
("full_input", "full_load"),
|
|
("content_only", "content_load"),
|
|
):
|
|
modes = {}
|
|
for mode in ("token_weighted", "prompt_balanced"):
|
|
conditions = {}
|
|
for condition in CONDITIONS:
|
|
conditions[condition] = {}
|
|
for domain in DOMAIN_ORDER:
|
|
loads = np.asarray(
|
|
[
|
|
row["conditions"][condition][load_key]
|
|
for row in prompt_rows
|
|
if row["domain"] == domain
|
|
],
|
|
dtype=np.int64,
|
|
)
|
|
conditions[condition][domain] = bootstrap_domain(
|
|
loads,
|
|
mode,
|
|
replicates,
|
|
seed,
|
|
(
|
|
f"layer={layer_index}|scope={load_scope}|mode={mode}|"
|
|
f"condition={condition}|domain={domain}"
|
|
),
|
|
)
|
|
|
|
comparisons = {}
|
|
for comparison, before_condition, after_condition in COMPARISONS:
|
|
comparisons[comparison] = {}
|
|
for domain in DOMAIN_ORDER:
|
|
rows = [row for row in prompt_rows if row["domain"] == domain]
|
|
before = np.asarray(
|
|
[
|
|
row["conditions"][before_condition][load_key]
|
|
for row in rows
|
|
],
|
|
dtype=np.int64,
|
|
)
|
|
after = np.asarray(
|
|
[
|
|
row["conditions"][after_condition][load_key]
|
|
for row in rows
|
|
],
|
|
dtype=np.int64,
|
|
)
|
|
comparisons[comparison][domain] = paired_domain(
|
|
before,
|
|
after,
|
|
mode,
|
|
replicates,
|
|
seed,
|
|
(
|
|
f"layer={layer_index}|scope={load_scope}|mode={mode}|"
|
|
f"comparison={comparison}|domain={domain}"
|
|
),
|
|
)
|
|
modes[mode] = {
|
|
"conditions": conditions,
|
|
"comparisons": comparisons,
|
|
}
|
|
scopes[load_scope] = {"modes": modes}
|
|
return scopes
|
|
|
|
|
|
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",
|
|
root / "tokenizer_config.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.content_tokens < 4:
|
|
raise ValueError("content prefix must contain at least four tokens")
|
|
if args.batch_prompts < 1:
|
|
raise ValueError("batch-prompts must be positive")
|
|
if args.bootstrap < 100:
|
|
raise ValueError("bootstrap replicates must be at least 100")
|
|
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 not tokenizer.is_fast:
|
|
raise RuntimeError("offset alignment requires a fast tokenizer")
|
|
if tokenizer.pad_token_id is None:
|
|
tokenizer.pad_token = tokenizer.eos_token
|
|
tokenizer.padding_side = "right"
|
|
tokenizer_config = json.loads((root / "tokenizer_config.json").read_text())
|
|
if tokenizer.chat_template != tokenizer_config.get("chat_template"):
|
|
raise RuntimeError("loaded chat template differs from tokenizer_config.json")
|
|
|
|
candidates = load_candidates(args)
|
|
samples, corpus_counts = select_samples(
|
|
tokenizer,
|
|
candidates,
|
|
args.per_domain,
|
|
args.content_tokens,
|
|
args.sample_salt,
|
|
)
|
|
batches = make_batches(samples, args.batch_prompts, 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 = []
|
|
causal_invariant = {
|
|
"shared_prefix_tokens": 0,
|
|
"ordered_topk_exact": 0,
|
|
"violations": 0,
|
|
}
|
|
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 = captures["topk_ids"].view(
|
|
len(batch["variants"]),
|
|
sequence,
|
|
config.num_experts_per_tok,
|
|
)
|
|
by_prompt: dict[str, dict[str, tuple[dict[str, Any], torch.Tensor]]] = {}
|
|
for variant_index, variant in enumerate(batch["variants"]):
|
|
sample = variant["sample"]
|
|
condition = variant["condition"]
|
|
routes = topk[variant_index, : variant["tokens"]]
|
|
by_prompt.setdefault(sample["id"], {})[condition] = (variant, routes)
|
|
|
|
for sample in batch["prompt_rows"]:
|
|
variants = by_prompt[sample["id"]]
|
|
condition_rows = {}
|
|
for condition in CONDITIONS:
|
|
variant, routes = variants[condition]
|
|
full_load = torch.bincount(
|
|
routes.flatten(),
|
|
minlength=config.n_routed_experts,
|
|
)
|
|
content_routes = routes[
|
|
variant["aligned_content_positions"]
|
|
]
|
|
content_load = torch.bincount(
|
|
content_routes.flatten(),
|
|
minlength=config.n_routed_experts,
|
|
)
|
|
condition_rows[condition] = {
|
|
"input_tokens": variant["tokens"],
|
|
"content_tokens": variant["content_tokens"],
|
|
"aligned_content_tokens": variant[
|
|
"aligned_content_tokens"
|
|
],
|
|
"wrapper_tokens": variant["wrapper_tokens"],
|
|
"boundary_crossing_tokens": variant[
|
|
"boundary_crossing_tokens"
|
|
],
|
|
"routes": int(full_load.sum()),
|
|
"content_routes": int(content_load.sum()),
|
|
"full_load": full_load.tolist(),
|
|
"content_load": content_load.tolist(),
|
|
"topk_sha256": canonical_hash(routes.tolist()),
|
|
"content_topk_sha256": canonical_hash(
|
|
content_routes.tolist()
|
|
),
|
|
}
|
|
|
|
raw_variant, raw_routes = variants["raw"]
|
|
user_variant, user_routes = variants["user"]
|
|
generation_variant, generation_routes = variants["generation"]
|
|
raw_user_pairs = aligned_pairs(raw_variant, user_variant)
|
|
raw_user = route_alignment(
|
|
raw_routes,
|
|
user_routes,
|
|
raw_user_pairs,
|
|
)
|
|
prefix_length = common_prefix_length(
|
|
user_variant["token_ids"],
|
|
generation_variant["token_ids"],
|
|
)
|
|
prefix_left = user_routes[:prefix_length]
|
|
prefix_right = generation_routes[:prefix_length]
|
|
exact_prefix = int(
|
|
torch.equal(prefix_left, prefix_right)
|
|
)
|
|
causal_invariant["shared_prefix_tokens"] += prefix_length
|
|
causal_invariant["ordered_topk_exact"] += (
|
|
prefix_length if exact_prefix else int(
|
|
(prefix_left == prefix_right).all(dim=1).sum()
|
|
)
|
|
)
|
|
causal_invariant["violations"] += int(not exact_prefix)
|
|
prompt_rows.append(
|
|
{
|
|
"id": sample["id"],
|
|
"domain": sample["domain"],
|
|
"conditions": condition_rows,
|
|
"alignments": {
|
|
"raw_to_user_content": {
|
|
**raw_user,
|
|
"raw_content_tokens": raw_variant[
|
|
"content_tokens"
|
|
],
|
|
"user_content_tokens": user_variant[
|
|
"content_tokens"
|
|
],
|
|
"raw_alignment_coverage": (
|
|
raw_user["aligned_tokens"]
|
|
/ raw_variant["content_tokens"]
|
|
),
|
|
"user_alignment_coverage": (
|
|
raw_user["aligned_tokens"]
|
|
/ user_variant["content_tokens"]
|
|
),
|
|
},
|
|
"user_to_generation_prefix": {
|
|
"shared_prefix_tokens": prefix_length,
|
|
"user_tokens": user_variant["tokens"],
|
|
"generation_tokens": generation_variant["tokens"],
|
|
"ordered_topk_exact": exact_prefix == 1,
|
|
"ordered_topk_exact_tokens": (
|
|
prefix_length if exact_prefix else int(
|
|
(prefix_left == prefix_right)
|
|
.all(dim=1)
|
|
.sum()
|
|
)
|
|
),
|
|
},
|
|
},
|
|
}
|
|
)
|
|
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"],
|
|
)
|
|
)
|
|
causal_invariant["exact_rate"] = (
|
|
causal_invariant["ordered_topk_exact"]
|
|
/ causal_invariant["shared_prefix_tokens"]
|
|
)
|
|
result["prompts"] = prompt_rows
|
|
result["causal_suffix_invariant"] = causal_invariant
|
|
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()
|
|
model_index = json.loads((root / "model.safetensors.index.json").read_text())
|
|
selected_identity = []
|
|
for sample in samples:
|
|
selected_identity.append(
|
|
{
|
|
"id": sample["id"],
|
|
"domain": sample["domain"],
|
|
"within_domain_index": sample["within_domain_index"],
|
|
"selection_rank": sample["selection_rank"],
|
|
"text_sha256": sample["text_sha256"],
|
|
"source_characters": sample["source_characters"],
|
|
"source_tokens": sample["source_tokens"],
|
|
"content_sha256": sample["content_sha256"],
|
|
"content_characters": sample["content_characters"],
|
|
"canonical_content_tokens": args.content_tokens,
|
|
"canonical_content_token_ids_sha256": sample[
|
|
"canonical_content_token_ids_sha256"
|
|
],
|
|
"conditions": {
|
|
condition: {
|
|
key: sample["variants"][condition][key]
|
|
for key in (
|
|
"rendered_sha256",
|
|
"tokens",
|
|
"token_ids_sha256",
|
|
"content_tokens",
|
|
"aligned_content_tokens",
|
|
"wrapper_tokens",
|
|
"boundary_crossing_tokens",
|
|
)
|
|
}
|
|
for condition in CONDITIONS
|
|
},
|
|
"user_generation_common_prefix_tokens": common_prefix_length(
|
|
sample["variants"]["user"]["token_ids"],
|
|
sample["variants"]["generation"]["token_ids"],
|
|
),
|
|
"raw_user_aligned_content_tokens": len(
|
|
aligned_pairs(
|
|
sample["variants"]["raw"],
|
|
sample["variants"]["user"],
|
|
)
|
|
),
|
|
}
|
|
)
|
|
|
|
result = {
|
|
"schema_version": 1,
|
|
"captured_at": captured_at,
|
|
"evidence_identity": (
|
|
"X / official BF16 weights, official tokenizer chat template, "
|
|
"paired 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,
|
|
"task_performance": False,
|
|
"training_or_online_load": False,
|
|
"expert_semantics_inferred": False,
|
|
"causal_claim": (
|
|
"the user-to-generation shared-prefix equality is a causal-mask "
|
|
"implementation invariant; raw-to-user differences are descriptive "
|
|
"protocol sensitivity, not a capability effect"
|
|
),
|
|
"population": (
|
|
f"{len(samples)} fixed public prompts across four domains; "
|
|
"not training data, online traffic, or a task benchmark"
|
|
),
|
|
"code_execution": False,
|
|
"answers_used": False,
|
|
},
|
|
"provenance": {
|
|
"model": {
|
|
"huggingface_model": "deepseek-ai/DeepSeek-V2-Lite",
|
|
"huggingface_revision": CHAT_TEMPLATE_REVISION,
|
|
"sha256": {
|
|
"config": sha256(root / "config.json"),
|
|
"modeling_code": sha256(root / "modeling_deepseek.py"),
|
|
"tokenizer": sha256(root / "tokenizer.json"),
|
|
"tokenizer_config": sha256(root / "tokenizer_config.json"),
|
|
"index": sha256(root / "model.safetensors.index.json"),
|
|
"shard_1": sha256(shard),
|
|
},
|
|
"checkpoint_tensor_bytes": model_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__,
|
|
"jinja2": __import__("jinja2").__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,
|
|
},
|
|
"template_contract": {
|
|
"chat_template_revision": CHAT_TEMPLATE_REVISION,
|
|
"chat_template": tokenizer.chat_template,
|
|
"chat_template_sha256": text_sha256(tokenizer.chat_template),
|
|
"bos_token": tokenizer.bos_token,
|
|
"bos_token_id": tokenizer.bos_token_id,
|
|
"eos_token": tokenizer.eos_token,
|
|
"eos_token_id": tokenizer.eos_token_id,
|
|
"conditions": {
|
|
condition: CONDITION_LABELS[condition]
|
|
for condition in CONDITIONS
|
|
},
|
|
"comparisons": [
|
|
{
|
|
"name": name,
|
|
"before": before,
|
|
"after": after,
|
|
}
|
|
for name, before, after in COMPARISONS
|
|
],
|
|
"scope_split": {
|
|
"full_input": "all BOS, wrapper, content, newline, and generation-prompt tokens",
|
|
"content_only": (
|
|
"the exact intersection of (relative character span, token ID) "
|
|
"inside source content across all three conditions; boundary-"
|
|
"crossing and unaligned tokens are excluded from every condition"
|
|
),
|
|
},
|
|
},
|
|
"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,
|
|
"canonical_content_tokens": args.content_tokens,
|
|
"content_prefix": (
|
|
"source text cut at the end offset of the fixed regular-tokenizer "
|
|
"content-token prefix"
|
|
),
|
|
"counts": corpus_counts,
|
|
"selected": selected_identity,
|
|
},
|
|
"inference_contract": {
|
|
"batch_prompts": args.batch_prompts,
|
|
"variants_per_prompt": len(CONDITIONS),
|
|
"rows_per_full_batch": args.batch_prompts * len(CONDITIONS),
|
|
"batches": len(batches),
|
|
"batch_grouping": (
|
|
"all raw/user/generation variants of one source prompt execute "
|
|
"in the same padded batch"
|
|
),
|
|
"attention": "official eager causal mask",
|
|
"dtype": "BF16",
|
|
"total_source_prompts": len(samples),
|
|
"total_prompt_variants": len(samples) * len(CONDITIONS),
|
|
"input_tokens_by_condition": {
|
|
condition: sum(
|
|
sample["variants"][condition]["tokens"] for sample in samples
|
|
)
|
|
for condition in CONDITIONS
|
|
},
|
|
"content_span_tokens_by_condition": {
|
|
condition: sum(
|
|
sample["variants"][condition]["content_tokens"]
|
|
for sample in samples
|
|
)
|
|
for condition in CONDITIONS
|
|
},
|
|
"aligned_content_tokens_by_condition": {
|
|
condition: sum(
|
|
sample["variants"][condition]["aligned_content_tokens"]
|
|
for sample in samples
|
|
)
|
|
for condition in CONDITIONS
|
|
},
|
|
"routes_per_condition_all_moe_layers": {
|
|
condition: sum(
|
|
sample["variants"][condition]["tokens"] for sample in samples
|
|
)
|
|
* config.num_experts_per_tok
|
|
* (args.layers - 1)
|
|
for condition in CONDITIONS
|
|
},
|
|
"total_routes_all_conditions_all_moe_layers": sum(
|
|
sum(sample["variants"][condition]["tokens"] for sample in samples)
|
|
* config.num_experts_per_tok
|
|
* (args.layers - 1)
|
|
for condition in CONDITIONS
|
|
),
|
|
},
|
|
"statistical_contract": {
|
|
"unit": "source prompt",
|
|
"bootstrap_replicates": args.bootstrap,
|
|
"seed": args.seed,
|
|
"paired_indices": (
|
|
"the same resampled source-prompt indices are used for before and "
|
|
"after within each domain/layer/scope/mode"
|
|
),
|
|
"modes": {
|
|
"token_weighted": "sum selected routes, then normalize",
|
|
"prompt_balanced": (
|
|
"normalize each prompt load, then average prompts equally"
|
|
),
|
|
},
|
|
"interval": "2.5th and 97.5th percentiles",
|
|
"multiple_comparison_correction": False,
|
|
"hypothesis_test": False,
|
|
},
|
|
"layers": layer_results,
|
|
}
|
|
args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
args.output.write_text(
|
|
json.dumps(result, indent=2, ensure_ascii=False) + "\n",
|
|
encoding="utf-8",
|
|
)
|
|
print(json.dumps(result, indent=2, ensure_ascii=False))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|