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
1import torch2import torch.distributed as dist3from typing import Tuple, Optional4 5 6def forward(7 loss: torch.Tensor,8 local_valid_tokens: torch.Tensor,9 global_valid_tokens: torch.Tensor,10) -> Tuple[torch.Tensor, torch.Tensor]:11 if local_valid_tokens.item() == 0:12 loss = torch.nan_to_num(loss)13 14 loss_sum = loss * local_valid_tokens15 dist.all_reduce(loss_sum, op=dist.ReduceOp.SUM)16 17 normalized_loss = loss_sum / global_valid_tokens18 return normalized_loss, loss_sum19 20 21def backward(22 local_valid_tokens: torch.Tensor,23 global_valid_tokens: torch.Tensor,24 grad_normalized_loss: torch.Tensor,25 grad_loss_sum: Optional[torch.Tensor],26) -> torch.Tensor:27 grad_from_normalized = grad_normalized_loss * local_valid_tokens / global_valid_tokens28 29 if grad_loss_sum is not None:30 grad_from_sum = grad_loss_sum * local_valid_tokens31 else:32 grad_from_sum = torch.zeros_like(grad_normalized_loss, device=grad_normalized_loss.device)33 34 return grad_from_normalized + grad_from_sum35 36 37def solution(38 loss: torch.Tensor,39 local_valid_tokens: torch.Tensor,40 global_valid_tokens: torch.Tensor,41 grad_normalized_loss: torch.Tensor,42 grad_loss_sum: Optional[torch.Tensor] = None,43) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:44 normalized_loss, loss_sum = forward(loss, local_valid_tokens, global_valid_tokens)45 46 grad_loss = backward(47 local_valid_tokens,48 global_valid_tokens,49 grad_normalized_loss,50 grad_loss_sum,51 )52 53 return normalized_loss, loss_sum, grad_loss54 