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.
0123
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 _pp_recv_forward(104 pp_group: dist.ProcessGroup,105 tensor_shape: Tuple[int, ...],106 dtype: torch.dtype,107 device: torch.device,108) -> torch.Tensor:109 prev_rank = dist.get_global_rank(110 pp_group, (dist.get_rank(pp_group) - 1) % dist.get_world_size(pp_group)111 )112 buf = torch.empty(tensor_shape, dtype=dtype, device=device)113 reqs = dist.batch_isend_irecv([dist.P2POp(dist.irecv, buf, prev_rank, group=pp_group)])114 for r in reqs:115 r.wait()116 return buf117 118 119def _pp_send_forward(120 pp_group: dist.ProcessGroup,121 tensor: torch.Tensor,122) -> None:123 next_rank = dist.get_global_rank(124 pp_group, (dist.get_rank(pp_group) + 1) % dist.get_world_size(pp_group)125 )126 reqs = dist.batch_isend_irecv([dist.P2POp(dist.isend, tensor.contiguous(), next_rank, group=pp_group)])127 for r in reqs:128 r.wait()129 130 131def _attention_block(132 hidden: torch.Tensor, w_qkv: torch.Tensor, w_o: torch.Tensor,133 num_heads: int, scale: float, causal: bool,134 cp_group: dist.ProcessGroup,135) -> torch.Tensor:136 B, S, D = hidden.shape137 head_dim = w_qkv.shape[0] // 3 // num_heads138 qkv = F.linear(hidden, w_qkv).view(B, S, 3, num_heads, head_dim)139 q, k, v = qkv.unbind(dim=2)140 ctx = _ring_attn_forward(cp_group, q.contiguous(), k.contiguous(), v.contiguous(),141 scale, causal)142 return F.linear(ctx.reshape(B, S, -1), w_o)143 144 145def solution(146 hidden_states: torch.Tensor,147 w_qkv: torch.Tensor,148 w_o: torch.Tensor,149 num_heads: int,150 softmax_scale: Optional[float] = None,151 causal: bool = False,152 cp_group: Optional[dist.ProcessGroup] = None,153 pp_group: Optional[dist.ProcessGroup] = None,154) -> torch.Tensor:155 cp_group = cp_group or dist.group.WORLD156 head_dim = w_qkv.shape[0] // 3 // num_heads157 scale = float(softmax_scale if softmax_scale is not None else head_dim ** -0.5)158 159 is_first = True160 is_last = True161 if pp_group is not None and dist.get_world_size(pp_group) > 1:162 pp_rank = dist.get_rank(pp_group)163 pp_size = dist.get_world_size(pp_group)164 is_first = (pp_rank == 0)165 is_last = (pp_rank == pp_size - 1)166 167 if is_first:168 stage_input = hidden_states169 else:170 stage_input = _pp_recv_forward(171 pp_group, tuple(hidden_states.shape), hidden_states.dtype, hidden_states.device,172 )173 174 stage_output = _attention_block(stage_input, w_qkv, w_o, num_heads, scale, causal, cp_group)175 176 if not is_last and pp_group is not None:177 _pp_send_forward(pp_group, stage_output)178 179 return stage_output180 