| 218 | |
| 219 | MIN_BLOCK_SIZE_V = 128 |
| 220 | |
| 221 | |
| 222 | # @torch.compile(fullgraph=True) |
| 223 | @nvtx.annotate() |
| 224 | def 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. |