File size: 5,072 Bytes
df43f42 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 | """
clankerDiffusion — RAG continuation fine-tune.
Loads the latest base checkpoint and continues training on the RAG corpus
(data/rag_corpus.txt) so the model actually learns to (a) emit
<tool name="retrieve">q</tool> and use the <result>, and (b) read
<context>...</context> injected mid-response. Run this AFTER the base
24h run (or anytime you have a checkpoint).
python rag_finetune.py --hours 3 --batch 24 --ckpt-every 500
It reuses the hybrid AR/DIFF loss from train.py's logic (self-contained copy
below so it doesn't import the training script's __main__).
"""
import os, json, argparse, time, random
import torch
import torch.nn.functional as F
from torch.utils.data import IterableDataset
from model import YKDiff
from tokenizer import YKTokenizer
from rag import DEFAULT_INDEX, KnowledgeBase
HERE = os.path.dirname(os.path.abspath(__file__))
DATADIR = os.path.join(HERE, "data")
CKPTDIR = os.path.join(HERE, "checkpoints")
CORPUS = os.path.join(DATADIR, "rag_corpus.txt")
SEQ = 1024
# ---- hybrid loss (mirrors train.py) ---------------------------------------
def ar_loss(logits, ids, mask):
sl = ids[:, 1:]; lm = logits[:, :-1]
return F.cross_entropy(lm.reshape(-1, lm.size(-1)), sl.reshape(-1), reduction="none").mean()
def diff_loss(logits, ids, mask, t):
B, L, V = logits.shape
ce = F.cross_entropy(logits.reshape(-1, V), ids.reshape(-1), reduction="none").view(B, L)
w = (mask * (1.0 - t)[:, None]).reshape(B, L)
return (ce * w).sum() / w.sum().clamp_min(1.0)
# ---- corpus -> token stream ----------------------------------------------
class RagTokens(IterableDataset):
def __init__(self, path, tok, seq):
self.lines = [l.rstrip("\n") for l in open(path, encoding="utf-8") if l.strip()]
self.tok = tok
self.seq = seq
self.buf = []
def _fill(self):
while len(self.buf) < self.seq + 1:
line = random.choice(self.lines)
self.buf.extend(self.tok.encode(line))
def __iter__(self):
while True:
self._fill()
chunk = self.buf[: self.seq + 1]
self.buf = self.buf[self.seq:]
yield torch.tensor(chunk, dtype=torch.long)
def pack_batch(ds, n):
ids = torch.stack([next(iter(ds)) for _ in range(n)]) # [B, L+1]
mask = (ids != ds.tok.pad_id).long()
return ids, mask
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--hours", type=float, default=3.0)
ap.add_argument("--batch", type=int, default=24)
ap.add_argument("--ckpt-every", type=int, default=500)
ap.add_argument("--log-every", type=int, default=25)
ap.add_argument("--ckpt", default=None)
a = ap.parse_args()
tok = YKTokenizer.load(os.path.join(DATADIR, "tokenizer.json"))
if a.ckpt is None:
ck = sorted([f for f in os.listdir(CKPTDIR) if f.endswith(".pt")])
if not ck:
raise SystemExit("no checkpoint")
a.ckpt = os.path.join(CKPTDIR, ck[-1])
sd = torch.load(a.ckpt, map_location="cuda")
model = YKDiff(sd["cfg"]).cuda().train()
model.load_state_dict(sd["model"])
step0 = sd.get("step", 0)
print(f"[rag-ft] resume {a.ckpt} step={step0} params={sum(p.numel() for p in model.parameters())/1e6:.1f}M")
opt = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.01, betas=(0.9, 0.95))
sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=max(1, int(a.hours * 60)), eta_min=1e-5)
ds = RagTokens(CORPUS, tok, SEQ)
t0 = time.time()
limit = a.hours * 3600
step = step0
while time.time() - t0 < limit:
step += 1
ids, mask = pack_batch(ds, a.batch)
x = ids[:, :-1].cuda(); y = ids[:, 1:].cuda(); m = mask[:, 1:].cuda()
if random.random() < 0.5:
mode = torch.zeros(a.batch, dtype=torch.long, device="cuda")
t = None
loss = ar_loss(model(x, mode), y, m)
else:
mode = torch.ones(a.batch, dtype=torch.long, device="cuda")
t = torch.rand(a.batch, device="cuda")
r = (0.1 + 0.9 * torch.rand(a.batch, device="cuda"))
inps = torch.where(torch.rand_like(y.float()) < r[:, None], tok.mask_id, y)
loss = diff_loss(model(inps, mode, t=t), y, m, t)
opt.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step(); sched.step()
if step % a.log_every == 0:
print(f"[rag-ft] step {step} loss={loss.item():.3f} t={ (time.time()-t0)/60:.1f}m", flush=True)
if step % a.ckpt_every == 0:
torch.save({"model": model.state_dict(), "cfg": sd["cfg"], "step": step,
"rag_ft": True}, os.path.join(CKPTDIR, f"ragft_{step:06d}.pt"))
print(f"[rag-ft] saved ragft_{step:06d}.pt", flush=True)
torch.save({"model": model.state_dict(), "cfg": sd["cfg"], "step": step, "rag_ft": True},
os.path.join(CKPTDIR, "ragft_final.pt"))
print("[rag-ft] done", flush=True)
if __name__ == "__main__":
main()
|