MCPcopy Create free account
hub / github.com/FlashSampling/FlashSampling / greedy_sample

Function greedy_sample

src/fused_mm_sampling/core.py:139–151  ·  view source on GitHub ↗

Baseline: matmul for logits followed by argmax. Returns [n_hidden_states, 1].

(
    weights: torch.Tensor,  # [V, D]
    hidden_states: torch.Tensor,  # [n_hidden_states, D]
    num_samples: int,  # ignored (always returns 1 sample per row)
    temperature: torch.Tensor,  # ignored (greedy)
    tp: "TPInfo" = TP1,
    **_kwargs,
)

Source from the content-addressed store, hash-verified

137 )
138 logits = apply_top_k_top_p_triton(logits, k, p)
139 return logits.softmax(dim=-1)
140
141
142@nvtx.annotate()
143def greedy_sample(
144 weights: torch.Tensor, # [V, D]
145 hidden_states: torch.Tensor, # [n_hidden_states, D]
146 num_samples: int, # ignored (always returns 1 sample per row)
147 temperature: torch.Tensor, # ignored (greedy)
148 tp: "TPInfo" = TP1,
149 **_kwargs,
150) -> torch.Tensor:
151 """Baseline: matmul for logits followed by argmax. Returns [n_hidden_states, 1]."""
152 logits = hidden_states @ weights.T # [n_hidden_states, V]
153 if tp.size > 1:
154 logits = _allgather_logits(logits)

Callers

nothing calls this directly

Calls 1

_allgather_logitsFunction · 0.85

Tested by

no test coverage detected