| 68 | |
| 69 | |
| 70 | def get_logits(tokenizer, model, inputs: List[str]): |
| 71 | input_ids = tokenizer(inputs, padding=False)["input_ids"] |
| 72 | input_ids = torch.tensor(input_ids, device=model.device) |
| 73 | |
| 74 | if input_ids.shape[1] > args.max_seq_len: |
| 75 | input_ids = input_ids[:, input_ids.shape[1] - args.max_seq_len + 1 :] |
| 76 | tokens = {"input_ids": input_ids} |
| 77 | |
| 78 | outputs = model(input_ids)["logits"] |
| 79 | logits = outputs[:, -1, :] |
| 80 | log_probs = torch.nn.functional.softmax(logits, dim=-1) |
| 81 | return log_probs, {"tokens": tokens} |
| 82 | |
| 83 | |
| 84 | @torch.no_grad() |