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