Team Ai
Datasetpublic

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.

sourceHugging Faceapache-2.0updated 5mo agoView on Hugging Face
0likes123downloads
82_vocab_parallel_log_prob_topk.py112 linesDownload Raw Back to reference
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)