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