Download generation_m31.py from eshanized/M31Tesla: direct link, hf CLI and curl.
- Browser
- Download file 2.04 kB
-
https://huggingface.co/eshanized/M31Tesla/resolve/main/generation_m31.py
- Command line
-
hf download hf://eshanized/M31Tesla/generation_m31.py
-
curl -L -o generation_m31.py https://huggingface.co/eshanized/M31Tesla/resolve/main/generation_m31.py
2.04 kB
| import torch | |
| def generate(model, tokenizer, prompt, max_new_tokens=512, temperature=0.7, top_p=0.9, | |
| repetition_penalty=1.05, max_context=512, device=None): | |
| if device is None: | |
| device = 'cuda' if torch.cuda.is_available() else 'cpu' | |
| model.eval() | |
| ids = tokenizer.encode(prompt).ids | |
| if not ids: | |
| raise ValueError('Prompt produced zero tokenizer tokens.') | |
| x = torch.tensor([ids[-max_context:]], dtype=torch.long, device=device) | |
| eos_id = tokenizer.token_to_id('<eos>') | |
| for _ in range(int(max_new_tokens)): | |
| x_in = x[:, -max_context:] | |
| with torch.autocast(device_type='cuda', dtype=torch.float16, enabled=(device == 'cuda')): | |
| logits, _ = model(x_in) | |
| logits = logits[:, -1, :].float() | |
| if repetition_penalty and repetition_penalty > 1.0: | |
| for token_id in torch.unique(x).tolist(): | |
| score = logits[0, token_id] | |
| logits[0, token_id] = score / repetition_penalty if score > 0 else score * repetition_penalty | |
| if temperature is None or temperature <= 0: | |
| nxt = torch.argmax(logits, dim=-1, keepdim=True) | |
| else: | |
| logits = logits / max(float(temperature), 1e-5) | |
| probs = torch.softmax(logits, dim=-1) | |
| if top_p is not None and 0 < top_p < 1.0: | |
| sorted_probs, sorted_idx = torch.sort(probs, descending=True, dim=-1) | |
| cumulative = torch.cumsum(sorted_probs, dim=-1) | |
| remove = cumulative > float(top_p) | |
| remove[..., 0] = False | |
| sorted_probs = sorted_probs.masked_fill(remove, 0.0) | |
| probs = torch.zeros_like(probs).scatter(-1, sorted_idx, sorted_probs) | |
| probs = probs / probs.sum(dim=-1, keepdim=True) | |
| nxt = torch.multinomial(probs, 1) | |
| x = torch.cat([x, nxt], dim=1) | |
| if eos_id is not None and int(nxt.item()) == int(eos_id): | |
| break | |
| return tokenizer.decode(x[0].tolist(), skip_special_tokens=False) | |