submission 68043
shellsmile15795 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 266 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-vectoradd-v2-68043?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:2b2411b6dac8fab4e6f5a4a137176f9e3a3704932922da1aed319d6b23623d74
license declaredunknown
license concludedunknown
authorsshellsmile15795
imported2026-08-15
Kernel source
submission.py266 lines
from typing import Any, Dict, Tuple
import torch
from task import input_t, output_t
import cutlass.cute as cute
from cutlass.cute.runtime import from_dlpack
from torch.utils.cpp_extension import load_inline
_KERNEL_CACHE: Dict[Tuple[Tuple[int, ...], torch.dtype], Any] = {}
_VECADD_EXT = None
def _get_vecadd_extension():
global _VECADD_EXT
if _VECADD_EXT is not None:
return _VECADD_EXT
cuda_src = r"""
#include <cuda_fp16.h>
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
namespace {
__device__ __forceinline__ ulonglong4 add_ulonglong4(ulonglong4 va, ulonglong4 vb) {
ulonglong4 vc;
__half2* ha = reinterpret_cast<__half2*>(&va);
__half2* hb = reinterpret_cast<__half2*>(&vb);
__half2* hc = reinterpret_cast<__half2*>(&vc);
#pragma unroll
for (int i = 0; i < 8; ++i) {
hc[i] = __hadd2(ha[i], hb[i]);
}
return vc;
}
__global__ void vecadd_fp16_kernel(
const half* __restrict__ a,
const half* __restrict__ b,
half* __restrict__ out,
size_t elements
) {
constexpr int pack_elems = 16; // 16 fp16 values per ulonglong4
const size_t tid = threadIdx.x + blockIdx.x * blockDim.x;
const size_t stride = static_cast<size_t>(blockDim.x) * gridDim.x;
const size_t total_packs = elements / pack_elems;
auto* a_vec = reinterpret_cast<const ulonglong4*>(a);
auto* b_vec = reinterpret_cast<const ulonglong4*>(b);
auto* o_vec = reinterpret_cast<ulonglong4*>(out);
const size_t stride_packs = stride;
auto* a_ptr = a_vec + tid;
auto* b_ptr = b_vec + tid;
auto* o_ptr = o_vec + tid;
for (size_t pack = tid; pack < total_packs; pack += stride_packs * 2) {
if (pack < total_packs) {
ulonglong4 va0 = *a_ptr;
ulonglong4 vb0 = *b_ptr;
*o_ptr = add_ulonglong4(va0, vb0);
}
a_ptr += stride_packs;
b_ptr += stride_packs;
o_ptr += stride_packs;
const size_t pack1 = pack + stride_packs;
if (pack1 < total_packs) {
ulonglong4 va1 = *a_ptr;
ulonglong4 vb1 = *b_ptr;
*o_ptr = add_ulonglong4(va1, vb1);
}
a_ptr += stride_packs;
b_ptr += stride_packs;
o_ptr += stride_packs;
}
const size_t tail_start = total_packs * pack_elems;
for (size_t idx = tail_start + tid; idx < elements; idx += stride) {
out[idx] = __hadd(a[idx], b[idx]);
}
}
} // namespace
void vecadd_fp16(torch::Tensor a, torch::Tensor b, torch::Tensor out) {
TORCH_CHECK(a.is_cuda(), "Input tensor A must be CUDA");
TORCH_CHECK(b.is_cuda(), "Input tensor B must be CUDA");
TORCH_CHECK(out.is_cuda(), "Output tensor must be CUDA");
TORCH_CHECK(a.dtype() == torch::kFloat16, "A must be float16");
TORCH_CHECK(b.dtype() == torch::kFloat16, "B must be float16");
TORCH_CHECK(out.dtype() == torch::kFloat16, "output must be float16");
TORCH_CHECK(a.numel() == b.numel(), "Input sizes must match");
TORCH_CHECK(a.numel() == out.numel(), "Output size mismatch");
auto stream = at::cuda::getCurrentCUDAStream();
const auto elements = static_cast<size_t>(out.numel());
constexpr int block = 256;
const int sms = at::cuda::getCurrentDeviceProperties()->multiProcessorCount;
constexpr int pack_elems = 16;
const size_t total_packs = elements / pack_elems;
int grid = static_cast<int>((total_packs + block - 1) / block);
grid = std::max(1, grid);
vecadd_fp16_kernel<<<grid, block, 0, stream>>>(
reinterpret_cast<const half*>(a.data_ptr<at::Half>()),
reinterpret_cast<const half*>(b.data_ptr<at::Half>()),
reinterpret_cast<half*>(out.data_ptr<at::Half>()),
elements
);
AT_CUDA_CHECK(cudaGetLastError());
}
"""
_VECADD_EXT = load_inline(
name="vectoradd_fp16_ext",
cpp_sources="#include <torch/extension.h>\nvoid vecadd_fp16(torch::Tensor a, torch::Tensor b, torch::Tensor out);\n",
cuda_sources=cuda_src,
functions=["vecadd_fp16"],
with_cuda=True,
extra_cuda_cflags=[
"-O3",
"--use_fast_math",
"-U__CUDA_NO_HALF_OPERATORS__",
"-U__CUDA_NO_HALF_CONVERSIONS__",
],
)
return _VECADD_EXT
def _is_aligned(tensor: torch.Tensor, alignment: int = 16) -> bool:
return tensor.data_ptr() % alignment == 0
@cute.kernel
def _tile_elementwise_add_kernel(
gA: cute.Tensor,
gB: cute.Tensor,
gC: cute.Tensor,
tv_layout,
) -> None:
"""Block-tiled elementwise add with per-thread vectorization."""
tidx, _, _ = cute.arch.thread_idx()
bidx, _, _ = cute.arch.block_idx()
grid_x, _, _ = cute.arch.grid_dim()
num_tiles = cute.size(gC, mode=[1])
tile_idx = bidx
while tile_idx < num_tiles:
blk_coord = ((None, None), tile_idx)
blkA = gA[blk_coord]
blkB = gB[blk_coord]
blkC = gC[blk_coord]
tidfrgA = cute.composition(blkA, tv_layout)
tidfrgB = cute.composition(blkB, tv_layout)
tidfrgC = cute.composition(blkC, tv_layout)
thr_coord = (tidx, None)
thrA = tidfrgA[thr_coord]
thrB = tidfrgB[thr_coord]
thrC = tidfrgC[thr_coord]
thrC[None] = thrA.load() + thrB.load()
tile_idx += grid_x
@cute.jit
def _tiled_elementwise_add(
mA: cute.Tensor,
mB: cute.Tensor,
mC: cute.Tensor,
) -> None:
"""Launch tiled kernel honoring remainder shapes."""
element_type = mA.element_type
bytes_per_element = element_type.width // 8
coalesced_ldst_bytes = bytes_per_element
thr_layout = cute.make_ordered_layout((4, 64), order=(1, 0))
val_layout = cute.make_ordered_layout((16, coalesced_ldst_bytes), order=(1, 0))
val_layout = cute.recast_layout(element_type.width, 8, val_layout)
tiler_mn, tv_layout = cute.make_layout_tv(thr_layout, val_layout)
gA = cute.zipped_divide(mA, tiler_mn)
gB = cute.zipped_divide(mB, tiler_mn)
gC = cute.zipped_divide(mC, tiler_mn)
remap_block = cute.make_ordered_layout(
cute.select(gA.shape[1], mode=[1, 0]), order=(1, 0)
)
gA = cute.composition(gA, (None, remap_block))
gB = cute.composition(gB, (None, remap_block))
gC = cute.composition(gC, (None, remap_block))
total_tiles = int(cute.size(gC, mode=[1]))
block_dim_x = int(cute.size(tv_layout, mode=[0]))
max_ctas = 128
if torch.cuda.is_available():
props = torch.cuda.get_device_properties(torch.cuda.current_device())
max_ctas = max(1, props.multi_processor_count * 2)
grid_dim_x = min(total_tiles, max_ctas) if total_tiles > 0 else 1
_tile_elementwise_add_kernel(gA, gB, gC, tv_layout).launch(
grid=(grid_dim_x, 1, 1),
block=(block_dim_x, 1, 1),
)
def _ensure_contiguous(tensor: torch.Tensor) -> torch.Tensor:
if tensor.is_contiguous():
return tensor
return tensor.contiguous()
def _torch_to_cute(tensor: torch.Tensor) -> cute.Tensor:
return from_dlpack(tensor, assumed_align=16)
def _get_compiled_kernel(
key: Tuple[Tuple[int, ...], torch.dtype],
a_: cute.Tensor,
b_: cute.Tensor,
c_: cute.Tensor,
):
compiled = _KERNEL_CACHE.get(key)
if compiled is None:
compiled = cute.compile(_tiled_elementwise_add, a_, b_, c_)
_KERNEL_CACHE[key] = compiled
return compiled
def custom_kernel(data: input_t) -> output_t:
"""Run vector addition via a tiled Cutlass (CuTe) kernel."""
A, B, output = data
A_contig = _ensure_contiguous(A)
B_contig = _ensure_contiguous(B)
output_contig = output if output.is_contiguous() else output.contiguous()
a_cute = _torch_to_cute(A_contig)
b_cute = _torch_to_cute(B_contig)
c_cute = _torch_to_cute(output_contig)
if (
A_contig.is_cuda
and A_contig.dtype == torch.float16
and _is_aligned(A_contig, 32)
and _is_aligned(B_contig, 32)
and _is_aligned(output_contig, 32)
):
ext = _get_vecadd_extension()
ext.vecadd_fp16(A_contig, B_contig, output_contig)
else:
kernel_key = (tuple(A_contig.shape), A_contig.dtype)
compiled_kernel = _get_compiled_kernel(
kernel_key, a_cute, b_cute, c_cute
)
compiled_kernel(a_cute, b_cute, c_cute)
if output_contig is not output:
output.copy_(output_contig)
return output
scrolls · 266 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 67297.
- from __future__ import annotations-from typing import Any, Dict, Tupleimport torch+ from task import input_t, output_t++ import cutlass.cute as cute+ from cutlass.cute.runtime import from_dlpackfrom torch.utils.cpp_extension import load_inline- from task import input_t, output_t+ _KERNEL_CACHE: Dict[Tuple[Tuple[int, ...], torch.dtype], Any] = {}+ _VECADD_EXT = None- CPP_DECL = r"""- #include <torch/extension.h>- void launch_vector_add(torch::Tensor A, torch::Tensor B, torch::Tensor C);- """+ def _get_vecadd_extension():+ global _VECADD_EXT+ if _VECADD_EXT is not None:+ return _VECADD_EXT- CUDA_SRC = r"""- #include <ATen/cuda/CUDAContext.h>- #include <c10/cuda/CUDAGuard.h>- #include <c10/cuda/CUDAException.h>+ cuda_src = r"""#include <cuda_fp16.h>- #include <cstdint>+ #include <torch/extension.h>+ #include <ATen/cuda/CUDAContext.h>namespace {- constexpr int kBlockSize = 512;- constexpr int kMaxBlocksPerSM = 32;+ __device__ __forceinline__ ulonglong4 add_ulonglong4(ulonglong4 va, ulonglong4 vb) {+ ulonglong4 vc;+ __half2* ha = reinterpret_cast<__half2*>(&va);+ __half2* hb = reinterpret_cast<__half2*>(&vb);+ __half2* hc = reinterpret_cast<__half2*>(&vc);+ #pragma unroll+ for (int i = 0; i < 8; ++i) {+ hc[i] = __hadd2(ha[i], hb[i]);+ }+ return vc;+ }- __global__ __launch_bounds__(kBlockSize, 2)- void vector_add_half2_kernel(const __half2* __restrict__ A,- const __half2* __restrict__ B,- __half2* __restrict__ C,- int64_t pair_count,- bool store_tail,- const __half* __restrict__ A_scalar,- const __half* __restrict__ B_scalar,- __half* __restrict__ C_scalar,- int64_t numel) {- const int64_t stride = static_cast<int64_t>(blockDim.x) * gridDim.x;- int64_t idx = blockIdx.x * blockDim.x + threadIdx.x;+ __global__ void vecadd_fp16_kernel(+ const half* __restrict__ a,+ const half* __restrict__ b,+ half* __restrict__ out,+ size_t elements+ ) {+ constexpr int pack_elems = 16; // 16 fp16 values per ulonglong4+ const size_t tid = threadIdx.x + blockIdx.x * blockDim.x;+ const size_t stride = static_cast<size_t>(blockDim.x) * gridDim.x;- while (idx < pair_count) {- __half2 aval = __ldg(A + idx);- __half2 bval = __ldg(B + idx);- C[idx] = __hadd2(aval, bval);- idx += stride;- }+ const size_t total_packs = elements / pack_elems;+ auto* a_vec = reinterpret_cast<const ulonglong4*>(a);+ auto* b_vec = reinterpret_cast<const ulonglong4*>(b);+ auto* o_vec = reinterpret_cast<ulonglong4*>(out);- if (store_tail && blockIdx.x == 0 && threadIdx.x == 0) {- const int64_t tail_index = numel - 1;- C_scalar[tail_index] = __hadd(A_scalar[tail_index], B_scalar[tail_index]);+ const size_t stride_packs = stride;+ auto* a_ptr = a_vec + tid;+ auto* b_ptr = b_vec + tid;+ auto* o_ptr = o_vec + tid;+ for (size_t pack = tid; pack < total_packs; pack += stride_packs * 2) {+ if (pack < total_packs) {+ ulonglong4 va0 = *a_ptr;+ ulonglong4 vb0 = *b_ptr;+ *o_ptr = add_ulonglong4(va0, vb0);+ }+ a_ptr += stride_packs;+ b_ptr += stride_packs;+ o_ptr += stride_packs;++ const size_t pack1 = pack + stride_packs;+ if (pack1 < total_packs) {+ ulonglong4 va1 = *a_ptr;+ ulonglong4 vb1 = *b_ptr;+ *o_ptr = add_ulonglong4(va1, vb1);+ }+ a_ptr += stride_packs;+ b_ptr += stride_packs;+ o_ptr += stride_packs;}- }- __global__ void vector_add_scalar_kernel(const __half* __restrict__ A,- const __half* __restrict__ B,- __half* __restrict__ C,- int64_t numel) {- const int64_t stride = static_cast<int64_t>(blockDim.x) * gridDim.x;- int64_t idx = blockIdx.x * blockDim.x + threadIdx.x;- while (idx < numel) {- C[idx] = __hadd(A[idx], B[idx]);- idx += stride;+ const size_t tail_start = total_packs * pack_elems;+ for (size_t idx = tail_start + tid; idx < elements; idx += stride) {+ out[idx] = __hadd(a[idx], b[idx]);}}- inline bool is_aligned_for_half2(const void* ptr) {- constexpr std::uintptr_t alignment = alignof(__half2);- return (reinterpret_cast<std::uintptr_t>(ptr) & (alignment - 1)) == 0;+ } // namespace++ void vecadd_fp16(torch::Tensor a, torch::Tensor b, torch::Tensor out) {+ TORCH_CHECK(a.is_cuda(), "Input tensor A must be CUDA");+ TORCH_CHECK(b.is_cuda(), "Input tensor B must be CUDA");+ TORCH_CHECK(out.is_cuda(), "Output tensor must be CUDA");+ TORCH_CHECK(a.dtype() == torch::kFloat16, "A must be float16");+ TORCH_CHECK(b.dtype() == torch::kFloat16, "B must be float16");+ TORCH_CHECK(out.dtype() == torch::kFloat16, "output must be float16");+ TORCH_CHECK(a.numel() == b.numel(), "Input sizes must match");+ TORCH_CHECK(a.numel() == out.numel(), "Output size mismatch");++ auto stream = at::cuda::getCurrentCUDAStream();+ const auto elements = static_cast<size_t>(out.numel());++ constexpr int block = 256;+ const int sms = at::cuda::getCurrentDeviceProperties()->multiProcessorCount;+ constexpr int pack_elems = 16;+ const size_t total_packs = elements / pack_elems;+ int grid = static_cast<int>((total_packs + block - 1) / block);+ grid = std::max(1, grid);++ vecadd_fp16_kernel<<<grid, block, 0, stream>>>(+ reinterpret_cast<const half*>(a.data_ptr<at::Half>()),+ reinterpret_cast<const half*>(b.data_ptr<at::Half>()),+ reinterpret_cast<half*>(out.data_ptr<at::Half>()),+ elements+ );+ AT_CUDA_CHECK(cudaGetLastError());}+ """- } // namespace+ _VECADD_EXT = load_inline(+ name="vectoradd_fp16_ext",+ cpp_sources="#include <torch/extension.h>\nvoid vecadd_fp16(torch::Tensor a, torch::Tensor b, torch::Tensor out);\n",+ cuda_sources=cuda_src,+ functions=["vecadd_fp16"],+ with_cuda=True,+ extra_cuda_cflags=[+ "-O3",+ "--use_fast_math",+ "-U__CUDA_NO_HALF_OPERATORS__",+ "-U__CUDA_NO_HALF_CONVERSIONS__",+ ],+ )+ return _VECADD_EXT- void launch_vector_add(torch::Tensor A, torch::Tensor B, torch::Tensor C) {- TORCH_CHECK(A.is_cuda(), "A must be a CUDA tensor");- TORCH_CHECK(B.is_cuda(), "B must be a CUDA tensor");- TORCH_CHECK(C.is_cuda(), "C must be a CUDA tensor");- TORCH_CHECK(A.scalar_type() == at::kHalf, "Expected torch.float16 tensors");- TORCH_CHECK(B.scalar_type() == at::kHalf, "Expected torch.float16 tensors");- TORCH_CHECK(C.scalar_type() == at::kHalf, "Expected torch.float16 tensors");+ def _is_aligned(tensor: torch.Tensor, alignment: int = 16) -> bool:+ return tensor.data_ptr() % alignment == 0- TORCH_CHECK(A.is_contiguous(), "A must be contiguous");- TORCH_CHECK(B.is_contiguous(), "B must be contiguous");- TORCH_CHECK(C.is_contiguous(), "C must be contiguous");- TORCH_CHECK(A.sizes() == B.sizes() && A.sizes() == C.sizes(),- "Input and output tensors must have identical shapes");+ @cute.kernel+ def _tile_elementwise_add_kernel(+ gA: cute.Tensor,+ gB: cute.Tensor,+ gC: cute.Tensor,+ tv_layout,+ ) -> None:+ """Block-tiled elementwise add with per-thread vectorization."""+ tidx, _, _ = cute.arch.thread_idx()+ bidx, _, _ = cute.arch.block_idx()+ grid_x, _, _ = cute.arch.grid_dim()+ num_tiles = cute.size(gC, mode=[1])- c10::cuda::CUDAGuard device_guard(A.device());- auto stream = at::cuda::getCurrentCUDAStream();+ tile_idx = bidx+ while tile_idx < num_tiles:+ blk_coord = ((None, None), tile_idx)+ blkA = gA[blk_coord]+ blkB = gB[blk_coord]+ blkC = gC[blk_coord]- const int64_t numel = A.numel();- if (numel == 0) {- return;- }+ tidfrgA = cute.composition(blkA, tv_layout)+ tidfrgB = cute.composition(blkB, tv_layout)+ tidfrgC = cute.composition(blkC, tv_layout)- const __half* A_ptr = reinterpret_cast<const __half*>(A.data_ptr<at::Half>());- const __half* B_ptr = reinterpret_cast<const __half*>(B.data_ptr<at::Half>());- __half* C_ptr = reinterpret_cast<__half*>(C.data_ptr<at::Half>());+ thr_coord = (tidx, None)+ thrA = tidfrgA[thr_coord]+ thrB = tidfrgB[thr_coord]+ thrC = tidfrgC[thr_coord]- const bool use_half2 = numel > 1 &&- is_aligned_for_half2(A_ptr) &&- is_aligned_for_half2(B_ptr) &&- is_aligned_for_half2(C_ptr);+ thrC[None] = thrA.load() + thrB.load()- auto* props = at::cuda::getCurrentDeviceProperties();- const int max_blocks = props->multiProcessorCount * kMaxBlocksPerSM;+ tile_idx += grid_x- if (use_half2) {- const int64_t pair_count = numel >> 1;- const bool has_tail = (numel & 1) != 0;- int grid = static_cast<int>((pair_count + kBlockSize - 1) / kBlockSize);- if (grid < 1) {- grid = 1;- }- if (grid > max_blocks) {- grid = max_blocks;- }- vector_add_half2_kernel<<<grid, kBlockSize, 0, stream>>>(- reinterpret_cast<const __half2*>(A_ptr),- reinterpret_cast<const __half2*>(B_ptr),- reinterpret_cast<__half2*>(C_ptr),- pair_count,- has_tail,- A_ptr,- B_ptr,- C_ptr,- numel);- } else {- int grid = static_cast<int>((numel + kBlockSize - 1) / kBlockSize);- if (grid < 1) {- grid = 1;- }- if (grid > max_blocks) {- grid = max_blocks;- }- vector_add_scalar_kernel<<<grid, kBlockSize, 0, stream>>>(- A_ptr, B_ptr, C_ptr, numel);- }- C10_CUDA_KERNEL_LAUNCH_CHECK();- }- """+ @cute.jit+ def _tiled_elementwise_add(+ mA: cute.Tensor,+ mB: cute.Tensor,+ mC: cute.Tensor,+ ) -> None:+ """Launch tiled kernel honoring remainder shapes."""+ element_type = mA.element_type+ bytes_per_element = element_type.width // 8+ coalesced_ldst_bytes = bytes_per_element- _kernel_cache: Dict[Tuple[int, int], Any] = {}+ thr_layout = cute.make_ordered_layout((4, 64), order=(1, 0))+ val_layout = cute.make_ordered_layout((16, coalesced_ldst_bytes), order=(1, 0))+ val_layout = cute.recast_layout(element_type.width, 8, val_layout)+ tiler_mn, tv_layout = cute.make_layout_tv(thr_layout, val_layout)+ gA = cute.zipped_divide(mA, tiler_mn)+ gB = cute.zipped_divide(mB, tiler_mn)+ gC = cute.zipped_divide(mC, tiler_mn)- def _load_kernel(device: torch.device):- capability = torch.cuda.get_device_capability(device=device)- key = capability- if key in _kernel_cache:- return _kernel_cache[key]+ remap_block = cute.make_ordered_layout(+ cute.select(gA.shape[1], mode=[1, 0]), order=(1, 0)+ )+ gA = cute.composition(gA, (None, remap_block))+ gB = cute.composition(gB, (None, remap_block))+ gC = cute.composition(gC, (None, remap_block))- arch_flag = f"-gencode=arch=compute_{capability[0]}{capability[1]},code=sm_{capability[0]}{capability[1]}"- module = load_inline(- name=f"vector_add_cuda_{capability[0]}{capability[1]}",- cpp_sources=CPP_DECL,- cuda_sources=CUDA_SRC,- functions=["launch_vector_add"],- extra_cuda_cflags=[- "-O3",- "-std=c++17",- "-lineinfo",- "-use_fast_math",- arch_flag,- ],+ total_tiles = int(cute.size(gC, mode=[1]))+ block_dim_x = int(cute.size(tv_layout, mode=[0]))+ max_ctas = 128+ if torch.cuda.is_available():+ props = torch.cuda.get_device_properties(torch.cuda.current_device())+ max_ctas = max(1, props.multi_processor_count * 2)+ grid_dim_x = min(total_tiles, max_ctas) if total_tiles > 0 else 1++ _tile_elementwise_add_kernel(gA, gB, gC, tv_layout).launch(+ grid=(grid_dim_x, 1, 1),+ block=(block_dim_x, 1, 1),)- _kernel_cache[key] = module- return module- def _validate_tensors(A: torch.Tensor, B: torch.Tensor, C: torch.Tensor) -> None:- if not (A.is_cuda and B.is_cuda and C.is_cuda):- raise RuntimeError("All tensors must reside on CUDA devices.")- if A.dtype != torch.float16 or B.dtype != torch.float16 or C.dtype != torch.float16:- raise RuntimeError("All tensors must use torch.float16 dtype.")- if A.shape != B.shape or A.shape != C.shape:- raise RuntimeError("Input and output tensors must share identical shapes.")- if not (A.is_contiguous() and B.is_contiguous() and C.is_contiguous()):- raise RuntimeError("Input and output tensors must be contiguous.")+ def _ensure_contiguous(tensor: torch.Tensor) -> torch.Tensor:+ if tensor.is_contiguous():+ return tensor+ return tensor.contiguous()+ def _torch_to_cute(tensor: torch.Tensor) -> cute.Tensor:+ return from_dlpack(tensor, assumed_align=16)+++ def _get_compiled_kernel(+ key: Tuple[Tuple[int, ...], torch.dtype],+ a_: cute.Tensor,+ b_: cute.Tensor,+ c_: cute.Tensor,+ ):+ compiled = _KERNEL_CACHE.get(key)+ if compiled is None:+ compiled = cute.compile(_tiled_elementwise_add, a_, b_, c_)+ _KERNEL_CACHE[key] = compiled+ return compiled++def custom_kernel(data: input_t) -> output_t:- try:- A, B, C = data # type: ignore[misc]- except ValueError as err:- raise RuntimeError("Expected (A, B, output) tensor tuple.") from err+ """Run vector addition via a tiled Cutlass (CuTe) kernel."""+ A, B, output = data- _validate_tensors(A, B, C)+ A_contig = _ensure_contiguous(A)+ B_contig = _ensure_contiguous(B)+ output_contig = output if output.is_contiguous() else output.contiguous()- numel = A.numel()- if numel >= (1 << 20):- torch.add(A, B, out=C)- return C+ a_cute = _torch_to_cute(A_contig)+ b_cute = _torch_to_cute(B_contig)+ c_cute = _torch_to_cute(output_contig)- kernel = _load_kernel(A.device)- kernel.launch_vector_add(A, B, C)- return C+ if (+ A_contig.is_cuda+ and A_contig.dtype == torch.float16+ and _is_aligned(A_contig, 32)+ and _is_aligned(B_contig, 32)+ and _is_aligned(output_contig, 32)+ ):+ ext = _get_vecadd_extension()+ ext.vecadd_fp16(A_contig, B_contig, output_contig)+ else:+ kernel_key = (tuple(A_contig.shape), A_contig.dtype)+ compiled_kernel = _get_compiled_kernel(+ kernel_key, a_cute, b_cute, c_cute+ )+ compiled_kernel(a_cute, b_cute, c_cute)++ if output_contig is not output:+ output.copy_(output_contig)++ return output
scrolls · 417 diff lines total
Best evidence level for this revision: reported
JSON