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 Optional, Tuple2 3import torch4import torch.distributed as dist5 6 7class RingComm:8 def __init__(self, group: dist.ProcessGroup):9 self._group = group10 self._ops: list = []11 self._reqs = None12 self.rank = dist.get_rank(group)13 self.world_size = dist.get_world_size(group)14 self.send_rank = dist.get_global_rank(group, (self.rank + 1) % self.world_size)15 self.recv_rank = dist.get_global_rank(group, (self.rank - 1) % self.world_size)16 17 def send_recv(self, to_send: torch.Tensor, recv_buf: Optional[torch.Tensor] = None) -> torch.Tensor:18 buf = recv_buf if recv_buf is not None else torch.empty_like(to_send)19 self._ops.append(dist.P2POp(dist.isend, to_send, self.send_rank, group=self._group))20 self._ops.append(dist.P2POp(dist.irecv, buf, self.recv_rank, group=self._group))21 return buf22 23 def commit(self):24 self._reqs = dist.batch_isend_irecv(self._ops)25 26 def wait(self):27 for r in self._reqs:28 r.wait()29 self._reqs = None30 self._ops = []31 32 def send_recv_kv(self, k: torch.Tensor, v: torch.Tensor):33 next_k = self.send_recv(k)34 next_v = self.send_recv(v)35 self.commit()36 return next_k, next_v37 38 39def _local_attn_backward(40 dout: torch.Tensor, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,41 out: torch.Tensor, softmax_lse: torch.Tensor,42 scale: float, causal: bool,43) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:44 qh = q.transpose(1, 2).float()45 kh = k.transpose(1, 2).float()46 vh = v.transpose(1, 2).float()47 doh = dout.transpose(1, 2).float()48 outh = out.transpose(1, 2).float()49 50 scores = torch.matmul(qh, kh.transpose(-2, -1)) * scale51 if causal:52 sq, sk = q.size(1), k.size(1)53 mask = torch.triu(torch.ones(sq, sk, device=q.device, dtype=torch.bool), 1)54 scores.masked_fill_(mask.unsqueeze(0).unsqueeze(0), float("-inf"))55 56 probs = torch.exp(scores - softmax_lse)57 dP = torch.matmul(doh, vh.transpose(-2, -1))58 row_dot = (doh * outh).sum(dim=-1, keepdim=True)59 dS = probs * (dP - row_dot)60 61 dQ = torch.matmul(dS, kh) * scale62 dK = torch.matmul(dS.transpose(-2, -1), qh) * scale63 dV = torch.matmul(probs.transpose(-2, -1), doh)64 65 return (66 dQ.transpose(1, 2).contiguous(),67 dK.transpose(1, 2).contiguous(),68 dV.transpose(1, 2).contiguous(),69 )70 71 72def _ring_attn_backward(73 group: dist.ProcessGroup,74 dout: torch.Tensor, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,75 out: torch.Tensor, softmax_lse: torch.Tensor,76 scale: float, causal: bool,77) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:78 world_size = dist.get_world_size(group)79 lse_4d = softmax_lse.unsqueeze(-1)80 81 if world_size == 1:82 dq, dk, dv = _local_attn_backward(dout, q, k, v, out, lse_4d, scale, causal)83 return dq.to(q.dtype), dk.to(k.dtype), dv.to(v.dtype)84 85 kv_comm = RingComm(group)86 d_kv_comm = RingComm(group)87 88 dq, dk, dv = None, None, None89 next_dk, next_dv = None, None90 next_k, next_v = None, None91 92 for step in range(kv_comm.world_size):93 if step + 1 != kv_comm.world_size:94 next_k, next_v = kv_comm.send_recv_kv(k, v)95 96 if step <= kv_comm.rank or not causal:97 block_dq, block_dk, block_dv = _local_attn_backward(98 dout, q, k, v, out, lse_4d, scale, causal=(causal and step == 0),99 )100 if dq is None:101 dq = block_dq.float()102 dk = block_dk.float()103 dv = block_dv.float()104 else:105 dq = dq + block_dq.float()106 d_kv_comm.wait()107 dk = block_dk.float() + next_dk108 dv = block_dv.float() + next_dv109 elif step != 0:110 d_kv_comm.wait()111 dk, dv = next_dk, next_dv112 113 if step + 1 != kv_comm.world_size:114 kv_comm.wait()115 k, v = next_k, next_v116 117 next_dk, next_dv = d_kv_comm.send_recv_kv(dk, dv)118 119 d_kv_comm.wait()120 121 return dq.to(q.dtype), next_dk.to(k.dtype), next_dv.to(v.dtype)122 123 124def solution(125 dout: torch.Tensor,126 q: torch.Tensor,127 k: torch.Tensor,128 v: torch.Tensor,129 out: torch.Tensor,130 softmax_lse: torch.Tensor,131 softmax_scale: Optional[float] = None,132 causal: bool = False,133 cp_group: Optional[dist.ProcessGroup] = None,134 dp_group: Optional[dist.ProcessGroup] = None,135) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:136 cp_group = cp_group or dist.group.WORLD137 if softmax_scale is None:138 softmax_scale = q.shape[-1] ** -0.5139 140 dq, dk, dv = _ring_attn_backward(141 cp_group, dout, q.contiguous(), k.contiguous(), v.contiguous(),142 out, softmax_lse, float(softmax_scale), causal,143 )144 145 if dp_group is not None and dist.get_world_size(dp_group) > 1:146 dp_size = dist.get_world_size(dp_group)147 for g in (dq, dk, dv):148 dist.all_reduce(g, op=dist.ReduceOp.SUM, group=dp_group)149 g.div_(dp_size)150 151 return dq, dk, dv152 