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
0likes123downloads
81_sam3_allgather_iou_suppression.py80 linesDownload Raw Back to reference
1from typing import List, Optional, Tuple2 3import torch4import torch.distributed as dist5 6_NO_OBJ_LOGIT = -10.07 8 9def _mask_iou(lhs: torch.Tensor, rhs: torch.Tensor) -> torch.Tensor:10    lhs_flat = lhs.flatten(1).float()11    rhs_flat = rhs.flatten(1).float()12    intersection = lhs_flat @ rhs_flat.T13    lhs_area = lhs_flat.sum(dim=1)14    rhs_area = rhs_flat.sum(dim=1)15    union = lhs_area[:, None] + rhs_area[None, :] - intersection16    return intersection / union.clamp_min(1.0)17 18 19def _all_gather_variable(20    tensor: torch.Tensor,21    counts: List[int],22    group: dist.ProcessGroup,23) -> torch.Tensor:24    recv = [tensor.new_empty((count, *tensor.shape[1:])) for count in counts]25    dist.all_gather(recv, tensor.contiguous(), group=group)26    return torch.cat(recv, dim=0)27 28 29def _suppression_mask(30    masks: torch.Tensor,31    last_occluded: torch.Tensor,32    iou_threshold: float,33    reverse: bool,34) -> torch.Tensor:35    num_objects = masks.shape[0]36    suppress = torch.zeros(num_objects, dtype=torch.bool, device=masks.device)37    if num_objects <= 1:38        return suppress39 40    overlaps = torch.triu(_mask_iou(masks, masks) >= iou_threshold, diagonal=1)41    last_i = last_occluded.view(num_objects, 1)42    last_j = last_occluded.view(1, num_objects)43    cmp = torch.lt if reverse else torch.gt44 45    suppress_i = overlaps & cmp(last_i, last_j) & (last_j > -1)46    suppress_j = overlaps & cmp(last_j, last_i) & (last_i > -1)47    return suppress_i.any(dim=1) | suppress_j.any(dim=0)48 49 50@torch.no_grad()51def solution(52    low_res_masks_local: torch.Tensor,53    obj_scores_local: torch.Tensor,54    num_obj_per_gpu: List[int],55    last_occluded: torch.Tensor,56    iou_threshold: float = 0.7,57    reverse: bool = False,58    group: Optional[dist.ProcessGroup] = None,59) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:60    group = group or dist.group.WORLD61    rank = dist.get_rank(group=group)62    expected = int(num_obj_per_gpu[rank])63    if low_res_masks_local.shape[0] != expected:64        raise ValueError("local mask count does not match num_obj_per_gpu")65    if obj_scores_local.shape[0] != expected:66        raise ValueError("local score count does not match num_obj_per_gpu")67 68    masks_local = low_res_masks_local.float().contiguous()69    scores_local = obj_scores_local.float().contiguous()70    masks_global = _all_gather_variable(masks_local, num_obj_per_gpu, group)71    scores_global = _all_gather_variable(scores_local, num_obj_per_gpu, group)72 73    last_occluded = last_occluded.to(device=masks_global.device, dtype=torch.long)74    binary_masks = masks_global > 075    to_suppress = _suppression_mask(76        binary_masks, last_occluded, iou_threshold=iou_threshold, reverse=reverse77    )78    masks_global[to_suppress] = _NO_OBJ_LOGIT79    return masks_global, scores_global, to_suppress80