M31Tesla / generation_m31.py
eshanized's picture
add M31 v5 standalone inference runtime
e92cf18 verified
Raw History Blame Contribute Delete
2.04 kB
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('<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)