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