Byrne-15M-Looped / generate.py
Quazim0t0's picture
Byrne-15M-Looped: inference package, scores, safetensors for leaderboard
be40882 verified
Raw
History Blame Contribute Delete
8.21 kB
"""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> 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|>", "<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)
# 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 <think>/</think>: they are registered special
# tokens, so skip_special_tokens=True deleted them and the reasoning block
# silently vanished from the output (measured on JEPA6: <think> 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()