import torch @torch.no_grad() 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('') 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)