feat(erp-tune): NVFP4A16 serving pipeline, and the MoE landmine it uncovered
Merge + quantize path for turning the Gemma-4 26B-A4B ERP/RP LoRA into a
servable NVFP4A16 seat, plus a playbook entry for the defect found while
validating it.
The landmine (playbook §3.15): a `targets=["Linear"]` NVFP4 recipe silently
misses every MoE expert on this architecture. Gemma-4 stores each layer's 128
experts as two fused 3-D nn.Parameter tensors, not nn.Linear modules, so the
recipe resolves 205 of 427 modules and ZERO experts — 22.84 B params, 88.5% of
the model, left in BF16 with no warning. This is the same blind spot that
killed QLoRA here via bitsandbytes; the tool changed, the checkpoint layout did
not.
before linearize_moe: 427 Linears, 205 targeted, experts 0
after linearize_moe: 11,947 Linears, 11,725 targeted, experts 11,520
(30 layers x 128 experts x 3 projections)
llmcompressor's linearize_moe unfuses them; no registration needed because
Gemma-4 satisfies FusedExpertsProtocol structurally. Caught by an §4.1 dry run
that asserts the expert count before any GPU spend, which is now the documented
requirement rather than an optional step.
Scheme is NVFP4A16, deviating from the playbook's mixed-W4A4 default on
measured grounds: brokkr-smithy-dev benched the W4A4 quant of this checkpoint
at 12% on contradiction detection with CoT off against gen's 81%, the signature
of 4-bit input activations on a reasoning-dense task, and W4A4 KLD degrades
2-4x past ~10k ctx on sm_120. This is a 16,384-ctx RP seat. Marlin's prefill
cost is accepted.
Two further silent-failure guards, both from prior hard-won lessons:
- the merged model ships the UPSTREAM chat template, not the trainee base's
stale 365-line one, because training rendered through upstream and the
mismatch would present as a tuning failure
- calibration reads the run's own encode cache rather than re-tokenizing, which
sidesteps §3.14 (a fast tokenizer mutated by truncation=True and persisted by
save_pretrained clamps every prompt forever)
Merge-then-quantize rather than LoRA hot-swap, since hot-swap onto NVFP4 was a
silent no-op on vLLM 0.24.0 (#47639). merge_lora.py asserts sampled target
weights actually changed, so an inert adapter cannot ship as a tune.
This commit is contained in:
Executable
+110
@@ -0,0 +1,110 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Merge the ERP LoRA adapter into the bf16 base, producing servable weights.
|
||||
|
||||
WHY MERGE RATHER THAN HOT-SWAP. Serving NVFP4 base + LoRA at runtime was a
|
||||
silent no-op on vLLM 0.24.0 (#47639, proven quant-agnostic). Merging first
|
||||
sidesteps it entirely: the quantizer then sees ordinary bf16 weights and the
|
||||
served artifact needs no adapter machinery at all.
|
||||
|
||||
⚠⚠ CHAT TEMPLATE. The trainee base ships a STALE 365-line chat_template.jinja;
|
||||
upstream's is 390 lines. The harness deliberately trained through the UPSTREAM
|
||||
template (config key `chat_template_path`), so the merged model MUST ship that
|
||||
same upstream template. Shipping the base's own template here would be
|
||||
train/serve skew with no error — it presents as a tuning failure.
|
||||
|
||||
⚠ CPU merge. device_map=None keeps the 48 GiB on host RAM (566 GB total here)
|
||||
so this can run while GPU0 is training. Do not use device_map="auto".
|
||||
|
||||
⚠ Loader class. This checkpoint is Gemma4ForConditionalGeneration (vision +
|
||||
audio towers present). Loading it as a plain CausalLM is playbook §3.2 — a
|
||||
silent weight-load failure.
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import shutil
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
UPSTREAM_TEMPLATE = "/tank/aimodels/gemma4-26b-a4b-it-bf16/chat_template.jinja"
|
||||
|
||||
|
||||
def main() -> int:
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--base", required=True)
|
||||
ap.add_argument("--adapter", required=True)
|
||||
ap.add_argument("--out", required=True)
|
||||
ap.add_argument("--chat-template", default=UPSTREAM_TEMPLATE)
|
||||
a = ap.parse_args()
|
||||
|
||||
out = Path(a.out)
|
||||
if out.exists() and any(out.iterdir()):
|
||||
print(f"REFUSING: {out} exists and is non-empty", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
import torch
|
||||
from transformers import AutoTokenizer, Gemma4ForConditionalGeneration
|
||||
from peft import PeftModel
|
||||
|
||||
print(f"[merge] loading base on CPU: {a.base}", flush=True)
|
||||
model = Gemma4ForConditionalGeneration.from_pretrained(
|
||||
a.base, dtype=torch.bfloat16, device_map=None, trust_remote_code=True,
|
||||
)
|
||||
|
||||
# Count LoRA-target params before/after as a merge-actually-happened check.
|
||||
print(f"[merge] applying adapter: {a.adapter}", flush=True)
|
||||
before = {n: p.detach().clone() for n, p in model.named_parameters()
|
||||
if n.endswith("self_attn.q_proj.weight")
|
||||
and ".language_model.layers.0." in n}
|
||||
|
||||
model = PeftModel.from_pretrained(model, a.adapter, is_trainable=False)
|
||||
n_lora = sum(1 for n, _ in model.named_parameters() if "lora_" in n)
|
||||
print(f"[merge] adapter tensors seen: {n_lora}", flush=True)
|
||||
if n_lora == 0:
|
||||
print("REFUSING: adapter contributed 0 tensors", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
model = model.merge_and_unload()
|
||||
print("[merge] merged", flush=True)
|
||||
|
||||
# ⚠ Prove the merge changed weights. A no-op merge is the failure mode that
|
||||
# ships a base model wearing the tune's name, and nothing else would catch it.
|
||||
changed = 0
|
||||
for n, p in model.named_parameters():
|
||||
if n in before:
|
||||
if not torch.equal(p.detach(), before[n]):
|
||||
changed += 1
|
||||
if changed == 0:
|
||||
print("REFUSING: merge produced BIT-IDENTICAL weights on sampled "
|
||||
"LoRA-target modules — the adapter was inert or did not apply",
|
||||
file=sys.stderr)
|
||||
return 3
|
||||
print(f"[merge] verified {changed}/{len(before)} sampled target(s) changed", flush=True)
|
||||
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
print(f"[merge] saving to {out}", flush=True)
|
||||
model.save_pretrained(out, safe_serialization=True)
|
||||
|
||||
# Tokenizer straight from the base — never one that has been through
|
||||
# calibration (playbook §3.14).
|
||||
AutoTokenizer.from_pretrained(a.base, trust_remote_code=True).save_pretrained(out)
|
||||
|
||||
# ⚠ Ship the UPSTREAM chat template, matching what training rendered.
|
||||
src = Path(a.chat_template)
|
||||
if not src.exists():
|
||||
print(f"REFUSING: chat template missing at {src}", file=sys.stderr)
|
||||
return 4
|
||||
shutil.copy2(src, out / "chat_template.jinja")
|
||||
n_lines = len(src.read_text().splitlines())
|
||||
print(f"[merge] chat_template.jinja <- {src} ({n_lines} lines)", flush=True)
|
||||
|
||||
tj = out / "tokenizer.json"
|
||||
if tj.exists() and json.loads(tj.read_text()).get("truncation"):
|
||||
print("REFUSING: shipped tokenizer carries a truncation cap", file=sys.stderr)
|
||||
return 5
|
||||
|
||||
print(f"[merge] DONE -> {out}", flush=True)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user