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
61_gsplat_3d_gaussian_splatting.py451 linesDownload Raw Back to reference
1import math2from typing import List, Optional, Tuple, Union3 4import torch5import torch.distributed as dist6import torch.distributed.nn.functional as distF7import torch.nn.functional as F8from torch import Tensor9 10 11def _all_gather_int32(12    world_size: int, value: Union[int, Tensor], device: Optional[torch.device] = None13) -> List[int]:14    if world_size == 1:15        return [value]16 17    if isinstance(value, int):18        assert device is not None, "device is required for scalar input"19        value_tensor = torch.tensor(value, dtype=torch.int, device=device)20    else:21        value_tensor = value22 23    collected = torch.empty(24        world_size, dtype=value_tensor.dtype, device=value_tensor.device25    )26    dist.all_gather_into_tensor(collected, value_tensor)27 28    if isinstance(value, int):29        return collected.tolist()30    else:31        return collected.unbind()32 33 34def _all_to_all_int32(35    world_size: int,36    values: List[Union[int, Tensor]],37    device: Optional[torch.device] = None,38) -> List[int]:39    if world_size == 1:40        return values41 42    assert len(values) == world_size43 44    if any(isinstance(v, int) for v in values):45        assert device is not None, "device is required for scalar input"46 47    values_tensor = [48        (torch.tensor(v, dtype=torch.int, device=device) if isinstance(v, int) else v)49        for v in values50    ]51 52    collected = [torch.empty_like(v) for v in values_tensor]53    dist.all_to_all(collected, values_tensor)54 55    return [56        v.item() if isinstance(tensor, int) else v57        for v, tensor in zip(collected, values)58    ]59 60 61def _all_gather_tensor_list(world_size: int, tensor_list: List[Tensor]) -> List[Tensor]:62    if world_size == 1:63        return tensor_list64 65    N = len(tensor_list[0])66    for tensor in tensor_list:67        assert len(tensor) == N, "All tensors should have the same first dimension size"68 69    data = torch.cat([t.reshape(N, -1) for t in tensor_list], dim=-1)70    sizes = [t.numel() // N for t in tensor_list]71 72    if data.requires_grad:73        collected = distF.all_gather(data)74    else:75        collected = [torch.empty_like(data) for _ in range(world_size)]76        dist.all_gather(collected, data)77    collected = torch.cat(collected, dim=0)78 79    out_tensor_tuple = torch.split(collected, sizes, dim=-1)80    out_tensor_list = []81    for out_tensor, tensor in zip(out_tensor_tuple, tensor_list):82        out_tensor = out_tensor.view(-1, *tensor.shape[1:])83        out_tensor_list.append(out_tensor)84    return out_tensor_list85 86 87def _all_to_all_tensor_list(88    world_size: int,89    tensor_list: List[Tensor],90    splits: List[Union[int, Tensor]],91    output_splits: Optional[List[Union[int, Tensor]]] = None,92) -> List[Tensor]:93    if world_size == 1:94        return tensor_list95 96    N = len(tensor_list[0])97    for tensor in tensor_list:98        assert len(tensor) == N, "All tensors should have the same first dimension size"99 100    assert len(splits) == world_size101 102    data = torch.cat([t.reshape(N, -1) for t in tensor_list], dim=-1)103    sizes = [t.numel() // N for t in tensor_list]104 105    if output_splits is not None:106        collected_splits = output_splits107    else:108        collected_splits = _all_to_all_int32(world_size, splits, device=data.device)109    collected = [110        torch.empty((l, *data.shape[1:]), dtype=data.dtype, device=data.device)111        for l in collected_splits112    ]113    splits = [s.item() if isinstance(s, Tensor) else s for s in splits]114    if data.requires_grad:115        distF.all_to_all(collected, data.split(splits, dim=0))116    else:117        dist.all_to_all(collected, list(data.split(splits, dim=0)))118    collected = torch.cat(collected, dim=0)119 120    out_tensor_tuple = torch.split(collected, sizes, dim=-1)121    out_tensor_list = []122    for out_tensor, tensor in zip(out_tensor_tuple, tensor_list):123        out_tensor = out_tensor.view(-1, *tensor.shape[1:])124        out_tensor_list.append(out_tensor)125    return out_tensor_list126 127 128def _quat_to_rotmat(quats: Tensor) -> Tensor:129    quats = F.normalize(quats, p=2, dim=-1)130    w, x, y, z = torch.unbind(quats, dim=-1)131    R = torch.stack(132        [133            1 - 2 * (y**2 + z**2),134            2 * (x * y - w * z),135            2 * (x * z + w * y),136            2 * (x * y + w * z),137            1 - 2 * (x**2 + z**2),138            2 * (y * z - w * x),139            2 * (x * z - w * y),140            2 * (y * z + w * x),141            1 - 2 * (x**2 + y**2),142        ],143        dim=-1,144    )145    return R.reshape(quats.shape[:-1] + (3, 3))146 147 148def _quat_scale_to_covar_preci(149    quats: Tensor,150    scales: Tensor,151    compute_covar: bool = True,152    compute_preci: bool = True,153    triu: bool = False,154) -> Tuple[Optional[Tensor], Optional[Tensor]]:155    batch_dims = quats.shape[:-1]156    assert quats.shape == batch_dims + (4,), quats.shape157    assert scales.shape == batch_dims + (3,), scales.shape158    R = _quat_to_rotmat(quats)159 160    if compute_covar:161        M = R * scales[..., None, :]162        covars = torch.einsum("...ij,...kj -> ...ik", M, M)163        if triu:164            covars = covars.reshape(batch_dims + (9,))165            covars = (166                covars[..., [0, 1, 2, 4, 5, 8]] + covars[..., [0, 3, 6, 4, 7, 8]]167            ) / 2.0168    if compute_preci:169        P = R * (1 / scales[..., None, :])170        precis = torch.einsum("...ij,...kj -> ...ik", P, P)171        if triu:172            precis = precis.reshape(batch_dims + (9,))173            precis = (174                precis[..., [0, 1, 2, 4, 5, 8]] + precis[..., [0, 3, 6, 4, 7, 8]]175            ) / 2.0176 177    return covars if compute_covar else None, precis if compute_preci else None178 179 180def _world_to_cam(181    means: Tensor,182    covars: Tensor,183    viewmats: Tensor,184) -> Tuple[Tensor, Tensor]:185    batch_dims = means.shape[:-2]186    N = means.shape[-2]187    C = viewmats.shape[-3]188    assert means.shape == batch_dims + (N, 3), means.shape189    assert covars.shape == batch_dims + (N, 3, 3), covars.shape190    assert viewmats.shape == batch_dims + (C, 4, 4), viewmats.shape191 192    R = viewmats[..., :3, :3]193    t = viewmats[..., :3, 3]194    means_c = (195        torch.einsum("...cij,...nj->...cni", R, means) + t[..., None, :]196    )197    covars_c = torch.einsum(198        "...cij,...njk,...clk->...cnil", R, covars, R199    )200    return means_c, covars_c201 202 203def _persp_proj(204    means: Tensor,205    covars: Tensor,206    Ks: Tensor,207    width: int,208    height: int,209) -> Tuple[Tensor, Tensor]:210    batch_dims = means.shape[:-3]211    C, N = means.shape[-3:-1]212    assert means.shape == batch_dims + (C, N, 3), means.shape213    assert covars.shape == batch_dims + (C, N, 3, 3), covars.shape214    assert Ks.shape == batch_dims + (C, 3, 3), Ks.shape215 216    tx, ty, tz = torch.unbind(means, dim=-1)217    tz2 = tz**2218 219    fx = Ks[..., 0, 0, None]220    fy = Ks[..., 1, 1, None]221    cx = Ks[..., 0, 2, None]222    cy = Ks[..., 1, 2, None]223    tan_fovx = 0.5 * width / fx224    tan_fovy = 0.5 * height / fy225 226    lim_x_pos = (width - cx) / fx + 0.3 * tan_fovx227    lim_x_neg = cx / fx + 0.3 * tan_fovx228    lim_y_pos = (height - cy) / fy + 0.3 * tan_fovy229    lim_y_neg = cy / fy + 0.3 * tan_fovy230    tx = tz * torch.clamp(tx / tz, min=-lim_x_neg, max=lim_x_pos)231    ty = tz * torch.clamp(ty / tz, min=-lim_y_neg, max=lim_y_pos)232 233    O = torch.zeros(batch_dims + (C, N), device=means.device, dtype=means.dtype)234    J = torch.stack(235        [fx / tz, O, -fx * tx / tz2, O, fy / tz, -fy * ty / tz2], dim=-1236    ).reshape(batch_dims + (C, N, 2, 3))237 238    cov2d = torch.einsum("...ij,...jk,...kl->...il", J, covars, J.transpose(-1, -2))239    means2d = torch.einsum(240        "...ij,...nj->...ni", Ks[..., :2, :3], means241    )242    means2d = means2d / tz[..., None]243    return means2d, cov2d244 245 246def _fully_fused_projection(247    means: Tensor,248    covars: Tensor,249    viewmats: Tensor,250    Ks: Tensor,251    width: int,252    height: int,253    eps2d: float = 0.3,254    near_plane: float = 0.01,255    far_plane: float = 1e10,256    calc_compensations: bool = False,257    camera_model: str = "pinhole",258) -> Tuple[Tensor, Tensor, Tensor, Tensor, Optional[Tensor]]:259    batch_dims = means.shape[:-2]260    N = means.shape[-2]261    C = viewmats.shape[-3]262    assert means.shape == batch_dims + (N, 3), means.shape263    assert covars.shape == batch_dims + (N, 3, 3), covars.shape264    assert viewmats.shape == batch_dims + (C, 4, 4), viewmats.shape265    assert Ks.shape == batch_dims + (C, 3, 3), Ks.shape266    assert camera_model == "pinhole", "only pinhole supported"267 268    means_c, covars_c = _world_to_cam(means, covars, viewmats)269    means2d, covars2d = _persp_proj(means_c, covars_c, Ks, width, height)270 271    det_orig = (272        covars2d[..., 0, 0] * covars2d[..., 1, 1]273        - covars2d[..., 0, 1] * covars2d[..., 1, 0]274    )275    covars2d = covars2d + torch.eye(2, device=means.device, dtype=means.dtype) * eps2d276 277    det = (278        covars2d[..., 0, 0] * covars2d[..., 1, 1]279        - covars2d[..., 0, 1] * covars2d[..., 1, 0]280    )281    det = det.clamp(min=1e-10)282 283    if calc_compensations:284        compensations = torch.sqrt(torch.clamp(det_orig / det, min=0.0))285    else:286        compensations = None287 288    conics = torch.stack(289        [290            covars2d[..., 1, 1] / det,291            -(covars2d[..., 0, 1] + covars2d[..., 1, 0]) / 2.0 / det,292            covars2d[..., 0, 0] / det,293        ],294        dim=-1,295    )296 297    depths = means_c[..., 2]298 299    # CUDA fully_fused_projection can use opacities for tighter radii; this torch300    # implementation follows gsplat/cuda/_torch_impl.py and uses covariance only.301    radius_x = torch.ceil(3.33 * torch.sqrt(covars2d[..., 0, 0]))302    radius_y = torch.ceil(3.33 * torch.sqrt(covars2d[..., 1, 1]))303    radius = torch.stack([radius_x, radius_y], dim=-1)304 305    valid = (depths > near_plane) & (depths < far_plane)306    radius[~valid] = 0.0307 308    inside = (309        (means2d[..., 0] + radius[..., 0] > 0)310        & (means2d[..., 0] - radius[..., 0] < width)311        & (means2d[..., 1] + radius[..., 1] > 0)312        & (means2d[..., 1] - radius[..., 1] < height)313    )314    radius[~inside] = 0.0315 316    radii = radius.int()317    return radii, means2d, depths, conics, compensations318 319 320def _pack_projection_results(321    radii: Tensor,322    means2d: Tensor,323    depths: Tensor,324    conics: Tensor,325    compensations: Optional[Tensor],326) -> Tuple[Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Optional[Tensor]]:327    C, N = radii.shape[:2]328    device = radii.device329 330    valid = (radii > 0).all(dim=-1)331    camera_ids, gaussian_ids = torch.where(valid)332    camera_ids = camera_ids.int()333    gaussian_ids = gaussian_ids.int()334 335    radii_packed = radii[valid]336    means2d_packed = means2d[valid]337    depths_packed = depths[valid]338    conics_packed = conics[valid]339    compensations_packed = compensations[valid] if compensations is not None else None340 341    counts = torch.bincount(camera_ids.long(), minlength=C)342    indptr = torch.zeros(C + 1, dtype=torch.int32, device=device)343    indptr[1:] = torch.cumsum(counts, dim=0).int()344 345    return (346        camera_ids, gaussian_ids, indptr,347        radii_packed, means2d_packed, depths_packed, conics_packed,348        compensations_packed,349    )350 351 352def solution(353    means: Tensor,354    quats: Tensor,355    scales: Tensor,356    opacities: Tensor,357    colors: Tensor,358    viewmats: Tensor,359    Ks: Tensor,360    image_width: int,361    image_height: int,362    eps2d: float = 0.3,363    near_plane: float = 0.01,364    far_plane: float = 1e10,365    camera_model: str = "pinhole",366) -> Tuple[Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor]:367    world_rank = dist.get_rank()368    world_size = dist.get_world_size()369    device = means.device370 371    N = means.shape[0]372    C = viewmats.shape[0]373    D = colors.shape[1]374 375    N_world = _all_gather_int32(world_size, N, device=device)376    C_world = [C] * world_size377 378    viewmats, Ks = _all_gather_tensor_list(world_size, [viewmats, Ks])379 380    C = len(viewmats)381 382    covars, _ = _quat_scale_to_covar_preci(383        quats, scales, compute_covar=True, compute_preci=False, triu=False384    )385 386    radii, means2d, depths, conics, compensations = _fully_fused_projection(387        means, covars, viewmats, Ks, image_width, image_height,388        eps2d=eps2d, near_plane=near_plane, far_plane=far_plane,389        calc_compensations=False,390        camera_model=camera_model,391    )392 393    (394        camera_ids, gaussian_ids, _indptr,395        radii, means2d, depths, conics, _compensations,396    ) = _pack_projection_results(radii, means2d, depths, conics, compensations)397 398    opacities = opacities[gaussian_ids.long()]399    colors = colors[gaussian_ids.long()]400 401    cnts = torch.bincount(camera_ids.long(), minlength=C)402    cnts = cnts.split(C_world, dim=0)403    cnts = [cuts.sum() for cuts in cnts]404 405    collected_splits = _all_to_all_int32(world_size, cnts, device=device)406 407    (radii,) = _all_to_all_tensor_list(408        world_size, [radii], cnts, output_splits=collected_splits409    )410 411    (means2d, depths, conics, opacities, colors) = _all_to_all_tensor_list(412        world_size,413        [means2d, depths, conics, opacities, colors],414        cnts,415        output_splits=collected_splits,416    )417 418    offsets = torch.tensor(419        [0] + C_world[:-1], device=camera_ids.device, dtype=camera_ids.dtype420    )421    offsets = torch.cumsum(offsets, dim=0)422    offsets = offsets.repeat_interleave(torch.stack(cnts))423    camera_ids = camera_ids - offsets424 425    offsets = torch.tensor(426        [0] + N_world[:-1],427        device=gaussian_ids.device,428        dtype=gaussian_ids.dtype,429    )430    offsets = torch.cumsum(offsets, dim=0)431    offsets = offsets.repeat_interleave(torch.stack(cnts))432    gaussian_ids = gaussian_ids + offsets433 434    (camera_ids, gaussian_ids) = _all_to_all_tensor_list(435        world_size,436        [camera_ids, gaussian_ids],437        cnts,438        output_splits=collected_splits,439    )440 441    return (442        camera_ids,443        gaussian_ids,444        radii,445        means2d,446        depths,447        conics,448        opacities,449        colors,450    )451