traj-mc / code /common.py
ttishere's picture
Publish code/common.py
950e94a verified
Raw History Blame Contribute Delete
23.3 kB
"""
common.py -- shared utilities for the Traj-MC vs SVD-LLM main experiment.
Only three objects exist in this experiment (see README):
Dense (REF) | BASE (SVD-LLM t=0) | OURS (Traj-MC t~U[0,1]).
Everything here is arm-agnostic. The ONLY module allowed to differ between
BASE and OURS is calib/build_calib.py (the noise switch). Compression, eval,
and analysis share a single code path across all arms.
"""
import os
import sys
import json
import hashlib
import subprocess
import torch
import torch.nn as nn
# ── fixed constants (never overridden per-arm) ────────────────────────────────
MASK_ID = 126336
# Main experiment target model = LLaDA-8B-BASE (see README section 1):
# - same model + eval protocol as Sink-Aware / LLaDA official EVAL.md,
# - Base has a conditional-likelihood (ppl) eval path; Instruct does not.
DEFAULT_MODEL_PATH = os.path.expanduser("~/LLaDA-8B-Base")
DEFAULT_MODEL_ID = "GSAI-ML/LLaDA-8B-Base"
# Instruct table (see LMEVAL_TASKS_INSTRUCT). Same architecture and same 224 target
# linears as Base; its tokenizer is a strict superset (adds only the three chat tokens
# at 126346-126348), which is why the Base calibration tensors are reused verbatim.
INSTRUCT_MODEL_PATH = os.path.expanduser("~/LLaDA-8B-Instruct")
INSTRUCT_MODEL_ID = "GSAI-ML/LLaDA-8B-Instruct"
SEQLEN = 2048
NSAMPLES = 1400
# The 7 target Linear suffixes inside each LLaDALlamaBlock.
# block_type=llama -> separate q/k/v; 32 blocks x 7 = 224 target linears.
ATTN_SUFFIXES = ("q_proj", "k_proj", "v_proj", "attn_out")
MLP_SUFFIXES = ("ff_proj", "up_proj", "ff_out")
ALL_SUFFIXES = ATTN_SUFFIXES + MLP_SUFFIXES
# A Linear is a compression target iff its name lives inside transformer.blocks.
# This deliberately EXCLUDES the top-level unembed head model.transformer.ff_out
# (which has the same leaf name 'ff_out' but no '.blocks.' in its path).
BLOCKS_MARKER = ".blocks."
# ── OFFICIAL LLaDA-8B-Base lm-eval per-task protocol (single source of truth) ──
# Verified byte-for-byte from scripts/eval_llada_lm_eval.sh (== Sink-Aware
# eval_llada.sh). Per README section 0: on any protocol conflict, follow LLaDA
# official (SVD-LLM is authoritative ONLY for the compression algorithm). The
# PRIMARY metric is pre-registered per README section 2 and is the ONLY column
# McNemar runs; acc/acc_norm are both logged for the appendix.
# fs=num_fewshot, cfg=classifier-free guidance, mc=Monte-Carlo iterations
LMEVAL_TASKS = {
# ppl (conditional likelihood) tasks
"arc_challenge": dict(fs=0, cfg=0.5, mc_num=128, batch_size=8, metric="acc", gen=False),
"arc_easy": dict(fs=0, cfg=0.5, mc_num=128, batch_size=8, metric="acc", gen=False),
"hellaswag": dict(fs=0, cfg=0.5, mc_num=128, batch_size=8, metric="acc_norm", gen=False),
"piqa": dict(fs=0, cfg=0.5, mc_num=128, batch_size=8, metric="acc_norm", gen=False),
"winogrande": dict(fs=5, cfg=0.0, mc_num=128, batch_size=8, metric="acc", gen=False),
"mmlu": dict(fs=5, cfg=0.0, mc_num=1, batch_size=1, metric="acc", gen=False),
# gen (conditional generation) task -- Base uses block_length == gen_length
# (full diffusion, NOT the Instruct block-diffusion block=8). EVAL.md 256/256
# reproduces GSM8K 70.0. num_fewshot=5 = lm-eval gsm8k task default.
"gsm8k": dict(fs=5, gen_length=256, steps=256, block_length=256,
batch_size=1, metric="exact_match", gen=True),
# SVAMP: math word problems (same family as gsm8k). Custom lm-eval task
# (eval/lm_tasks/svamp.yaml), reuses gsm8k's number-extraction + gen config.
# gen_length/steps/block_length identical to gsm8k so the 3 arms are comparable.
"svamp": dict(fs=5, gen_length=256, steps=256, block_length=256,
batch_size=1, metric="exact_match", gen=True),
# MATH-500: competition math (HuggingFaceH4/MATH-500). Custom lm-eval task
# (eval/lm_tasks/math500.yaml) reusing minerva_math boxed extraction + is_equiv.
# Longer generation than gsm8k (512) for multi-step solutions.
"math500": dict(fs=4, gen_length=512, steps=512, block_length=512,
batch_size=1, metric="exact_match", gen=True),
# code generation, pass@1 (EXECUTES generated code). STOCK lm-eval tasks
# (humaneval.yaml / mbpp.yaml, unsafe_code:true) -> build_cmd adds
# --confirm_run_unsafe_code and env sets HF_ALLOW_CODE_EVAL=1. code=True marks them.
# gen config = GSM8K reference (256/256/256 full diffusion), per README
# pre-reg "参考 GSM8K". code completions are short, 256 gen_length is ample.
"humaneval": dict(fs=0, gen_length=256, steps=256, block_length=256,
batch_size=1, metric="pass@1", gen=True, code=True),
"mbpp": dict(fs=3, gen_length=256, steps=256, block_length=256,
batch_size=1, metric="pass_at_1", gen=True, code=True),
# IFEval: instruction-following (google/IFEval, 541 prompts). Generation task,
# rule-based verifier (lm-eval ifeval). Primary = prompt-level strict accuracy.
"ifeval": dict(fs=0, gen_length=512, steps=512, block_length=512,
batch_size=1, metric="prompt_level_strict_acc", gen=True),
# BBH: BIG-Bench Hard, 3-shot DIRECT answer (no CoT), exact_match (lukaemon/
# bbh, 27 subtasks, 6511 docs). LOCAL copy of lm-eval's bbh_fewshot group
# (eval/lm_tasks/bbh_llada/) with one fix: a remove_whitespace filter --
# prompts end in "A:" so continuations carry a leading space that stock
# strict exact_match scores 0 (dense scored 0/54 without it). Few-shot
# exemplars fixed in the yamls. Targets are short -- measured max 48 tokens
# (word_sorting), all other subtasks <=7 -- so 64/64/64 full diffusion
# (block==gen, Base convention) suffices without truncation.
"bbh_llada": dict(fs=3, gen_length=64, steps=64, block_length=64,
batch_size=1, metric="exact_match", gen=True),
}
# ── OFFICIAL LLaDA-8B-INSTRUCT protocol ───────────────────────────────────────
# Source: evaluation/EVAL.md, the Instruct table. Two things differ structurally
# from the Base protocol above and are NOT stylistic choices:
#
# 1. Instruct is "evaluated using only conditional generation" -- it has NO ppl
# path at all. So MMLU is the GENERATIVE lm-eval task (mmlu_generative,
# output_type: generate_until, prompt ends in "Answer:", target "A".."D"),
# not the multiple-choice likelihood task the Base row uses.
# 2. block_length == gen_length here TOO. EVAL.md is explicit that the paper's
# Tab.1/Tab.2 Instruct numbers use "pure diffusion sampling without any
# autoregressive elements". Block diffusion (block=8 on GSM8K, 64 on Math) is
# a SEPARATE follow-up experiment that helps those two tasks and, in the
# authors' words, lowers accuracy elsewhere. We follow the headline setting.
#
# gen_length/logits_eos_inf/confidence_eos_eot_inf are copied verbatim from that
# table. steps == gen_length keeps our one-token-per-step convention (EVAL.md
# tabulates gen_length and block_length but not steps).
#
# num_fewshot is NOT given by EVAL.md; we keep the Base row's values so the two
# tables stay comparable and so each task keeps its lm-eval default.
#
# bbh_llada and ifeval have NO official Instruct reference point -- they are not
# in the EVAL.md table at all. Their gen config is inherited from our Base row and
# their EOS switches are left off; both are marked no_official_target so the dense
# gate does not pretend to validate them.
LMEVAL_TASKS_INSTRUCT = {
"mmlu_generative": dict(fs=5, gen_length=3, steps=3, block_length=3,
batch_size=1, metric="exact_match", gen=True,
logits_eos_inf=False, confidence_eos_eot_inf=False),
"gsm8k": dict(fs=5, gen_length=512, steps=512, block_length=512,
batch_size=1, metric="exact_match", gen=True,
logits_eos_inf=False, confidence_eos_eot_inf=True),
# LOCAL chat-aware variants (eval/lm_tasks/code_instruct/). The stock lm-eval
# humaneval/mbpp are raw-COMPLETION tasks and score ~0 on a chat model: it answers
# in prose + a markdown fence, which is not valid Python once concatenated onto the
# signature, and humaneval's keyword stop list truncates the fence mid-answer.
# Measured 2026-08-18 on dense Instruct: stock humaneval pass@1 = 0/4.
"humaneval_instruct": dict(fs=0, gen_length=512, steps=512, block_length=512,
batch_size=1, metric="pass@1", gen=True, code=True,
logits_eos_inf=True, confidence_eos_eot_inf=False),
"mbpp_instruct": dict(fs=3, gen_length=256, steps=256, block_length=256,
batch_size=1, metric="pass_at_1", gen=True, code=True,
logits_eos_inf=False, confidence_eos_eot_inf=True),
# no official Instruct setting -- inherited from the Base row, EOS switches off
"bbh_llada": dict(fs=3, gen_length=64, steps=64, block_length=64,
batch_size=1, metric="exact_match", gen=True,
logits_eos_inf=False, confidence_eos_eot_inf=False,
no_official_target=True),
"ifeval": dict(fs=0, gen_length=512, steps=512, block_length=512,
batch_size=1, metric="prompt_level_strict_acc", gen=True,
logits_eos_inf=False, confidence_eos_eot_inf=False,
no_official_target=True),
}
FAMILIES = ("base", "instruct")
def tasks_for(family="base"):
"""Per-task protocol table for a model family. Base is the default so every
pre-existing caller (the finished Base table) keeps its exact behaviour."""
if family == "base":
return LMEVAL_TASKS
if family == "instruct":
return LMEVAL_TASKS_INSTRUCT
raise ValueError(f"unknown family {family!r}; expected one of {FAMILIES}")
def primary_metric(task, family="base"):
return tasks_for(family)[task]["metric"]
# ── provenance ────────────────────────────────────────────────────────────────
def git_hash(short=True):
"""Short git hash of the repo, for stamping every artifact."""
try:
repo = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
out = subprocess.check_output(
["git", "-C", repo, "rev-parse"] + (["--short"] if short else []) + ["HEAD"],
stderr=subprocess.DEVNULL,
)
return out.decode().strip()
except Exception:
return "nogit"
def sha256_ids(ids):
"""Deterministic hash of a 1-D token-id sequence (python list or tensor)."""
if torch.is_tensor(ids):
ids = ids.detach().cpu().to(torch.int64).tolist()
b = ",".join(str(int(x)) for x in ids).encode()
return hashlib.sha256(b).hexdigest()[:16]
def sha256_text(text):
return hashlib.sha256(text.encode("utf-8")).hexdigest()[:16]
# ── model loading ─────────────────────────────────────────────────────────────
def load_model(model_path=DEFAULT_MODEL_PATH, dtype=torch.bfloat16, device="cuda"):
"""Load the dense LLaDA model + tokenizer. Used identically by all arms."""
from transformers import AutoTokenizer, AutoModel
model = (
AutoModel.from_pretrained(model_path, trust_remote_code=True, torch_dtype=dtype)
.to(device)
.eval()
)
tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
return model, tokenizer
# ── target-linear enumeration (the 224 block-internal Linears) ────────────────
def is_target_linear(name, module, layer_type="all"):
if not isinstance(module, nn.Linear):
return False
if BLOCKS_MARKER not in name:
return False # excludes the unembed head model.transformer.ff_out
suffixes = ALL_SUFFIXES if layer_type == "all" else (
ATTN_SUFFIXES if layer_type == "attn" else MLP_SUFFIXES
)
return name.endswith(suffixes)
def iter_target_linears(model, layer_type="all"):
"""Yield (name, module) for every compression-target Linear."""
for name, module in model.named_modules():
if is_target_linear(name, module, layer_type):
yield name, module
def count_target_linears(model, layer_type="all"):
return sum(1 for _ in iter_target_linears(model, layer_type))
def get_parent_attr(model, name):
parts = name.split(".")
parent = model
for p in parts[:-1]:
parent = getattr(parent, p)
return parent, parts[-1]
# ── rank formula (identical to official SVD-LLM) ──────────────────────────────
def rank_from_ratio(ratio, out_dim, in_dim):
"""
Official: k = int(out*in*ratio / (out+in)). We keep max(1,.) as a guard;
at the experiment's dims (d=4096/12288) and ratios (0.7/0.8) it never binds.
"""
return max(1, int(ratio * out_dim * in_dim / (out_dim + in_dim)))
def is_compression_beneficial(k, out_dim, in_dim):
"""Skip layers where the low-rank factorization would not save params."""
return k * (out_dim + in_dim) < out_dim * in_dim
# ── FULL-MODEL compression ratio (ζ–Ήζ‘ˆB, the LoRAP convention) ─────────────────
# The reported ratio is a WHOLE-CHECKPOINT number: the denominator is every
# parameter in the model (embedding + unembed head + norms + biases + every
# transformer weight), not only the subset we factorise. Only the target Linears
# inside the transformer blocks are compressible; the rest is kept verbatim but
# STILL COUNTED in the denominator. To land a full-model retention of
# Ratio_model, the target set must therefore be squeezed HARDER than Ratio_model:
#
# P_fixed = Param_total - P_targets (uncompressible, kept)
# budget = Ratio_model * Param_total - P_fixed (params left for targets)
# Ratio_layer = budget / P_targets (uniform over targets)
#
# This is LoRAP's rule -- "only transformer layers are compressed, embedding and
# lm_head are untouched; to reach a specified model-level compression rate the
# layers must take a HIGHER layer-level compression rate" -- written in this
# repo's RETENTION convention rather than LoRAP's removal convention.
#
# ⚠ SIGN/DIRECTION: `ratio` here always means "fraction of parameters KEPT", so
# Ratio_layer comes out BELOW Ratio_model (e.g. 0.770 < 0.80). Read as a removal
# rate that is 1-0.770 = 23.0% > 20.0%, i.e. precisely LoRAP's "higher layer-level
# compression rate". Lower retention == stronger compression; same statement.
#
# Because two different models have different P_fixed shares (LLaDA's untied
# 126464-row embedding + head is 12.9% of the checkpoint, Dream's 152064-row pair
# is 14.3%), a SHARED Ratio_layer would mean different real compression strengths.
# Fixing Ratio_model instead and deriving Ratio_layer per model is the whole point
# of this convention: the two models are then compared at equal true strength.
def param_inventory(model, layer_type="all"):
"""
Whole-checkpoint parameter census, split into the compressible target set and
the fixed remainder. Works on a `meta`-device model (numel needs no storage),
so the pre-flight report runs on a login node with no GPU and no weight load.
P_targets counts ONLY the 2-D weight matrices of the target Linears -- the
part low-rank factorisation actually replaces. A target Linear's bias (Dream's
q/k/v_proj) is re-attached verbatim, so it lands in P_fixed, where it belongs.
"""
param_total = int(sum(p.numel() for p in model.parameters()))
targets = {name: (mod.out_features, mod.in_features)
for name, mod in iter_target_linears(model, layer_type)}
p_targets = int(sum(o * i for o, i in targets.values()))
return {
"param_total": param_total,
"p_targets": p_targets,
"p_fixed": param_total - p_targets,
"n_targets": len(targets),
"targets": targets,
}
def realize_ranks(targets, layer_ratio):
"""
Apply the official integer rank formula at `layer_ratio` to every target and
report the params the target set would then occupy.
A layer whose rank-k form would not save params is SKIPPED (stays dense) --
the same rule compress.py applies -- so its full out*in is charged here too.
Returns (ranks {name: k or None}, realized_params).
"""
ranks, tot = {}, 0
for name, (o, i) in targets.items():
k = rank_from_ratio(layer_ratio, o, i)
if is_compression_beneficial(k, o, i):
ranks[name] = int(k)
tot += k * (o + i)
else:
ranks[name] = None # stays dense
tot += o * i
return ranks, int(tot)
def rank_plan(model, model_ratio, layer_type="all"):
"""
Derive the per-layer ranks that hit a FULL-MODEL retention of `model_ratio`.
ONE plan per model, used for every benchmark and both arms -- the ranks depend
only on (architecture, layer_type, model_ratio), never on the calibration data,
so BASE and OURS are guaranteed rank-identical by construction.
`layer_ratio` is uniform across targets (LoRAP): the rank formula makes each
target's retained fraction k*(o+i)/(o*i) ~= layer_ratio, so a single scalar
spends the budget proportionally. The only slack is int() flooring, which is
bounded by sum(o+i) / Param_total (~5e-4 here) and always errs toward MORE
compression, so the reported effective ratio never overstates compression.
"""
inv = param_inventory(model, layer_type)
p_total, p_fixed, p_tgt = inv["param_total"], inv["p_fixed"], inv["p_targets"]
budget = model_ratio * p_total - p_fixed
if budget <= 0:
raise ValueError(
f"model_ratio={model_ratio} is unreachable: the uncompressible part "
f"(P_fixed={p_fixed:,} = {p_fixed/p_total:.2%} of the checkpoint) already "
f"exceeds the whole-model budget {model_ratio*p_total:,.0f}. The lowest "
f"attainable full-model ratio is {p_fixed/p_total:.4f} (targets -> rank 0)."
)
layer_ratio = budget / p_tgt
ranks, realized = realize_ranks(inv["targets"], layer_ratio)
post_total = p_fixed + realized
plan = {
"ratio_mode": "model",
"model_ratio": model_ratio,
"layer_ratio": layer_ratio,
"layer_type": layer_type,
"param_total": p_total,
"p_fixed": p_fixed,
"p_targets": p_tgt,
"p_fixed_share": p_fixed / p_total,
"target_budget_params": int(round(budget)),
"realized_target_params": realized,
"post_compression_total_params": post_total,
"effective_model_ratio": post_total / p_total,
"target_kept_fraction": realized / p_tgt,
"n_targets": inv["n_targets"],
"n_would_skip": sum(1 for v in ranks.values() if v is None),
"ranks": ranks,
}
plan["rank_plan_hash"] = rank_plan_hash(ranks)
return plan
def rank_plan_hash(ranks):
"""Order-independent fingerprint of a rank plan; one field for the BASE/OURS
gate to compare instead of eyeballing 224 integers."""
blob = ";".join(f"{n}:{ranks[n]}" for n in sorted(ranks))
return hashlib.sha256(blob.encode()).hexdigest()[:16]
# Pre-registered: |effective - nominal| must be <= this.
# NB this tolerance is calibrated for the 7-8B checkpoints in this experiment, where
# the int()-flooring slack is ~5e-4 (30x margin). The slack scales as
# sum(out+in)/Param_total, so a TOY model can legitimately exceed 0.005 without any
# bug -- do not "fix" the allocator against a small-model failure of this gate.
EFFECTIVE_RATIO_TOL = 0.005
def assert_effective_ratio(plan, tol=EFFECTIVE_RATIO_TOL):
"""Hard gate: a plan whose measured full-model ratio misses the nominal one is
an allocation bug, not a rounding artifact. Abort before spending GPU hours."""
eff, nom = plan["effective_model_ratio"], plan["model_ratio"]
if abs(eff - nom) > tol:
raise RuntimeError(
f"RANK ALLOCATION BUG: effective_model_ratio={eff:.6f} deviates from "
f"nominal model_ratio={nom} by {abs(eff-nom):.2e} > tol={tol}."
)
return True
# ── low-rank replacement module (packing-only diff vs official SVD_Llama*) ─────
class LowRankLinear(nn.Module):
"""
forward(x) = A(B(x)). B: in->k, A: k->out. Mathematically identical to the
official svd_u @ (svd_v @ x); only the module packing differs (see port_diff).
"""
def __init__(self, A, B, bias=None):
super().__init__()
k, in_dim = B.shape
out_dim, _ = A.shape
self.B = nn.Linear(in_dim, k, bias=False)
self.A = nn.Linear(k, out_dim, bias=bias is not None)
self.B.weight.data = B.to(torch.bfloat16)
self.A.weight.data = A.to(torch.bfloat16)
if bias is not None:
self.A.bias.data = bias.to(torch.bfloat16)
def forward(self, x):
return self.A(self.B(x))
# ── lm_head hard guard ────────────────────────────────────────────────────────
def get_output_head(model):
"""Return the module used as the output projection (unembed)."""
try:
return model.get_output_embeddings()
except Exception:
return getattr(model.model.transformer, "ff_out", None)
def assert_head_dense(model):
"""
Hard defense line: the unembed head must remain a plain Linear/Embedding.
If it was ever replaced by our LowRankLinear/Identity, abort.
"""
head = get_output_head(model)
if isinstance(head, (LowRankLinear, nn.Identity)):
raise RuntimeError(
"lm_head DEFENSE TRIGGERED: output head was replaced by "
f"{type(head).__name__}; it MUST stay dense."
)
if head is not None and not isinstance(head, (nn.Linear, nn.Embedding)):
raise RuntimeError(
f"lm_head DEFENSE: unexpected head type {type(head).__name__}"
)
return True
# ── json helpers ──────────────────────────────────────────────────────────────
def dump_json(obj, path):
os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True)
with open(path, "w") as f:
json.dump(obj, f, indent=2)
def load_json(path):
with open(path) as f:
return json.load(f)
def peak_rss_gb():
"""Peak resident set size of this process in GB (Linux ru_maxrss is KB)."""
import resource
return resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / (1024.0 * 1024.0)