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
0likes266downloads
55_ring_attention_tp.py133 linesDownload Raw Back to reference
1from typing import Optional, Tuple2 3import torch4import torch.distributed as dist5import torch.nn.functional as F6 7 8@torch.jit.script9def _update_out_and_lse(10    out: torch.Tensor, lse: torch.Tensor,11    block_out: torch.Tensor, block_lse: torch.Tensor,12) -> Tuple[torch.Tensor, torch.Tensor]:13    block_out = block_out.to(torch.float32)14    block_lse = block_lse.transpose(-2, -1).unsqueeze(dim=-1)15    out = out - F.sigmoid(block_lse - lse) * (out - block_out)16    lse = lse - F.logsigmoid(lse - block_lse)17    return out, lse18 19 20def _merge_out_lse(21    out: Optional[torch.Tensor], lse: Optional[torch.Tensor],22    block_out: torch.Tensor, block_lse: torch.Tensor,23) -> Tuple[torch.Tensor, torch.Tensor]:24    if out is None:25        return block_out.to(torch.float32), block_lse.transpose(-2, -1).unsqueeze(-1)26    return _update_out_and_lse(out, lse, block_out, block_lse)27 28 29class RingComm:30    def __init__(self, group: dist.ProcessGroup):31        self._group = group32        self._ops = []33        self._reqs = None34        self.rank = dist.get_rank(group)35        self.world_size = dist.get_world_size(group)36        self.send_rank = dist.get_global_rank(group, (self.rank + 1) % self.world_size)37        self.recv_rank = dist.get_global_rank(group, (self.rank - 1) % self.world_size)38 39    def send_recv(self, to_send: torch.Tensor, recv_buf: Optional[torch.Tensor] = None) -> torch.Tensor:40        buf = recv_buf if recv_buf is not None else torch.empty_like(to_send)41        self._ops.append(dist.P2POp(dist.isend, to_send, self.send_rank, group=self._group))42        self._ops.append(dist.P2POp(dist.irecv, buf, self.recv_rank, group=self._group))43        return buf44 45    def commit(self):46        self._reqs = dist.batch_isend_irecv(self._ops)47 48    def wait(self):49        for r in self._reqs:50            r.wait()51        self._reqs = None52        self._ops = []53 54    def send_recv_kv(self, k: torch.Tensor, v: torch.Tensor):55        next_k = self.send_recv(k)56        next_v = self.send_recv(v)57        self.commit()58        return next_k, next_v59 60 61def _local_attn(62    q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,63    scale: float, causal: bool,64) -> Tuple[torch.Tensor, torch.Tensor]:65    qh = q.transpose(1, 2).float()66    kh = k.transpose(1, 2).float()67    vh = v.transpose(1, 2).float()68    scores = torch.matmul(qh, kh.transpose(-2, -1)) * scale69    if causal:70        mask = torch.triu(torch.ones(q.size(1), k.size(1), device=q.device, dtype=torch.bool), 1)71        scores.masked_fill_(mask.unsqueeze(0).unsqueeze(0), float("-inf"))72    block_lse = torch.logsumexp(scores, dim=-1)73    block_out = torch.matmul(torch.softmax(scores, dim=-1), vh).transpose(1, 2).contiguous()74    return block_out, block_lse75 76 77def _ring_attn_forward(78    group: dist.ProcessGroup,79    q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,80    scale: float, causal: bool,81) -> torch.Tensor:82    world_size = dist.get_world_size(group)83    if world_size == 1:84        out, lse = _merge_out_lse(None, None, *_local_attn(q, k, v, scale, causal))85        return out.to(q.dtype)86 87    comm = RingComm(group)88    out, lse = None, None89 90    for step in range(world_size):91        if step + 1 != world_size:92            next_k, next_v = comm.send_recv_kv(k, v)93        if (not causal) or step <= comm.rank:94            block_out, block_lse = _local_attn(q, k, v, scale, causal=(causal and step == 0))95            out, lse = _merge_out_lse(out, lse, block_out, block_lse)96        if step + 1 != world_size:97            comm.wait()98            k, v = next_k, next_v99 100    return out.to(q.dtype)101 102 103def solution(104    hidden_states: torch.Tensor,105    w_qkv: torch.Tensor,106    w_o: torch.Tensor,107    num_heads: int,108    softmax_scale: Optional[float] = None,109    causal: bool = False,110    tp_group: Optional[dist.ProcessGroup] = None,111    cp_group: Optional[dist.ProcessGroup] = None,112) -> torch.Tensor:113    tp_group = tp_group or dist.group.WORLD114    cp_group = cp_group or dist.group.WORLD115 116    tp_size = dist.get_world_size(tp_group)117    heads_local = num_heads // tp_size118    head_dim = w_qkv.shape[0] // 3 // heads_local119    if softmax_scale is None:120        softmax_scale = head_dim ** -0.5121 122    B, S = hidden_states.shape[:2]123    qkv = F.linear(hidden_states, w_qkv).view(B, S, 3, heads_local, head_dim)124    q, k, v = qkv.unbind(dim=2)125 126    context = _ring_attn_forward(cp_group, q.contiguous(), k.contiguous(), v.contiguous(),127                                 float(softmax_scale), causal)128 129    out = F.linear(context.reshape(B, S, -1), w_o)130    if tp_size > 1:131        dist.all_reduce(out, op=dist.ReduceOp.SUM, group=tp_group)132    return out133