Team Ai
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 5mo agoView on Hugging Face
0likes146downloads
40_ddp.py80 linesDownload Raw Back to reference
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