CoolFace
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 4mo agoView on Hugging Face
0likes263downloads
83_vocab_parallel_log_prob_topk_chunked.py110 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 > 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