"""Standalone generation for this package. No external project files needed. python generate.py -p "The capital of France is" python generate.py -p "..." --max-new 120 --temp 0.7 --no-cache """ import argparse, json, os, sys os.environ.setdefault("TRANSFORMERS_NO_TF", "1") os.environ.setdefault("USE_TF", "0") import torch import torch.nn.functional as F HERE = os.path.dirname(os.path.abspath(__file__)) sys.path.insert(0, HERE) MANIFEST = json.load(open(os.path.join(HERE, "package.json"))) def load(device="cpu", ckpt_override=None): from model_v2 import SpikeWhaleLM from spike_tokenizer import SpikeTokenizer from config import SpikeWhaleConfig tok = SpikeTokenizer(os.path.join(HERE, "tokenizer.json")) ckpt = ckpt_override or MANIFEST["ckpt"] path = ckpt if os.path.isabs(ckpt) else os.path.join(HERE, ckpt) drop = ( "architectures", "transformers_version", "dtype", "torch_dtype", "id2label", "label2id", "_name_or_path", ) if os.path.isdir(path): model = SpikeWhaleLM.from_pretrained(path) else: blob = torch.load(path, map_location="cpu", weights_only=False) raw = dict(blob["config"]) for k in drop: raw.pop(k, None) if "vocab_size" not in raw: raw["vocab_size"] = MANIFEST.get("vocab_size", 16512) cfg = SpikeWhaleConfig(**raw) model = SpikeWhaleLM(cfg) sd = blob["model_state"] msd = model.state_dict() model.load_state_dict({k: v for k, v in sd.items() if k in msd and msd[k].shape == v.shape}, strict=False) if getattr(cfg, "tie_word_embeddings", True): model.tie_weights() return model.to(device).float().eval(), tok def build_prompt(tok, text, chat): if chat: try: from chat_format import format_chat return format_chat([{"role": "user", "content": text}], add_generation_prompt=True) except Exception: pass return text @torch.no_grad() def generate(model, tok, prompt, max_new=80, temp=0.0, top_k=20, rp=1.15, device="cpu", use_cache=True, seed=0): try: from model_v2 import reset_memory_cache reset_memory_cache(model) except Exception: pass g = torch.Generator(device=device); g.manual_seed(seed) ids = tok.encode(prompt) ids = ids.tolist() if hasattr(ids, "tolist") else list(ids) # to match training: chat_format.tokenize_chatml builds every training # sequence with add_bos=True, so the model has only ever seen sequences # starting with it. Omitting it shifts every token one position left and puts # a token at position 0 that the model never saw there. _bos = getattr(tok, "bos_token_id", None) if _bos is not None and (not ids or ids[0] != _bos): ids = [_bos] + ids stop = set() for name in ("<|im_end|>", ""): try: i = tok.convert_tokens_to_ids(name) if i is not None and i >= 0: stop.add(int(i)) except Exception: pass e = getattr(tok, "eos_token_id", None) if e is not None: stop.add(int(e)) cfg = model.config use_eng = bool(getattr(cfg, "use_engram", False)) and MANIFEST.get("engram_kwarg") nctx = max(1, int(getattr(cfg, "engram_max_ngram", 3)) - 1) # Positions run 0..ctx-1. Going past that indexes the RoPE cache out of bounds # and raises a CUDA device-side assert that kills the process -- it does NOT # degrade gracefully. Keep the MOST RECENT tokens. # # Truncation is the right answer here, not a fallback: measured at N=32768 on # both Mark2 trees, truncating to the window gave ppl 8.4, while every # position-aliasing scheme tried (JetLong G=2 21.6/23.8, bifocal 20.4/22.5, # clamp 19.0/20.1) was 2.3-2.8x worse. The rope is already extended via # rope_theta, and aliasing on top of that fights it. max_ctx = int(getattr(cfg, "max_position_embeddings", 4096)) keep = max_ctx - int(max_new) if keep > 0 and len(ids) > keep: ids = ids[-keep:] if use_cache: out = model(torch.tensor([ids], device=device), use_cache=True) past, lg = out.past_key_values, out.logits[:, -1, :] else: past = None lg = model(torch.tensor([ids], device=device)).logits[:, -1, :] gen, new = list(ids), [] for _ in range(max_new): l = lg[0].float().clone() if new and rp != 1.0: idx = torch.tensor(sorted(set(new)), device=device) v = l[idx]; l[idx] = torch.where(v > 0, v / rp, v * rp) if temp <= 0: nxt = int(l.argmax()) else: v, i2 = l.topk(min(top_k, l.numel())) nxt = int(i2[torch.multinomial(F.softmax(v / temp, -1), 1, generator=g)]) if nxt in stop: break gen.append(nxt); new.append(nxt) if len(gen) >= max_ctx: # never index past the RoPE cache break if use_cache: kw = {} if use_eng and len(gen) > 1: kw["engram_context_ids"] = torch.tensor([gen[-(nctx + 1):-1]], device=device) out = model(torch.tensor([[nxt]], device=device), past_key_values=past, use_cache=True, **kw) past, lg = out.past_key_values, out.logits[:, -1, :] else: lg = model(torch.tensor([gen], device=device)).logits[:, -1, :] # Keep CONTENT markers such as /: they are registered special # tokens, so skip_special_tokens=True deleted them and the reasoning block # silently vanished from the output (measured on JEPA6: is the # TOP-RANKED token at step 0 on reasoning prompts). Drop only framing ids. _drop = set(stop) for _a in ("bos_token_id", "eos_token_id"): _v = getattr(tok, _a, None) if _v is not None: _drop.add(int(_v)) for _n in ("<|im_start|>", "<|im_end|>"): try: _v = tok.convert_tokens_to_ids(_n) if _v is not None and _v >= 0: _drop.add(int(_v)) except Exception: pass return tok.decode([t for t in new if t not in _drop], skip_special_tokens=False) def main(): ap = argparse.ArgumentParser() ap.add_argument("-p", "--prompt", default=MANIFEST.get("example_prompt", "The capital of France is")) _dd = MANIFEST.get("decoding_defaults", {}) ap.add_argument("--max-new", type=int, default=80) ap.add_argument("--temp", type=float, default=_dd.get("temp", 0.7)) ap.add_argument("--top-k", type=int, default=_dd.get("top_k", 40)) ap.add_argument("--rp", type=float, default=_dd.get("rep_pen", 1.3)) ap.add_argument("--device", default="cpu") ap.add_argument("--threads", type=int, default=4) ap.add_argument("--no-cache", action="store_true") ap.add_argument("--chat", action="store_true", help="wrap the prompt in this model's chat template") ap.add_argument("--ckpt", default=None, help="checkpoint override, e.g. checkpoints/base_62k.pt | " "checkpoints/sft_7100.pt | checkpoints/dpo_3200.pt " "(default: package.json). Use --chat with sft/dpo.") a = ap.parse_args() torch.set_num_threads(a.threads) model, tok = load(a.device, a.ckpt) n = sum(p.numel() for p in model.parameters()) / 1e6 print(f"{MANIFEST['name']} {n:.1f}M params device={a.device} " f"cache={not a.no_cache}") txt = generate(model, tok, build_prompt(tok, a.prompt, a.chat), a.max_new, a.temp, a.top_k, a.rp, a.device, not a.no_cache) print("-" * 70) print(a.prompt) print("-" * 70) print(txt) if __name__ == "__main__": main()