Team Ai
Datasetpublic

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.

sourceHugging Faceapache-2.0updated 4mo agoView on Hugging Face
0likes177downloads
77_magi1_cso_async_attention.py186 linesDownload Raw Back to reference
1from typing import List, Optional2 3import torch4import torch.distributed as dist5import torch.nn.functional as F6 7 8def _a2a_rows(9    tensor: torch.Tensor,10    split_sizes: List[int],11    group: dist.ProcessGroup,12) -> torch.Tensor:13    out = torch.empty_like(tensor)14    dist.all_to_all_single(15        out,16        tensor.contiguous(),17        output_split_sizes=split_sizes,18        input_split_sizes=split_sizes,19        group=group,20    )21    return out22 23 24def _a2a_async(25    tensor: torch.Tensor,26    split_sizes: List[int],27    group: dist.ProcessGroup,28) -> tuple[torch.Tensor, dist.Work]:29    out = torch.empty_like(tensor)30    handle = dist.all_to_all_single(31        out,32        tensor.contiguous(),33        output_split_sizes=split_sizes,34        input_split_sizes=split_sizes,35        group=group,36        async_op=True,37    )38    return out, handle39 40 41def _redistribute_kv(42    key_value: torch.Tensor,43    world_size: int,44    group: dist.ProcessGroup,45) -> torch.Tensor:46    tokens, heads, width = key_value.shape47    if heads < world_size and world_size % heads == 0:48        key_value = key_value.repeat_interleave(world_size // heads, dim=1)49        heads = key_value.shape[1]50    if heads % world_size != 0:51        raise ValueError("KV heads must divide evenly across context ranks")52 53    local_heads = heads // world_size54    packed = key_value.reshape(tokens, world_size, local_heads, width)55    packed = packed.permute(1, 0, 2, 3).reshape(world_size * tokens, local_heads, width)56    return _a2a_rows(packed.contiguous(), [tokens] * world_size, group)57 58 59def _kv_by_range(60    kv: torch.Tensor,61    world_size: int,62    ranges: int,63    spb: int,64    clip_token_nums: int,65) -> torch.Tensor:66    _, heads, width = kv.shape67    kv = kv.reshape(world_size, ranges, spb, heads, width)68    kv = kv.permute(1, 0, 2, 3, 4).contiguous()69    kv = kv.reshape(ranges, world_size * spb, heads, width)70    return kv[:, :clip_token_nums].reshape(ranges * clip_token_nums, heads, width)71 72 73def _split_query(query: torch.Tensor, world_size: int, ranges: int) -> List[torch.Tensor]:74    tokens, heads, head_dim = query.shape75    if tokens % ranges != 0:76        raise ValueError("query token count must divide cp_shuffle_num")77    if heads % world_size != 0:78        raise ValueError("query heads must divide evenly across context ranks")79 80    spb = tokens // ranges81    local_heads = heads // world_size82    query = query.reshape(ranges, spb, world_size, local_heads, head_dim)83    query = query.permute(0, 2, 1, 3, 4).contiguous()84    query = query.reshape(ranges, world_size * spb, local_heads, head_dim)85    return [query[idx] for idx in range(ranges)]86 87 88def _restore_output(89    chunks: List[torch.Tensor],90    world_size: int,91    spb: int,92) -> torch.Tensor:93    out = torch.stack(chunks, dim=0)94    ranges, _, heads, head_dim = out.shape95    out = out.reshape(ranges, world_size, spb, heads, head_dim)96    out = out.permute(0, 2, 1, 3, 4).contiguous()97    return out.reshape(ranges * spb, world_size * heads, head_dim)98 99 100def _sdpa(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:101    q = q.unsqueeze(0).transpose(1, 2)102    k = k.unsqueeze(0).transpose(1, 2)103    v = v.unsqueeze(0).transpose(1, 2)104    if k.shape[1] < q.shape[1]:105        repeat = q.shape[1] // k.shape[1]106        k = k.repeat_interleave(repeat, dim=1)107        v = v.repeat_interleave(repeat, dim=1)108    return F.scaled_dot_product_attention(q, k, v).squeeze(0).transpose(0, 1).contiguous()109 110 111@torch.no_grad()112def solution(113    query: torch.Tensor,114    key_value: torch.Tensor,115    k_ranges: torch.Tensor,116    cp_shuffle_num: int,117    clip_token_nums: Optional[int] = None,118    group: Optional[dist.ProcessGroup] = None,119) -> torch.Tensor:120    group = group or dist.group.WORLD121    world_size = dist.get_world_size(group=group)122    ranges = cp_shuffle_num123    tokens, _, head_dim = query.shape124    if tokens % ranges != 0:125        raise ValueError("query token count must divide cp_shuffle_num")126    spb = tokens // ranges127    clip_token_nums = int(clip_token_nums or world_size * spb)128 129    kv = _redistribute_kv(key_value, world_size, group)130    kv = _kv_by_range(kv, world_size, ranges, spb, clip_token_nums)131    key = kv[..., :head_dim]132    value = kv[..., head_dim:]133 134    q_chunks = _split_query(query, world_size, ranges)135    split_sizes = [spb] * world_size136    if ranges == 1:137        q_local, handle_q = _a2a_async(q_chunks[0], split_sizes, group)138        handle_q.wait()139        start = int(k_ranges[0, 0])140        end = int(k_ranges[0, 1])141        out = _sdpa(q_local, key[start:end], value[start:end])142        out, handle_o = _a2a_async(out, split_sizes, group)143        handle_o.wait()144        return _restore_output([out], world_size, spb)145 146    outputs: List[torch.Tensor] = []147    q_chunks[0], handle_q = _a2a_async(q_chunks[0], split_sizes, group)148    loop_var: Optional[torch.Tensor] = None149    loop_handle: Optional[dist.Work] = None150    prev_out: Optional[torch.Tensor] = None151    for idx in range(ranges):152        if idx == 0:153            handle_q.wait()154            q_local = q_chunks[0]155            loop_var, loop_handle = _a2a_async(q_chunks[1], split_sizes, group)156        else:157            assert loop_var is not None and loop_handle is not None158            loop_handle.wait()159            if loop_var.numel() == q_chunks[0].numel():160                q_local = loop_var161            else:162                q_local, ready_out = torch.chunk(loop_var, 2, dim=-1)163                outputs.append(ready_out)164 165            assert prev_out is not None166            send = (167                torch.cat([q_chunks[idx + 1], prev_out], dim=-1)168                if idx < ranges - 1169                else prev_out170            )171            loop_var, loop_handle = _a2a_async(send, split_sizes, group)172 173        start = int(k_ranges[idx, 0])174        end = int(k_ranges[idx, 1])175        prev_out = _sdpa(q_local, key[start:end], value[start:end])176 177        if idx == ranges - 1:178            assert loop_var is not None and loop_handle is not None179            loop_handle.wait()180            outputs.append(loop_var)181            last_out, handle_out = _a2a_async(prev_out, split_sizes, group)182            handle_out.wait()183            outputs.append(last_out)184 185    return _restore_output(outputs, world_size, spb)186