Files
qwen3-4b-blockdist/eval/scripts/distill_lora.py
ModelHub XC b315e39b60 初始化项目,由ModelHub XC社区提供模型
Model: hxia7/qwen3-4b-blockdist
Source: Original Platform
2026-07-27 06:09:10 +08:00

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()