初始化项目,由ModelHub XC社区提供模型
Model: hxia7/qwen3-4b-blockdist Source: Original Platform
This commit is contained in:
218
eval/scripts/distill_lora.py
Normal file
218
eval/scripts/distill_lora.py
Normal file
@@ -0,0 +1,218 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user