File size: 23,346 Bytes
950e94a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
"""
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)