MCPcopy Create free account
hub / github.com/MLSysU/TD-Pipe / forward

Method forward

TD_Pipe/model_executor/layers/sampler.py:30–85  ·  view source on GitHub ↗
(
        self,
        logits: torch.Tensor, # LM head
        sampling_metadata: SamplingMetadata,
    )

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 8

_apply_logits_processorsFunction · 0.85
_apply_penaltiesFunction · 0.85
_apply_top_p_top_kFunction · 0.85
_apply_min_pFunction · 0.85
_sampleFunction · 0.85
_get_logprobsFunction · 0.85
_build_sampler_outputFunction · 0.85

Tested by

no test coverage detected