| """ |
| 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:] |
| 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 |
| n_mask = int(round((steps - step - 1) / steps * gen_len)) |
| order = conf.argsort().tolist() |
| 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 = generate_ar(model, tok, prompt_ids, max_new=64, **kw) |
| |
| ids = generate_diff(model, tok, think, **kw) |
| else: |
| raise ValueError(mode) |
| return tok.decode(ids) |
|
|
|
|
| |
| _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) |
|
|