Files
llm-atlas/experiments/deepseek/v2_lite_chat_full_depth_trace.py
2026-07-30 00:24:27 +08:00

873 lines
29 KiB
Python

#!/usr/bin/env python3
"""Trace all 27 Chat decoder layers and all 26 MoE gates.
The probe reuses the fixed 16-source, eight-condition completion cohort but
runs prompt-only full-depth forwards. It stores hashes and summaries rather
than raw hidden tensors. Target-content comparisons use only byte-identical
interior content tokens; boundary-crossing tokens are explicitly excluded.
"""
from __future__ import annotations
import argparse
import hashlib
import json
import math
import os
import platform
import resource
import time
from datetime import datetime, timezone
from pathlib import Path
from statistics import mean
from typing import Any
import accelerate
import safetensors
import torch
import torch.nn.functional as F
import transformers
from transformers import AutoTokenizer
import v2_lite_chat_special_token_behavior_probe as behavior
import v2_lite_routing_special_token_family_control as special
EDGE_PAIRS = {
"system_eos": ("s0_eos", "s1_eos"),
"system_bos": ("s0_bos", "s1_bos"),
"system_x": ("s0_x", "s1_x"),
"system_period": ("s0_period", "s1_period"),
"bos_at_s0": ("s0_eos", "s0_bos"),
"bos_at_s1": ("s1_eos", "s1_bos"),
"x_at_s0": ("s0_eos", "s0_x"),
"x_at_s1": ("s1_eos", "s1_x"),
"period_at_s0": ("s0_eos", "s0_period"),
"period_at_s1": ("s1_eos", "s1_period"),
}
SCOPES = ("target_content", "full_input")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--artifact-dir", type=Path, required=True)
parser.add_argument("--reference-routing-json", 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(
"--domains",
nargs="+",
choices=behavior.DEFAULT_DOMAINS,
default=list(behavior.DEFAULT_DOMAINS),
)
parser.add_argument("--per-domain", type=int, default=1)
parser.add_argument("--gpu-memory", default="28GiB")
parser.add_argument("--cpu-memory", default="80GiB")
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 tensor_hash(tensor: torch.Tensor) -> str:
value = tensor.detach().cpu().contiguous()
header = json.dumps(
{
"shape": list(value.shape),
"dtype": str(value.dtype),
},
separators=(",", ":"),
).encode()
raw = value.view(torch.uint8).numpy().tobytes()
return hashlib.sha256(header + b"\0" + raw).hexdigest()
def validate_args(args: argparse.Namespace) -> None:
if args.per_domain <= 0:
raise ValueError("--per-domain must be positive")
for path in (
args.artifact_dir,
args.reference_routing_json,
args.human_eval,
args.gsm8k,
args.tnews,
args.tnews_archive,
args.wikitext,
):
if not path.exists():
raise FileNotFoundError(path)
missing = [
name
for name in behavior.MODEL_FILES
if not (args.artifact_dir / name).is_file()
]
if missing:
raise FileNotFoundError(
f"artifact directory is incomplete: {missing}"
)
def condition_positions(
variants: dict[str, dict[str, Any]],
batch_tokens: int,
) -> dict[str, dict[str, list[int]]]:
result = {}
for condition in special.CONDITIONS:
variant = variants[condition]
padding = batch_tokens - variant["tokens"]
result[condition] = {
"target_content": [
padding + int(position)
for position in variant["content_positions"]
],
"full_input": list(range(padding, batch_tokens)),
}
return result
def hidden_stats(value: torch.Tensor) -> dict[str, Any]:
fp32 = value.float()
return {
"shape": list(value.shape),
"dtype": str(value.dtype),
"tensor_sha256": tensor_hash(value),
"mean": fp32.mean().item(),
"std": fp32.std(unbiased=False).item(),
"rms": fp32.square().mean().sqrt().item(),
"max_abs": fp32.abs().max().item(),
}
def compare_hidden(
left: torch.Tensor,
right: torch.Tensor,
) -> dict[str, Any]:
if left.shape != right.shape:
raise RuntimeError(
f"hidden comparison shape mismatch: {left.shape} != "
f"{right.shape}"
)
left_fp32 = left.float()
right_fp32 = right.float()
delta = right_fp32 - left_fp32
denominator = left_fp32.norm(dim=-1).clamp_min(1e-12)
relative = delta.norm(dim=-1) / denominator
cosine = F.cosine_similarity(
left_fp32,
right_fp32,
dim=-1,
eps=1e-12,
).clamp(-1.0, 1.0)
exact = (left == right).all(dim=-1)
return {
"tokens": left.shape[0],
"exact_hidden_rows": exact.sum().item(),
"mean_cosine_similarity": cosine.mean().item(),
"min_cosine_similarity": cosine.min().item(),
"mean_relative_l2": relative.mean().item(),
"max_abs_delta": delta.abs().max().item(),
}
def route_distribution(
route_ids: torch.Tensor,
route_weights: torch.Tensor,
experts: int,
) -> dict[str, Any]:
loads = torch.bincount(
route_ids.reshape(-1),
minlength=experts,
).to(torch.int64)
weighted = torch.zeros(experts, dtype=torch.float64)
weighted.scatter_add_(
0,
route_ids.reshape(-1),
route_weights.reshape(-1).to(torch.float64),
)
total = loads.sum().item()
probabilities = loads.to(torch.float64) / max(total, 1)
nonzero = probabilities[probabilities > 0]
entropy = (
-(nonzero * nonzero.log()).sum().item() / math.log(experts)
if len(nonzero)
else 0.0
)
mean_load = loads.double().mean()
cv = (
loads.double().std(unbiased=False) / mean_load
if mean_load > 0
else torch.tensor(0.0)
)
return {
"tokens": route_ids.shape[0],
"top_k": route_ids.shape[1],
"decisions": total,
"ordered_route_sha256": canonical_hash(route_ids.tolist()),
"route_weight_sha256": tensor_hash(route_weights),
"loads": loads.tolist(),
"weighted_loads": weighted.tolist(),
"unique_experts": int((loads > 0).sum().item()),
"load_cv": cv.item(),
"normalized_load_entropy": entropy,
"max_load_share": probabilities.max().item(),
}
def normalized_load(loads: torch.Tensor) -> torch.Tensor:
value = loads.to(torch.float64)
return value / value.sum().clamp_min(1e-12)
def compare_routes(
left_ids: torch.Tensor,
left_weights: torch.Tensor,
right_ids: torch.Tensor,
right_weights: torch.Tensor,
experts: int,
) -> dict[str, Any]:
if left_ids.shape != right_ids.shape:
raise RuntimeError(
f"route comparison shape mismatch: {left_ids.shape} != "
f"{right_ids.shape}"
)
ordered_exact = (left_ids == right_ids).all(dim=-1)
left_sorted = left_ids.sort(dim=-1).values
right_sorted = right_ids.sort(dim=-1).values
set_exact = (left_sorted == right_sorted).all(dim=-1)
left_hot = F.one_hot(
left_ids,
num_classes=experts,
).sum(dim=1).bool()
right_hot = F.one_hot(
right_ids,
num_classes=experts,
).sum(dim=1).bool()
intersection = (left_hot & right_hot).sum(dim=-1)
union = (left_hot | right_hot).sum(dim=-1).clamp_min(1)
left_token_weighted = torch.zeros(
left_ids.shape[0],
experts,
dtype=torch.float64,
)
right_token_weighted = torch.zeros_like(left_token_weighted)
left_token_weighted.scatter_add_(
1,
left_ids,
left_weights.to(torch.float64),
)
right_token_weighted.scatter_add_(
1,
right_ids,
right_weights.to(torch.float64),
)
token_weighted_tv = 0.5 * (
left_token_weighted - right_token_weighted
).abs().sum(dim=-1)
left_load = torch.bincount(
left_ids.reshape(-1),
minlength=experts,
)
right_load = torch.bincount(
right_ids.reshape(-1),
minlength=experts,
)
load_tv = 0.5 * (
normalized_load(left_load)
- normalized_load(right_load)
).abs().sum()
return {
"tokens": left_ids.shape[0],
"ordered_topk_exact_tokens": ordered_exact.sum().item(),
"set_exact_tokens": set_exact.sum().item(),
"mean_set_jaccard": (
intersection.to(torch.float64)
/ union.to(torch.float64)
).mean().item(),
"mean_token_weighted_tv": token_weighted_tv.mean().item(),
"aggregate_load_tv": load_tv.item(),
}
class TraceCapture:
def __init__(self, experts: int) -> None:
self.experts = experts
self.positions: dict[str, dict[str, list[int]]] = {}
self.batch_tokens = 0
self.hidden: dict[str, Any] = {}
self.routes: dict[str, Any] = {}
def start(
self,
positions: dict[str, dict[str, list[int]]],
batch_tokens: int,
) -> None:
self.positions = positions
self.batch_tokens = batch_tokens
self.hidden = {}
self.routes = {}
def hidden_hook(self, stage: str):
def hook(
_module: torch.nn.Module,
_inputs: tuple[Any, ...],
output: Any,
) -> None:
tensor = output[0] if isinstance(output, tuple) else output
if tensor.ndim != 3:
raise RuntimeError(
f"{stage} hidden rank {tensor.ndim}, expected 3"
)
cpu = tensor.detach().cpu()
condition_values: dict[str, dict[str, torch.Tensor]] = {}
condition_rows = {}
for row, condition in enumerate(special.CONDITIONS):
condition_values[condition] = {}
condition_rows[condition] = {}
for scope in SCOPES:
positions = self.positions[condition][scope]
value = cpu[row, positions, :].contiguous()
condition_values[condition][scope] = value
condition_rows[condition][scope] = hidden_stats(value)
comparisons = {}
for name, (left, right) in EDGE_PAIRS.items():
comparisons[name] = compare_hidden(
condition_values[left]["target_content"],
condition_values[right]["target_content"],
)
self.hidden[stage] = {
"conditions": condition_rows,
"target_comparisons": comparisons,
}
return hook
def gate_hook(self, layer: int):
def hook(
_module: torch.nn.Module,
_inputs: tuple[Any, ...],
output: Any,
) -> None:
route_ids, route_weights, _aux = output
top_k = route_ids.shape[-1]
ids = route_ids.detach().reshape(
len(special.CONDITIONS),
self.batch_tokens,
top_k,
).cpu()
weights = route_weights.detach().reshape(
len(special.CONDITIONS),
self.batch_tokens,
top_k,
).cpu()
values: dict[
str,
dict[str, tuple[torch.Tensor, torch.Tensor]],
] = {}
condition_rows = {}
for row, condition in enumerate(special.CONDITIONS):
values[condition] = {}
condition_rows[condition] = {}
for scope in SCOPES:
positions = self.positions[condition][scope]
scoped_ids = ids[row, positions, :].contiguous()
scoped_weights = weights[
row,
positions,
:,
].contiguous()
values[condition][scope] = (
scoped_ids,
scoped_weights,
)
condition_rows[condition][scope] = (
route_distribution(
scoped_ids,
scoped_weights,
self.experts,
)
)
comparisons = {}
for name, (left, right) in EDGE_PAIRS.items():
left_ids, left_weights = values[left][
"target_content"
]
right_ids, right_weights = values[right][
"target_content"
]
comparisons[name] = compare_routes(
left_ids,
left_weights,
right_ids,
right_weights,
self.experts,
)
self.routes[f"layer_{layer:02d}"] = {
"layer": layer,
"conditions": condition_rows,
"target_comparisons": comparisons,
}
return hook
def aggregate_trace(sources: list[dict[str, Any]]) -> dict[str, Any]:
hidden_stages = [
"embedding",
*(f"layer_{index:02d}" for index in range(27)),
"final_norm",
]
router_layers = [
f"layer_{index:02d}"
for index in range(1, 27)
]
hidden = {}
for stage in hidden_stages:
hidden[stage] = {}
for edge in EDGE_PAIRS:
rows = [
source["hidden_stages"][stage][
"target_comparisons"
][edge]
for source in sources
]
tokens = sum(row["tokens"] for row in rows)
hidden[stage][edge] = {
"sources": len(rows),
"tokens": tokens,
"exact_hidden_rows": sum(
row["exact_hidden_rows"] for row in rows
),
"token_weighted_mean_cosine_similarity": sum(
row["mean_cosine_similarity"] * row["tokens"]
for row in rows
) / tokens,
"source_mean_relative_l2": mean(
row["mean_relative_l2"] for row in rows
),
"max_abs_delta": max(
row["max_abs_delta"] for row in rows
),
}
routes = {}
for layer in router_layers:
routes[layer] = {}
for edge in EDGE_PAIRS:
rows = [
source["router_layers"][layer][
"target_comparisons"
][edge]
for source in sources
]
tokens = sum(row["tokens"] for row in rows)
routes[layer][edge] = {
"sources": len(rows),
"tokens": tokens,
"ordered_topk_exact_tokens": sum(
row["ordered_topk_exact_tokens"] for row in rows
),
"set_exact_tokens": sum(
row["set_exact_tokens"] for row in rows
),
"token_weighted_mean_set_jaccard": sum(
row["mean_set_jaccard"] * row["tokens"]
for row in rows
) / tokens,
"source_mean_token_weighted_tv": mean(
row["mean_token_weighted_tv"] for row in rows
),
"source_mean_aggregate_load_tv": mean(
row["aggregate_load_tv"] for row in rows
),
}
target_tokens = sum(
source["content_tokens"] for source in sources
)
return {
"sources": len(sources),
"conditions_per_source": len(special.CONDITIONS),
"decoder_layers_executed": 27,
"moe_gates_executed": 26,
"hidden_stages": len(hidden_stages),
"target_content_tokens_per_condition": target_tokens,
"target_route_decisions": (
target_tokens
* len(special.CONDITIONS)
* 26
* 6
),
"hidden": hidden,
"routes": routes,
}
def main() -> None:
args = parse_args()
validate_args(args)
captured_at = args.captured_at or datetime.now(
timezone.utc
).isoformat()
revision_contract = behavior.download_revision_contract(
args.artifact_dir
)
special.install_control_contract()
tokenizer = AutoTokenizer.from_pretrained(
args.artifact_dir,
trust_remote_code=True,
local_files_only=True,
use_fast=True,
)
if tokenizer.pad_token_id is None:
tokenizer.pad_token_id = tokenizer.eos_token_id
if special.special_family_token_ids(tokenizer) != {
"eos": 100001,
"bos": 100000,
"x": 87,
"period": 13,
}:
raise RuntimeError("pinned tokenizer ID contract changed")
source_rows, reference = behavior.selected_sources(args, tokenizer)
for source in source_rows:
source["variants"] = {
condition: special.prior.render_boundary_variant(
tokenizer,
source["content"],
condition,
)
for condition in special.CONDITIONS
}
content_sequences = {
tuple(
source["variants"][condition]["token_ids"][position]
for position in source["variants"][condition][
"content_positions"
]
)
for condition in special.CONDITIONS
}
if len(content_sequences) != 1:
raise RuntimeError(
f"{source['id']} target token sequence is not exact "
"across conditions"
)
source["content_token_ids_sha256"] = canonical_hash(
next(iter(content_sequences))
)
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
torch.cuda.synchronize()
load_started = time.perf_counter()
_, official_modeling = special.base.load_official_modules(
args.artifact_dir
)
model = official_modeling.DeepseekV2ForCausalLM.from_pretrained(
args.artifact_dir,
local_files_only=True,
torch_dtype=torch.bfloat16,
low_cpu_mem_usage=True,
device_map="auto",
max_memory={
0: args.gpu_memory,
"cpu": args.cpu_memory,
},
)
model.eval()
if torch.cuda.is_available():
torch.cuda.synchronize()
load_seconds = time.perf_counter() - load_started
load_peak_cuda = (
torch.cuda.max_memory_allocated()
if torch.cuda.is_available()
else None
)
input_device = model.get_input_embeddings().weight.device
capture = TraceCapture(model.config.n_routed_experts)
handles = [
model.model.embed_tokens.register_forward_hook(
capture.hidden_hook("embedding")
),
model.model.norm.register_forward_hook(
capture.hidden_hook("final_norm")
),
]
for index, layer in enumerate(model.model.layers):
handles.append(
layer.register_forward_hook(
capture.hidden_hook(f"layer_{index:02d}")
)
)
if hasattr(layer.mlp, "gate"):
handles.append(
layer.mlp.gate.register_forward_hook(
capture.gate_hook(index)
)
)
traced_sources = []
try:
for source in source_rows:
unpadded = [
source["variants"][condition]["token_ids"]
for condition in special.CONDITIONS
]
prompt_lengths = [len(row) for row in unpadded]
batch_tokens = max(prompt_lengths)
padded = [
[int(tokenizer.pad_token_id)]
* (batch_tokens - len(row))
+ row
for row in unpadded
]
masks = [
[0] * (batch_tokens - len(row))
+ [1] * len(row)
for row in unpadded
]
positions = condition_positions(
source["variants"],
batch_tokens,
)
capture.start(positions, batch_tokens)
input_ids = torch.tensor(
padded,
dtype=torch.long,
device=input_device,
)
attention_mask = torch.tensor(
masks,
dtype=torch.long,
device=input_device,
)
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats()
torch.cuda.synchronize()
started = time.perf_counter()
with torch.inference_mode():
output = model.model(
input_ids=input_ids,
attention_mask=attention_mask,
use_cache=False,
output_attentions=False,
output_hidden_states=False,
return_dict=True,
)
if torch.cuda.is_available():
torch.cuda.synchronize()
elapsed = time.perf_counter() - started
peak = (
torch.cuda.max_memory_allocated()
if torch.cuda.is_available()
else None
)
if output.last_hidden_state.shape[:2] != input_ids.shape:
raise RuntimeError(
"final hidden batch/sequence shape does not match input"
)
expected_hidden = {
"embedding",
*(f"layer_{index:02d}" for index in range(27)),
"final_norm",
}
expected_routes = {
f"layer_{index:02d}"
for index in range(1, 27)
}
if set(capture.hidden) != expected_hidden:
raise RuntimeError(
f"hidden stages mismatch: {set(capture.hidden)}"
)
if set(capture.routes) != expected_routes:
raise RuntimeError(
f"router layers mismatch: {set(capture.routes)}"
)
content_tokens = len(
positions[special.CONDITIONS[0]][
"target_content"
]
)
traced_sources.append(
{
key: value
for key, value in source.items()
if key not in {"content", "variants"}
}
| {
"prompt_tokens_by_condition": {
condition: prompt_lengths[index]
for index, condition in enumerate(
special.CONDITIONS
)
},
"batch_prompt_tokens": batch_tokens,
"content_tokens": content_tokens,
"content_token_ids_sha256": source[
"content_token_ids_sha256"
],
"boundary_crossing_tokens_by_condition": {
condition: source["variants"][condition][
"boundary_crossing_tokens"
]
for condition in special.CONDITIONS
},
"forward_seconds": elapsed,
"peak_cuda_memory_allocated_bytes": peak,
"hidden_stages": capture.hidden,
"router_layers": capture.routes,
}
)
finally:
for handle in handles:
handle.remove()
device_map = getattr(model, "hf_device_map", {})
peaks = [
source["peak_cuda_memory_allocated_bytes"]
for source in traced_sources
if source["peak_cuda_memory_allocated_bytes"] is not None
]
result = {
"schema_version": 1,
"captured_at": captured_at,
"model": {
"repo": "deepseek-ai/DeepSeek-V2-Lite-Chat",
"revision": behavior.MODEL_REVISION,
"checkpoint_identity": "SFT Chat",
"architecture": type(model).__name__,
"dtype": "bfloat16",
"download_revision_contract": revision_contract,
"files": {
name: {
"bytes": (args.artifact_dir / name).stat().st_size,
"sha256": behavior.sha256(
args.artifact_dir / name
),
}
for name in behavior.MODEL_FILES
},
},
"tokenizer_contract": {
"length": len(tokenizer),
"vocab_size": tokenizer.vocab_size,
"all_special_ids": tokenizer.all_special_ids,
"pad_token_id": tokenizer.pad_token_id,
"conditions": special.FACTORS,
},
"source_contract": {
**reference,
"domains": args.domains,
"per_domain": args.per_domain,
"sources": len(source_rows),
"full_source_text_used": True,
"target_scope": (
"byte-identical interior content tokens; tokens whose "
"offset crosses the content boundary are excluded"
),
},
"trace_contract": {
"conditions": list(special.CONDITIONS),
"scopes": list(SCOPES),
"edge_pairs": EDGE_PAIRS,
"decoder_layers": 27,
"moe_gate_layers": list(range(1, 27)),
"experts": model.config.n_routed_experts,
"top_k": model.config.num_experts_per_tok,
"use_cache": False,
"raw_hidden_tensors_saved": False,
"raw_route_ids_saved": False,
"load_vectors_saved": True,
"hashes_cover_exact_bfloat16_hidden_and_route_values": True,
},
"execution": {
"python": platform.python_version(),
"torch": torch.__version__,
"transformers": transformers.__version__,
"accelerate": accelerate.__version__,
"safetensors": safetensors.__version__,
"platform": platform.platform(),
"gpu": (
torch.cuda.get_device_name(0)
if torch.cuda.is_available()
else None
),
"gpu_memory_limit": args.gpu_memory,
"cpu_memory_limit": args.cpu_memory,
"pytorch_cuda_alloc_conf": os.environ.get(
"PYTORCH_CUDA_ALLOC_CONF"
),
"input_device": str(input_device),
"device_map": device_map,
"load_seconds": load_seconds,
"forward_seconds": sum(
source["forward_seconds"]
for source in traced_sources
),
"load_peak_cuda_memory_allocated_bytes": load_peak_cuda,
"peak_cuda_memory_allocated_bytes": max(
[load_peak_cuda, *peaks]
),
"process_max_rss_kib": resource.getrusage(
resource.RUSAGE_SELF
).ru_maxrss,
},
"sources": traced_sources,
"summary": aggregate_trace(traced_sources),
"claim_boundary": [
"The trace observes one fixed SFT Chat checkpoint, not the Base checkpoint.",
"Target comparisons exclude boundary-crossing tokens and require exact token IDs.",
"Router and hidden associations do not establish mediation or capability causality.",
"CPU in hf_device_map is offload residency, not proof of CPU matrix execution.",
"Prompt-only use_cache=False traces are not generation-time KV or serving traces.",
],
}
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(
json.dumps(result, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
payload = args.output.read_bytes()
print(
json.dumps(
{
"output": str(args.output),
"sha256": hashlib.sha256(payload).hexdigest(),
"bytes": len(payload),
"sources": len(traced_sources),
"decoder_layers": 27,
"moe_gates": 26,
"target_route_decisions": result["summary"][
"target_route_decisions"
],
"load_seconds": load_seconds,
"forward_seconds": result["execution"][
"forward_seconds"
],
"peak_cuda_memory_allocated_bytes": result[
"execution"
]["peak_cuda_memory_allocated_bytes"],
},
ensure_ascii=False,
indent=2,
)
)
if __name__ == "__main__":
main()