| """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)
|
|
|
|
|
|
|
|
|
| _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|>", "<eos>"):
|
| 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)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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:
|
| 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, :]
|
|
|
|
|
|
|
|
|
| _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()
|
|
|