togethercomputer/ParallelKernelBench_Problems
ParallelKernelBench (benchmark) Reference problems for ParallelKernelBench: a benchmark for LLM-generated multi-GPU CUDA kernels. This dataset contains 87 reference implementations in reference/ and the input tensor specification in utils/input_output_tensors.py. Inputs are deterministic — reproduce them with create_input_tensor(rank, world_size, problem_id, base_shape, dtype, trial) from that file; you do not need stored .pt files. Files Path Description… See the full description on the dataset page: https://huggingface.co/datasets/togethercomputer/ParallelKernelBench_Problems.
0177
1from typing import Optional2 3import torch4import torch.distributed as dist5import torch.nn.functional as F6 7 8def _apply_top_k_top_p(9 logits: torch.Tensor,10 top_k: Optional[int],11 top_p: float,12) -> torch.Tensor:13 need_k = top_k is not None and top_k > 0 14 need_p = top_p is not None and top_p < 1.0 15 16 if not need_k and not need_p: 17 return logits 18 19 original_shape = logits.shape 20 vocab_size = logits.shape[-1] 21 logits_2d = logits.reshape(-1, vocab_size)22 if need_k:23 top_k = min(int(top_k), vocab_size)24 25 if need_k and not need_p: 26 top_k_values, _ = torch.topk(logits_2d, top_k, dim=-1) 27 threshold = top_k_values[..., -1:].expand_as(logits_2d) 28 keep_mask = logits_2d >= threshold 29 filtered = torch.where( 30 keep_mask, 31 logits_2d, 32 torch.full_like(logits_2d, float("-inf")), 33 ) 34 return filtered.reshape(original_shape)35 36 logits_sort, logits_idx = logits_2d.sort(dim=-1, descending=False) 37 38 top_k_mask = None 39 if need_k: 40 top_k_index = logits_sort.size(-1) - top_k 41 threshold = logits_sort.gather( 42 -1, 43 torch.full( 44 logits_sort.shape[:-1], 45 top_k_index, 46 device=logits_2d.device, 47 dtype=torch.long, 48 ).unsqueeze(-1), 49 ) 50 top_k_mask = logits_sort >= threshold 51 logits_sort = logits_sort.masked_fill(~top_k_mask, float("-inf")) 52 53 probs_sort = logits_sort.softmax(dim=-1) 54 probs_sum = torch.cumsum(probs_sort, dim=-1) 55 top_p_mask = probs_sum > 1 - top_p 56 top_p_mask[..., -1] = True # always keep at least one token 57 logits_sort = logits_sort.masked_fill(~top_p_mask, float("-inf")) 58 59 filtered = logits_sort.scatter(dim=-1, index=logits_idx, src=logits_sort) 60 return filtered.reshape(original_shape)61 62 63def _all_to_all_vp_to_seq(64 vocab_parallel_logits: torch.Tensor,65 tp_group: dist.ProcessGroup, 66) -> torch.Tensor:67 world_size = dist.get_world_size(tp_group) 68 num_tokens, local_vocab = vocab_parallel_logits.shape69 local_tokens = num_tokens // world_size70 71 input_flat = vocab_parallel_logits.contiguous().flatten()72 output_flat = torch.empty_like(input_flat) 73 dist.all_to_all_single(output_flat, input_flat, group=tp_group) 74 75 output = output_flat.view(world_size, local_tokens, local_vocab)76 return output.permute(1, 0, 2).reshape(local_tokens, world_size * local_vocab)77 78 79@torch.no_grad() 80def solution( 81 vocab_parallel_logits: torch.Tensor,82 target: torch.Tensor,83 tp_group: Optional[dist.ProcessGroup] = None,84 top_k: Optional[int] = None, 85 top_p: float = 1.0, 86) -> torch.Tensor:87 tp_group = tp_group or dist.group.WORLD88 world_size = dist.get_world_size(tp_group) 89 rank = dist.get_rank(tp_group) 90 batch, seq_len, local_vocab = vocab_parallel_logits.shape91 num_tokens = batch * seq_len92 93 if num_tokens % world_size != 0:94 raise ValueError( 95 f"B*S={num_tokens} must be divisible by tensor parallel size {world_size}"96 ) 97 local_tokens = num_tokens // world_size98 99 logits_2d = vocab_parallel_logits.reshape(num_tokens, local_vocab)100 target_flat = target.reshape(-1)101 target_local = target_flat[rank * local_tokens : (rank + 1) * local_tokens]102 103 seq_parallel_logits = _all_to_all_vp_to_seq(logits_2d, tp_group)104 logits = _apply_top_k_top_p(seq_parallel_logits, top_k=top_k, top_p=top_p)105 log_probs = F.log_softmax(logits.to(dtype=torch.float32), dim=-1) 106 107 token_logprobs = torch.gather(log_probs, -1, target_local.unsqueeze(-1))108 token_logprobs = token_logprobs.squeeze(-1)109 110 gathered = [torch.empty_like(token_logprobs) for _ in range(world_size)]111 dist.all_gather(gathered, token_logprobs, group=tp_group) 112 return torch.cat(gathered, dim=0).reshape(batch, seq_len)