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.
0266
1import torch2import torch.nn.functional as F3import torch.distributed as dist4from typing import Tuple, Any5 6def solution(7 hidden_states: torch.Tensor,8 weight: torch.Tensor,9 labels: torch.Tensor,10 old_logprobs: torch.Tensor,11 advantages: torch.Tensor,12 ignore_index: int = -100,13) -> Tuple[torch.Tensor, Any, torch.Tensor, torch.Tensor, torch.Tensor]:14 logits = F.linear(hidden_states, weight)15 logits_flat = logits.view(-1, logits.size(-1))16 labels_flat = labels.view(-1)17 18 per_token_ce = F.cross_entropy(logits_flat, labels_flat, ignore_index=ignore_index, reduction='none')19 20 new_logprobs_flat = -per_token_ce.detach()21 old_logprobs_flat = old_logprobs.view(-1)22 advantages_flat = advantages.view(-1)23 24 valid_mask = (labels_flat != ignore_index)25 n_valid_local = valid_mask.sum().float()26 27 n_valid_global = n_valid_local.clone()28 dist.all_reduce(n_valid_global, op=dist.ReduceOp.SUM)29 n_valid_global_clamped = n_valid_global.clamp(min=1.0)30 31 delta = (new_logprobs_flat - old_logprobs_flat).masked_fill(~valid_mask, 0.0).clamp(min=-20.0, max=20.0)32 ratio = torch.exp(delta)33 34 per_token_pg = -(ratio * advantages_flat).masked_fill(~valid_mask, 0.0)35 local_pg_sum = per_token_pg.sum()36 37 global_pg_sum = local_pg_sum.clone()38 dist.all_reduce(global_pg_sum, op=dist.ReduceOp.SUM)39 true_pg = global_pg_sum / n_valid_global_clamped40 41 w = (ratio.detach() * advantages_flat).masked_fill(~valid_mask, 0.0)42 local_surrogate_sum = (w * per_token_ce).sum()43 44 surrogate = local_surrogate_sum / n_valid_global_clamped45 46 loss = true_pg.detach() + surrogate - surrogate.detach()47 48 ratio_valid = ratio.masked_fill(~valid_mask, 0.0)49 sum_ratio_local = ratio_valid.sum()50 dist.all_reduce(sum_ratio_local, op=dist.ReduceOp.SUM)51 ratio_mean = sum_ratio_local / n_valid_global_clamped52 53 ratio_for_min = ratio.masked_fill(~valid_mask, float('inf'))54 min_ratio_local = ratio_for_min.min() if n_valid_local > 0 else torch.tensor(float('inf'), device=ratio.device)55 dist.all_reduce(min_ratio_local, op=dist.ReduceOp.MIN)56 57 ratio_for_max = ratio.masked_fill(~valid_mask, float('-inf'))58 max_ratio_local = ratio_for_max.max() if n_valid_local > 0 else torch.tensor(float('-inf'), device=ratio.device)59 dist.all_reduce(max_ratio_local, op=dist.ReduceOp.MAX)60 61 k3_local = (ratio - delta - 1.0).masked_fill(~valid_mask, 0.0).sum()62 dist.all_reduce(k3_local, op=dist.ReduceOp.SUM)63 k3_mean = k3_local / n_valid_global_clamped64 65 entropy_local = per_token_ce.detach().masked_fill(~valid_mask, 0.0).sum()66 dist.all_reduce(entropy_local, op=dist.ReduceOp.SUM)67 entropy_mean = entropy_local / n_valid_global_clamped68 69 metrics = torch.stack([ratio_mean, min_ratio_local, max_ratio_local, k3_mean, entropy_mean])70 71 per_token_logprobs = new_logprobs_flat.view_as(labels)72 per_token_loss = per_token_pg.view_as(labels)73 74 return loss, None, per_token_logprobs, per_token_loss, metrics75 