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.
0146
1from __future__ import annotations2 3import math4 5import torch6import torch.distributed as dist7import torch.nn.functional as F8from torch import Tensor9from torch._utils import _flatten_dense_tensors, _unflatten_dense_tensors10 11 12def solution(13 X_local: Tensor,14 y_local: Tensor,15 W1: Tensor,16 b1: Tensor,17 W2: Tensor,18 b2: Tensor,19 exp_avg_W1: Tensor,20 exp_avg_b1: Tensor,21 exp_avg_W2: Tensor,22 exp_avg_b2: Tensor,23 exp_avg_sq_W1: Tensor,24 exp_avg_sq_b1: Tensor,25 exp_avg_sq_W2: Tensor,26 exp_avg_sq_b2: Tensor,27 lr: float,28 beta1: float,29 beta2: float,30 eps: float,31 step: int,32) -> tuple[Tensor, ...]:33 world_size = dist.get_world_size()34 35 params = [W1, b1, W2, b2]36 exp_avg = [exp_avg_W1, exp_avg_b1, exp_avg_W2, exp_avg_b2]37 exp_avg_sq = [exp_avg_sq_W1, exp_avg_sq_b1, exp_avg_sq_W2, exp_avg_sq_b2]38 39 flat_params = _flatten_dense_tensors(params)40 dist.broadcast(flat_params, src=0)41 broadcast_params = _unflatten_dense_tensors(flat_params, params)42 params = [t.detach().requires_grad_(True) for t in broadcast_params]43 44 flat_m = _flatten_dense_tensors(exp_avg)45 dist.broadcast(flat_m, src=0)46 exp_avg = list(_unflatten_dense_tensors(flat_m, exp_avg))47 48 flat_v = _flatten_dense_tensors(exp_avg_sq)49 dist.broadcast(flat_v, src=0)50 exp_avg_sq = list(_unflatten_dense_tensors(flat_v, exp_avg_sq))51 52 h = F.relu(F.linear(X_local, params[0], params[1]))53 out = F.linear(h, params[2], params[3])54 loss = F.mse_loss(out, y_local)55 loss.backward()56 57 grads = [p.grad for p in params]58 flat_grad = _flatten_dense_tensors(grads)59 dist.all_reduce(flat_grad, op=dist.ReduceOp.SUM)60 flat_grad.div_(world_size)61 avg_grads = _unflatten_dense_tensors(flat_grad, grads)62 for p, g in zip(params, avg_grads):63 p.grad.copy_(g)64 65 assert step >= 166 bc1 = 1.0 - math.pow(beta1, step)67 bc2 = 1.0 - math.pow(beta2, step)68 69 for p, m_buf, v_buf in zip(params, exp_avg, exp_avg_sq):70 g = p.grad71 m_buf.mul_(beta1).add_(g, alpha=1.0 - beta1)72 v_buf.mul_(beta2).addcmul_(g, g, value=1.0 - beta2)73 m_hat = m_buf / bc174 v_hat = v_buf / bc275 denom = v_hat.sqrt().add(eps)76 p.data.add_(m_hat.div(denom).mul(-lr))77 78 out_tensors = tuple(list(params) + exp_avg + exp_avg_sq)79 return out_tensors80 