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
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 