Download inference.py from eshanized/M31Genesis: direct link, hf CLI and curl.
- Browser
- Download file 856 Bytes
-
https://huggingface.co/eshanized/M31Genesis/resolve/main/inference.py
- Command line
-
hf download hf://eshanized/M31Genesis/inference.py
-
curl -L -o inference.py https://huggingface.co/eshanized/M31Genesis/resolve/main/inference.py
856 Bytes
| import torch | |
| from modeling_m31 import load_model_bundle | |
| def generate(model, tokenizer, prompt, max_new_tokens=256, temperature=0.7, top_p=0.95): | |
| ids = tokenizer.encode(prompt).ids | |
| x = torch.tensor([ids], device=next(model.parameters()).device) | |
| eos = tokenizer.token_to_id('<eos>') | |
| with torch.no_grad(): | |
| for _ in range(max_new_tokens): | |
| logits = model(x[:, -1024:])[:, -1, :] / max(temperature, 1e-5) | |
| probs = torch.softmax(logits, dim=-1) | |
| vals, idx = torch.sort(probs, descending=True) | |
| c = torch.cumsum(vals, dim=-1); vals[c > top_p] = 0; vals = vals / vals.sum(dim=-1, keepdim=True) | |
| nxt = idx.gather(-1, torch.multinomial(vals, 1)); x = torch.cat([x, nxt], dim=1) | |
| if eos is not None and int(nxt.item()) == eos: break | |
| return tokenizer.decode(x[0].tolist()) | |