File size: 5,039 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
131
132
133
134
135
136
"""
clankerDiffusion — inference.

Three generation modes, all from the SAME weights:
  ar      normal autoregressive next-token decoding (causal)
  diff    masked discrete diffusion: start the answer fully masked,
          iteratively unmask the most-confident tokens (LLaDA-style remasking)
  hybrid  a few AR steps to "think", then diffusion for the answer
  random  flips a coin each call -> "sometimes normal, sometimes diffusion"
"""
import os, json, argparse, random
import torch
import torch.nn.functional as F
from model import YKDiff
from tokenizer import YKTokenizer

HERE = os.path.dirname(os.path.abspath(__file__))
DATADIR = os.path.join(HERE, "data")
CKPTDIR = os.path.join(HERE, "checkpoints")


def load(tok_path=None, ckpt_path=None):
    tok_path = tok_path or os.path.join(DATADIR, "tokenizer.json")
    tok = YKTokenizer.load(tok_path)
    if ckpt_path is None:
        ck = sorted([f for f in os.listdir(CKPTDIR) if f.endswith(".pt")])
        if not ck:
            raise SystemExit("no checkpoint found in ./checkpoints")
        ckpt_path = os.path.join(CKPTDIR, ck[-1])
    sd = torch.load(ckpt_path, map_location="cuda")
    cfg = sd["cfg"]
    model = YKDiff(cfg).cuda().eval()
    model.load_state_dict(sd["model"])
    print(f"[infer] loaded {ckpt_path}  params={sum(p.numel() for p in model.parameters())/1e6:.1f}M")
    return model, tok


@torch.no_grad()
def _topp(logits, temp, top_p):
    logits = logits / max(temp, 1e-6)
    if top_p >= 1.0:
        return torch.multinomial(F.softmax(logits, -1), 1).item()
    s, order = torch.sort(logits, descending=True)
    p = F.softmax(s, -1)
    c = torch.cumsum(p, -1)
    keep = order[c <= top_p]
    if len(keep) == 0:
        keep = order[:1]
    return keep[torch.multinomial(F.softmax(logits[keep], -1), 1).item()].item()


@torch.no_grad()
def generate_ar(model, tok, prompt_ids, max_new=256, temp=0.9, top_p=0.95):
    ids = list(prompt_ids)
    max_len = model.max_len
    for _ in range(max_new):
        ctx = torch.tensor([ids[-max_len:]], device="cuda")
        logits = model(ctx, torch.zeros(1, dtype=torch.long, device="cuda"))[:, -1]
        nxt = _topp(logits[0], temp, top_p)
        if nxt == tok.eos_id:
            break
        ids.append(nxt)
    return ids


@torch.no_grad()
def generate_diff(model, tok, prompt_ids, gen_len=128, steps=24, temp=1.0, sample=True):
    L = len(prompt_ids) + gen_len
    seq = list(prompt_ids) + [tok.mask_id] * gen_len
    p0 = len(prompt_ids)
    for step in range(steps):
        x = torch.tensor([seq], device="cuda")
        t = torch.full((1,), (steps - step - 1) / steps, device="cuda")
        logits = model(x, torch.ones(1, dtype=torch.long, device="cuda"), t=t)
        gl = logits[0, p0:]                         # generated-region logits
        probs = F.softmax(gl / max(temp, 1e-6), -1)
        if sample:
            preds = torch.multinomial(probs, 1).squeeze(-1).tolist()
        else:
            preds = probs.argmax(-1).tolist()
        conf = probs.max(-1).values            # [gen_len]
        n_mask = int(round((steps - step - 1) / steps * gen_len))
        order = conf.argsort().tolist()        # ascending confidence
        mask_set = set(order[:n_mask])
        for j in range(gen_len):
            seq[p0 + j] = tok.mask_id if j in mask_set else preds[j]
    return seq


def build_prompt(tok, system, user):
    return [tok.bos_id] + tok.encode(f"<system>{system}</system><user>{user}</user><assistant>")


def generate(prompt_ids, mode="random", **kw):
    model, tok = _MODEL, _TOK
    if mode == "random":
        mode = "ar" if random.random() < 0.5 else "diff"
    if mode == "ar":
        ids = generate_ar(model, tok, prompt_ids, **kw)
    elif mode == "diff":
        ids = generate_diff(model, tok, prompt_ids, **kw)
    elif mode == "hybrid":
        # think a little in AR, then diffuse the answer
        think = generate_ar(model, tok, prompt_ids, max_new=64, **kw)
        # generate_diff returns the FULL sequence (think + answer)
        ids = generate_diff(model, tok, think, **kw)
    else:
        raise ValueError(mode)
    return tok.decode(ids)


# module-level cache so agent.py can call generate() directly
_MODEL, _TOK = None, None


def init(tok_path=None, ckpt_path=None):
    global _MODEL, _TOK
    _MODEL, _TOK = load(tok_path, ckpt_path)
    return _MODEL, _TOK


if __name__ == "__main__":
    ap = argparse.ArgumentParser()
    ap.add_argument("prompt", nargs="?", default="What is 17 * 23?")
    ap.add_argument("--mode", default="random")
    ap.add_argument("--max-new", type=int, default=200)
    ap.add_argument("--gen-len", type=int, default=160)
    ap.add_argument("--steps", type=int, default=24)
    ap.add_argument("--ckpt", default=None)
    a = ap.parse_args()
    m, t = init(ckpt_path=a.ckpt)
    ids = build_prompt(t, "You are clanker, a helpful assistant.", a.prompt)
    out = generate(ids, mode=a.mode, max_new=a.max_new,
                   gen_len=a.gen_len, steps=a.steps)
    print("\n=== clanker ===\n" + out)