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

Function fused_mm_sample_triton

src/fused_mm_sampling/core.py:220–332  ·  view source on GitHub ↗
(
    weights: torch.Tensor,  # [V_local, D] (may be a TP shard)
    hidden_states: torch.Tensor,  # [n_hidden_states, D]
    num_samples: int,
    temperature: torch.Tensor,  # scalar (0-d)
    seed: int,
    greedy_sampling: bool = False,
    tp: "TPInfo" = TP1,
    return_logits: bool = False,
)

Source from the content-addressed store, hash-verified

218
219MIN_BLOCK_SIZE_V = 128
220
221
222# @torch.compile(fullgraph=True)
223@nvtx.annotate()
224def fused_mm_sample_triton(
225 weights: torch.Tensor, # [V_local, D] (may be a TP shard)
226 hidden_states: torch.Tensor, # [n_hidden_states, D]
227 num_samples: int,
228 temperature: torch.Tensor, # scalar (0-d)
229 seed: int,
230 greedy_sampling: bool = False,
231 tp: "TPInfo" = TP1,
232 return_logits: bool = False,
233 p2p_no_overlap: bool = False,
234):
235 assert torch.cuda.is_available(), "fused_mm_sample_triton requires CUDA"
236 V, D = weights.shape # noqa: N806
237 H, D2 = hidden_states.shape # noqa: N806
238 if D2 != D:
239 raise ValueError(
240 f"hidden_states second dimension ({D2}) must match weights second dimension ({D})"
241 )
242
243 # The kernel uses TMA descriptors which need a runtime allocator. Some
244 # autotuner configs (notably the ones picked on B200/sm_100) request global
245 # scratch from this allocator; without it Triton raises a RuntimeError at launch.
246 set_torch_allocator_for_tma_descriptors_cached()
247
248 NUM_SMS = num_sms_cached(weights.device.index) # noqa: N806
249
250 max_grid_size_v = triton.cdiv(V, MIN_BLOCK_SIZE_V)
251 fan_out_tp = tp.size > 1 and not p2p_no_overlap
252 if tp.size > 1:
253 maxs, maxs_idx, symm_mem_hdl, storage_offset_maxs_idx = allocate_symm_mem_outputs(
254 num_samples=num_samples,
255 max_grid_size_v=max_grid_size_v,
256 H=H,
257 )
258 kernel_maxs = maxs[tp.rank]
259 kernel_maxs_idx = maxs_idx[tp.rank]
260 symm_mem_buffer_ptrs = symm_mem_hdl.buffer_ptrs_dev
261 else:
262 maxs = torch.empty(
263 (num_samples, max_grid_size_v, H),
264 dtype=torch.float32,
265 device=weights.device,
266 )
267 maxs_idx = torch.empty_like(maxs, dtype=torch.long)
268 kernel_maxs = maxs
269 kernel_maxs_idx = maxs_idx
270 storage_offset_maxs_idx = 0
271 symm_mem_buffer_ptrs = maxs
272
273 # logits_out is only read when RETURN_LOGITS=True. For the common path
274 # (return_logits=False), allocating a (V, H) fp32 buffer per call is wasted
275 # HBM (155 MB per decode step at Qwen3-1.7B / H=256) for a buffer the
276 # kernel never touches. Pass a 1-element dummy in that case so the kernel
277 # still has a valid pointer to receive.

Callers 5

test_greedy_samplingFunction · 0.90
mainFunction · 0.90
basic_usage.pyFile · 0.90
get_samplerFunction · 0.85

Calls 6

num_sms_cachedFunction · 0.85
tp_post_kernel_reduceFunction · 0.85
_local_reduceFunction · 0.85

Tested by 2

test_greedy_samplingFunction · 0.72