"""LoRA block distillation for Qwen3-14B on one H200. ONE frozen base model plays both roles via adapter toggling (memory: ~1 model, not 2): teacher = LoRA disabled + FULL attention (no_grad) student = LoRA enabled + BLOCK attention (grad, only LoRA params train) Layout per example: [system][block_1..n-1 isolated (+4 sink tokens each)][last segment attends to all]. Loss on the last segment: KL(teacher_full || student_block) + damage-weighted CE(student_block, gold), w = max(CE(teacher_block) - CE(teacher_full), 0) * alpha + beta. Data (real, segmented): LongbenchSeg / LoCoMoSeg : chunks=blocks, last segment = "Question: {q}\nAnswer: {a}" (CE on answer) SemanticSeg : markers -> blocks, last segment = final block (CE on whole block) Run: source env.sh && HF_HOME=/projects/bdjx/hxia3/hf_cache_proj \ python scripts/distill_lora.py --model Qwen/Qwen3-14B --steps 3000 --out checkpoints/qwen3-14b-blockdist-lora """ from __future__ import annotations import argparse import glob import json import random import re import torch import torch.nn.functional as F from peft import LoraConfig, get_peft_model from transformers import AutoModelForCausalLM, AutoTokenizer NEG = -1e9 LB = "/work/hdd/bdjx/hxia3/hf_cache/hub/datasets--Syon-Li--LongbenchSeg/snapshots/*/longbench_segmented.jsonl" LOCOMO = "/work/hdd/bdjx/hxia3/hf_cache/hub/datasets--Syon-Li--LoCoMoSeg/snapshots/*/locomo_segmented.json" SEMSEG = "/work/hdd/bdjx/hxia3/hf_cache/hub/datasets--Syon-Li--SemanticSeg/snapshots/*/cut_*.jsonl" def _tok(tok, text, special=False): return tok(text, add_special_tokens=special)["input_ids"] def emit(tok, sink_ids, system, blocks, last_seg, ce_from_ratio=0.0): """Return ids, blk, sink, kl_start, ce_start for one training example.""" ids, blk, sink = [], [], [] def add(text, b, is_sink=False, special=False): t = [text] if is_sink else _tok(tok, text, special) for x in t: ids.append(x); blk.append(b); sink.append(is_sink) if system: for x in _tok(tok, system + "\n", special=True): ids.append(x); blk.append(-1); sink.append(False) for bi, bt in enumerate(blocks): for s in sink_ids: ids.append(s); blk.append(bi); sink.append(True) add("\n" + bt.strip() + "\n", bi) kl_start = len(ids) seg_ids = _tok(tok, last_seg, special=False) ce_start = kl_start + int(len(seg_ids) * ce_from_ratio) for x in seg_ids: ids.append(x); blk.append(-2); sink.append(False) return ids, blk, sink, kl_start, ce_start def group_blocks(segs, target_chars=1000, max_blocks=8): """Merge tiny segments into ~target_chars blocks (char proxy for ~256 tok); cap count.""" out, cur, cur_n = [], [], 0 for c in segs: cur.append(c); cur_n += len(c) if cur_n >= target_chars: out.append(" ".join(cur)); cur, cur_n = [], 0 if len(out) >= max_blocks: break if cur and len(out) < max_blocks: out.append(" ".join(cur)) return out def load_data(tok, sink_ids, max_len, max_per_source, seed): rng = random.Random(seed) data = [] CHAR = max_len * 4 # ~4 chars/token budget for char prefilters # LongbenchSeg (QA) — most examples are long; cheap char prefilter before json.loads lb = 0 for line in open(glob.glob(LB)[0]): if lb >= max_per_source: break if len(line) > CHAR * 40: # skip only pathologically long docs; blocks are capped below continue try: r = json.loads(line) except Exception: continue if not r.get("chunks") or not r.get("answers") or len(r["chunks"]) < 3: continue blocks = [c[:CHAR // 4] for c in r["chunks"][1:9]] # cap each block's chars ex = emit(tok, sink_ids, r["chunks"][0][:600], blocks, f"Question: {r['input']}\nAnswer: {r['answers'][0]}") q_ids = _tok(tok, f"Question: {r['input']}\nAnswer:") ex = (ex[0], ex[1], ex[2], ex[3], ex[3] + len(q_ids)) if len(ex[0]) <= max_len and ex[4] < len(ex[0]): data.append((*ex, "lb")); lb += 1 print(f" LongbenchSeg: {lb} examples", flush=True) # SemanticSeg (text) — big diverse corpus, bounded scan per file, char-based blocks files = sorted(glob.glob(SEMSEG)); rng.shuffle(files) per_file = max(1, max_per_source // max(1, len(files))) for f in files: got = 0; scanned = 0 for line in open(f): if got >= per_file or scanned > per_file * 20: break scanned += 1 if len(line) > CHAR * 2: continue try: txt = json.loads(line).get("cut_item", [{}])[0].get("txt_marker", "") except Exception: continue segs = [s for s in re.split(r'', txt) if s.strip()] if len(segs) < 4: continue blocks = group_blocks(segs, target_chars=1000, max_blocks=8) if len(blocks) < 3: continue ex = emit(tok, sink_ids, "", blocks[:-1], blocks[-1]) if len(ex[0]) <= max_len and ex[3] < len(ex[0]) - 1: data.append((*ex, "ss")); got += 1 print(f" {f.split('cut_')[-1].split('.')[0]}: {got}", flush=True) rng.shuffle(data) return data def masks(blk, dev): n = len(blk); b = torch.tensor(blk, device=dev) causal = torch.tril(torch.ones(n, n, dtype=torch.bool, device=dev)) bi, bj = b.view(n, 1), b.view(1, n) full = torch.where(causal, 0.0, NEG).view(1, 1, n, n).float() allowed = ((bj == -1) | (bi == -2) | (bi == bj)) & causal block = torch.where(allowed, 0.0, NEG).view(1, 1, n, n).float() return full, block def ce_span(logits, ids, start): lp = F.log_softmax(logits[0, start - 1:-1].float(), -1) return -lp.gather(-1, ids[0, start:].view(-1, 1)).squeeze(-1) def main(): ap = argparse.ArgumentParser() ap.add_argument("--model", default="Qwen/Qwen3-14B") ap.add_argument("--steps", type=int, default=3000) ap.add_argument("--lr", type=float, default=1e-4) ap.add_argument("--accum", type=int, default=8) ap.add_argument("--rank", type=int, default=16) ap.add_argument("--n-sinks", type=int, default=4) ap.add_argument("--alpha", type=float, default=0.3) ap.add_argument("--beta", type=float, default=0.1) ap.add_argument("--max-len", type=int, default=1536) ap.add_argument("--max-per-source", type=int, default=6000) ap.add_argument("--seed", type=int, default=0) ap.add_argument("--out", default="/projects/bdjx/hxia3/lazy2/checkpoints/qwen3-14b-blockdist-lora") ap.add_argument("--save-every", type=int, default=500) args = ap.parse_args() dev = "cuda"; rng = random.Random(args.seed) print(f"loading {args.model} (bf16, eager, grad-checkpoint) ...") tok = AutoTokenizer.from_pretrained(args.model) base = AutoModelForCausalLM.from_pretrained( args.model, dtype=torch.bfloat16, attn_implementation="eager").to(dev) base.gradient_checkpointing_enable() base.enable_input_require_grads() lora = LoraConfig(r=args.rank, lora_alpha=args.rank * 2, lora_dropout=0.0, target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"]) model = get_peft_model(base, lora) model.print_trainable_parameters() sink_ids = tok("\n", add_special_tokens=False)["input_ids"] * args.n_sinks print("loading data ...") data = load_data(tok, sink_ids, args.max_len, args.max_per_source, args.seed) from collections import Counter print(f"examples: {len(data)} | sources: {dict(Counter(d[5] for d in data))}") opt = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad], lr=args.lr, betas=(0.9, 0.95)) opt.zero_grad() model.train() run = 0.0 for step in range(1, args.steps + 1): ids, blk, sink, kl0, ce0, src = data[rng.randrange(len(data))] t = torch.tensor([ids], device=dev) fm, bm = masks(blk, dev) with torch.no_grad(), model.disable_adapter(): # TEACHER (full attn, no LoRA) tf = model(input_ids=t, attention_mask=fm).logits tf_ce = ce_span(tf, t, ce0) tb_ce = ce_span(model(input_ids=t, attention_mask=bm).logits, t, ce0) w = torch.clamp(tb_ce - tf_ce, min=0) * args.alpha + args.beta sb = model(input_ids=t, attention_mask=bm).logits # STUDENT (block attn, LoRA) kl = F.kl_div(F.log_softmax(sb[0, kl0:].float(), -1), F.log_softmax(tf[0, kl0:].float(), -1), reduction="none", log_target=True).sum(-1).mean() wce = (w * ce_span(sb, t, ce0)).mean() loss = (kl + wce) / args.accum loss.backward(); run += loss.item() * args.accum if step % args.accum == 0: torch.nn.utils.clip_grad_norm_([p for p in model.parameters() if p.requires_grad], 1.0) opt.step(); opt.zero_grad() if step % 10 == 0: print(f"step {step:5d}/{args.steps} | loss {run/10:.3f} | KL {kl.item():.3f} | wCE {wce.item():.3f}", flush=True) run = 0.0 if step % args.save_every == 0: model.save_pretrained(args.out); print(f" saved adapter -> {args.out}", flush=True) model.save_pretrained(args.out) tok.save_pretrained(args.out) print(f"done. LoRA adapter at {args.out}") if __name__ == "__main__": main()