219 lines
9.4 KiB
Python
219 lines
9.4 KiB
Python
"""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 : <cut N> 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'<cut \d+>', 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()
|