CoolFace
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 4mo agoView on Hugging Face
0likes266downloads
25_importance_sampling_loss.py75 linesDownload Raw Back to reference
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