Skip to content
KernelIndex
Search⌘K

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
FP16 vector additionsuite of 5 cases
NVIDIA B200
235.5µs
#19 of 66
2025-11-08

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