(
self,
logits: torch.Tensor, # LM head
sampling_metadata: SamplingMetadata,
)
| 28 | super().__init__() |
| 29 | |
| 30 | def forward( |
| 31 | self, |
| 32 | logits: torch.Tensor, # LM head |
| 33 | sampling_metadata: SamplingMetadata, |
| 34 | ) -> Optional[SamplerOutput]: |
| 35 | # Only perform sampling in the driver worker. |
| 36 | # Note: `_get_logits` is still distributed across TP workers because |
| 37 | # the `embedding` weight is distributed across TP workers. |
| 38 | # TODO(zhuohan): Change the get_logits part to a separate stage. |
| 39 | if not sampling_metadata.perform_sampling: |
| 40 | return None |
| 41 | |
| 42 | assert logits is not None |
| 43 | _, vocab_size = logits.shape |
| 44 | |
| 45 | # Apply logits processors (if any). |
| 46 | logits = _apply_logits_processors(logits, sampling_metadata) |
| 47 | |
| 48 | # Prepare sampling tensors with pinned memory to avoid blocking. |
| 49 | (sampling_tensors, do_penalties, do_top_p_top_k, |
| 50 | do_min_p) = SamplingTensors.from_sampling_metadata( |
| 51 | sampling_metadata, vocab_size, logits.device, logits.dtype) |
| 52 | |
| 53 | # Apply presence and frequency penalties. |
| 54 | if do_penalties: |
| 55 | logits = _apply_penalties(logits, sampling_tensors.prompt_tokens, |
| 56 | sampling_tensors.output_tokens, |
| 57 | sampling_tensors.presence_penalties, |
| 58 | sampling_tensors.frequency_penalties, |
| 59 | sampling_tensors.repetition_penalties) |
| 60 | |
| 61 | # Apply temperature scaling. |
| 62 | # Use in-place division to avoid creating a new tensor. |
| 63 | logits.div_(sampling_tensors.temperatures.unsqueeze_(dim=1)) |
| 64 | |
| 65 | if do_top_p_top_k: |
| 66 | logits = _apply_top_p_top_k(logits, sampling_tensors.top_ps, |
| 67 | sampling_tensors.top_ks) |
| 68 | |
| 69 | if do_min_p: |
| 70 | logits = _apply_min_p(logits, sampling_tensors.min_ps) |
| 71 | |
| 72 | # We use float32 for probabilities and log probabilities. |
| 73 | # Compute the probabilities. |
| 74 | probs = torch.softmax(logits, dim=-1, dtype=torch.float) |
| 75 | # Compute the log probabilities. |
| 76 | # Use log_softmax to ensure numerical stability. |
| 77 | logprobs = torch.log_softmax(logits, dim=-1, dtype=torch.float) |
| 78 | |
| 79 | # Sample the next tokens. |
| 80 | sample_results = _sample(probs, logprobs, sampling_metadata) |
| 81 | # Get the logprobs query results. |
| 82 | prompt_logprobs, sample_logprobs = _get_logprobs( |
| 83 | logprobs, sampling_metadata, sample_results) |
| 84 | return _build_sampler_output(sample_results, sampling_metadata, |
| 85 | prompt_logprobs, sample_logprobs) |
| 86 | |
| 87 |
nothing calls this directly
no test coverage detected