""" 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}{user}") 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)