|
|
| import torch |
| import torch.nn.functional as F |
| from transformers import GPT2Tokenizer |
|
|
| |
| tokenizer = GPT2Tokenizer.from_pretrained("gpt2") |
| tokenizer.pad_token = tokenizer.eos_token |
|
|
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
|
|
| def generate(model, prompt, max_new_tokens=200, |
| temperature=0.8, top_k=50, |
| repetition_penalty=1.3): |
| model.eval() |
| input_ids = tokenizer.encode(prompt, return_tensors="pt").to(DEVICE) |
| generated = input_ids.clone() |
|
|
| with torch.no_grad(): |
| for _ in range(max_new_tokens): |
| x = generated[:, -512:] |
| logits = model(x)[:, -1, :].float() |
|
|
| for token_id in set(generated[0].tolist()): |
| if logits[0, token_id] > 0: |
| logits[0, token_id] /= repetition_penalty |
| else: |
| logits[0, token_id] *= repetition_penalty |
|
|
| logits = logits / temperature |
| k = min(top_k, logits.size(-1)) |
| topk_vals, _ = torch.topk(logits, k) |
| logits = logits.masked_fill(logits < topk_vals[:, -1:], -1e9) |
| probs = torch.softmax(logits, dim=-1).clamp(min=0) |
| probs = probs / probs.sum() |
| next_token = torch.multinomial(probs, num_samples=1) |
| generated = torch.cat([generated, next_token], dim=1) |
| if next_token.item() == tokenizer.eos_token_id: |
| break |
|
|
| return tokenizer.decode(generated[0], skip_special_tokens=True) |
|
|