willychan21/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. Files Path Description data/problems.parquet One row per problem (tabular access) reference/*.py Reference solution() implementations utils/input_output_tensors.py Input/output tensor… See the full description on the dataset page: https://huggingface.co/datasets/willychan21/ParallelKernelBench_Problems.
0263
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 > 014 need_p = top_p is not None and top_p < 1.015 if not need_k and not need_p:16 return logits17 18 original_shape = logits.shape19 vocab_size = logits.shape[-1]20 logits_2d = logits.reshape(-1, vocab_size)21 if need_k:22 top_k = min(int(top_k), vocab_size)23 24 if need_k and not need_p:25 top_k_values, _ = torch.topk(logits_2d, top_k, dim=-1)26 threshold = top_k_values[..., -1:].expand_as(logits_2d)27 filtered = logits_2d.masked_fill(logits_2d < threshold, float("-inf"))28 return filtered.reshape(original_shape)29 30 sorted_logits, sorted_idx = logits_2d.sort(dim=-1, descending=False)31 if need_k:32 top_k_index = sorted_logits.shape[-1] - top_k33 threshold = sorted_logits[..., top_k_index : top_k_index + 1]34 sorted_logits = sorted_logits.masked_fill(35 sorted_logits < threshold, float("-inf")36 )37 38 sorted_probs = sorted_logits.softmax(dim=-1)39 top_p_mask = torch.cumsum(sorted_probs, dim=-1) > 1 - top_p40 top_p_mask[..., -1] = True41 sorted_logits = sorted_logits.masked_fill(~top_p_mask, float("-inf"))42 filtered = sorted_logits.scatter(dim=-1, index=sorted_idx, src=sorted_logits)43 return filtered.reshape(original_shape)44 45 46def _all_to_all_vp_to_seq(47 logits: torch.Tensor,48 group: dist.ProcessGroup,49) -> torch.Tensor:50 world_size = dist.get_world_size(group=group)51 num_tokens, local_vocab = logits.shape52 local_tokens = num_tokens // world_size53 54 send = logits.contiguous().flatten()55 recv = torch.empty_like(send)56 dist.all_to_all_single(recv, send, group=group)57 recv = recv.view(world_size, local_tokens, local_vocab)58 return recv.permute(1, 0, 2).reshape(local_tokens, world_size * local_vocab)59 60 61@torch.no_grad()62def solution(63 vocab_parallel_logits: torch.Tensor,64 target: torch.Tensor,65 tp_group: Optional[dist.ProcessGroup] = None,66 top_k: Optional[int] = None,67 top_p: float = 1.0,68 chunk_size: int = 1,69) -> torch.Tensor:70 tp_group = tp_group or dist.group.WORLD71 world_size = dist.get_world_size(group=tp_group)72 rank = dist.get_rank(group=tp_group)73 batch, seq_len, local_vocab = vocab_parallel_logits.shape74 num_tokens = batch * seq_len75 chunk_tokens = batch * max(1, int(chunk_size))76 77 if num_tokens % world_size != 0:78 raise ValueError(79 f"B*S={num_tokens} must be divisible by tensor parallel size {world_size}"80 )81 if chunk_tokens % world_size != 0:82 raise ValueError(83 f"B*chunk_size={chunk_tokens} must be divisible by tp size {world_size}"84 )85 86 logits_2d = vocab_parallel_logits.reshape(num_tokens, local_vocab)87 target_flat = target.reshape(-1)88 pieces = []89 90 for start in range(0, num_tokens, chunk_tokens):91 end = min(start + chunk_tokens, num_tokens)92 current = end - start93 local_tokens = current // world_size94 logits_chunk = logits_2d[start:end]95 target_chunk = target_flat[start:end]96 target_local = target_chunk[rank * local_tokens : (rank + 1) * local_tokens]97 98 seq_logits = _all_to_all_vp_to_seq(logits_chunk, tp_group)99 filtered = _apply_top_k_top_p(seq_logits, top_k=top_k, top_p=top_p)100 log_probs = F.log_softmax(filtered.float(), dim=-1)101 local_logprobs = torch.gather(102 log_probs, -1, target_local.unsqueeze(-1)103 ).squeeze(-1)104 105 gathered = [torch.empty_like(local_logprobs) for _ in range(world_size)]106 dist.all_gather(gathered, local_logprobs, group=tp_group)107 pieces.append(torch.cat(gathered, dim=0))108 109 return torch.cat(pieces, dim=0).reshape(batch, seq_len)110 