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 Optional2 3import torch4import torch.distributed as dist5 6 7@torch.no_grad()8def solution(9 teacher_output: torch.Tensor,10 teacher_temp: float,11 n_masked_patches_tensor: torch.Tensor,12 n_iterations: int = 3,13 group: Optional[dist.ProcessGroup] = None,14) -> torch.Tensor:15 group = group or dist.group.WORLD16 q = torch.exp(teacher_output.float() / teacher_temp).T17 total_batch = n_masked_patches_tensor.to(device=q.device, dtype=q.dtype).clone()18 dist.all_reduce(total_batch, group=group)19 20 num_prototypes = q.shape[0]21 total_mass = q.sum()22 dist.all_reduce(total_mass, group=group)23 q /= total_mass24 25 for _ in range(n_iterations):26 row_sum = q.sum(dim=1, keepdim=True)27 dist.all_reduce(row_sum, group=group)28 q /= row_sum29 q /= num_prototypes30 31 q /= q.sum(dim=0, keepdim=True)32 q /= total_batch33 34 q *= total_batch35 return q.T.contiguous()36 