Skip to content
KernelIndex
Search⌘K

submission 333230

baicai1145 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 536 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-333230?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 dual GEMMsuite of 4 cases
NVIDIA B200
32.4µs
#261 of 420
2026-01-11

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:fa818177fe1f2a67089ae36f4189eecf3eeb680b93cbc3d1d54aff983a827757
license declaredunknown
license concludedunknown
authorsbaicai1145
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

clusterusing ClusterShape = Shape<_1, _1, _1>;
fp4using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
fused-epilogueusing CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<

Kernel source

submission.py536 lines
#!POPCORN leaderboard nvfp4_dual_gemm

from __future__ import annotations

import hashlib
import os
import textwrap
import threading
from dataclasses import dataclass

import torch

from task import input_t, output_t

# -----------------------------------------------------------------------------
# NOTE (anti-cheat):
# - 不要在代码中显式使用/传递 CUDA 队列句柄(可能触发反作弊静态检测)。
# - Python 侧不做任何会触发 CUDA kernel 的操作(如 contiguous / silu / matmul 等)。
# - 扩展内 kernel/CUTLASS 都使用 CUDA 默认队列(不显式指定)。
# -----------------------------------------------------------------------------

_EXT_LOCK = threading.Lock()
_EXT_MOD = None
_EXT_ERR: str | None = None


def _summarize_ext_error(err: BaseException) -> str:
    s = f"{type(err).__name__}: {err}"
    lines = [ln.strip() for ln in s.splitlines() if ln.strip()]
    key = []
    for ln in lines:
        low = ln.lower()
        if ("error:" in low) or ("fatal error" in low) or ("undefined reference" in low) or ("static assertion" in low):
            key.append(ln)
    return "\n".join(key[-40:]) if key else "\n".join(lines[-40:]) if lines else repr(err)


def _cutlass_include_paths() -> list[str]:
    # 评测机上 CUTLASS 安装在 /opt/cutlass/4.3.0/include(日志里可见)。
    candidates = [
        "/opt/cutlass/4.3.0/include",
        "/opt/cutlass/include",
    ]
    return [p for p in candidates if os.path.isdir(p)]


def _load_ext():
    global _EXT_MOD, _EXT_ERR
    if _EXT_MOD is not None or _EXT_ERR is not None:
        return _EXT_MOD

    with _EXT_LOCK:
        if _EXT_MOD is not None or _EXT_ERR is not None:
            return _EXT_MOD

        try:
            from torch.utils.cpp_extension import load_inline

            # B200 = sm_100a
            os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "10.0a")

            cuda_src = _CUDA_SRC
            name_src = _CPP_SRC + "\n/* --- */\n" + cuda_src
            name = "nvfp4_dual_gemm_ext_" + hashlib.sha256(name_src.encode("utf-8")).hexdigest()[:16]

            extra_cuda_cflags = [
                "-O3",
                "-std=c++17",
                "--expt-relaxed-constexpr",
                "--expt-extended-lambda",
            ]

            verbose = os.environ.get("NVFP4_EXT_VERBOSE", "0").strip().lower() in {"1", "true", "yes"}
            _EXT_MOD = load_inline(
                name=name,
                cpp_sources=_CPP_SRC,
                cuda_sources=cuda_src,
                functions=["pack_scale", "workspace_size_bytes", "run"],
                extra_cuda_cflags=extra_cuda_cflags,
                extra_include_paths=_cutlass_include_paths(),
                with_cuda=True,
                verbose=verbose,
            )
            return _EXT_MOD
        except Exception as e:
            _EXT_ERR = _summarize_ext_error(e)
            return None


def _require_ext(ext) -> None:
    if ext is not None:
        return
    err = _EXT_ERR
    if err is None:
        raise RuntimeError("extension not available")
    raise RuntimeError(f"extension not available (build/load failed):\n{err}")


@dataclass(frozen=True)
class _WorkspaceKey:
    device: str
    m: int
    n: int
    k: int
    l: int


_WORKSPACE: dict[_WorkspaceKey, torch.Tensor] = {}


def _get_workspace(device: torch.device, m: int, n: int, k: int, l: int) -> torch.Tensor:
    key = _WorkspaceKey(str(device), m, n, k, l)
    ws = _WORKSPACE.get(key)
    if ws is not None:
        return ws

    ext = _load_ext()
    _require_ext(ext)
    bytes_needed = int(ext.workspace_size_bytes(int(m), int(n), int(k), int(l)))
    ws = torch.empty((bytes_needed,), device=device, dtype=torch.uint8)
    _WORKSPACE[key] = ws
    return ws


def _sf_packed_view(sf_permuted: torch.Tensor) -> torch.Tensor:
    """
    利用 generate_input() 的构造方式:sf_permuted 实际来自一个 contiguous 基张量的 permute view。
    这里把维度 permute 回去得到 [L, rest_mn, rest_k, 32, 4, 4],然后 reshape 成 [L, -1]。

    这一步只做 view/stride 变换,不触发任何 CUDA kernel,也不会产生额外 CUDA 队列工作。
    """
    if sf_permuted is None or getattr(sf_permuted, "ndim", 0) != 6:
        raise RuntimeError("permuted scale factors must be 6D")
    # [32, 4, rest_mn, 4, rest_k, L] -> [L, rest_mn, rest_k, 32, 4, 4]
    base = sf_permuted.permute(5, 2, 4, 0, 1, 3)

    # 绝大多数情况下(评测 benchmark 的原始 data),base 是 contiguous 的:直接 view 成 [L, -1],零拷贝。
    if base.is_contiguous():
        return base.view(base.shape[0], -1)

    # correctness 阶段会 clone() 破坏上述“逆 permute 复原 contiguous”的性质;此时避免 reshape 触发隐式 copy,
    # 直接让扩展执行显式打包(不依赖当前 CUDA 队列)。
    ext = _load_ext()
    _require_ext(ext)
    return ext.pack_scale(sf_permuted)


@torch.no_grad()
def custom_kernel(data: input_t) -> output_t:
    a, b1, b2, _sfa, _sfb1, _sfb2, sfa_p, sfb1_p, sfb2_p, c = data

    ext = _load_ext()
    _require_ext(ext)

    if sfa_p is None or sfb1_p is None or sfb2_p is None:
        raise RuntimeError("permuted scale factors not available")
    if not (a.is_cuda and b1.is_cuda and b2.is_cuda and c.is_cuda):
        raise RuntimeError("CUDA tensors required")

    m, n, l = c.shape
    # a/b 是 torch.float4_e2m1fn_x2:K 维在 tensor.shape 中是 K/2(每元素包含 2 个 fp4)。
    k = int(a.shape[1]) * 2

    ws = _get_workspace(a.device, m, n, k, l)
    sfa_packed = _sf_packed_view(sfa_p)
    sfb1_packed = _sf_packed_view(sfb1_p)
    sfb2_packed = _sf_packed_view(sfb2_p)

    ext.run(a, b1, b2, sfa_packed, sfb1_packed, sfb2_packed, c, ws)
    return c


_CUDA_SRC = textwrap.dedent(
    r"""
    #include <torch/extension.h>
    #include <cuda_runtime.h>

    #include "cutlass/cutlass.h"
    #include "cute/tensor.hpp"
    #include "cutlass/gemm/dispatch_policy.hpp"
    #include "cutlass/gemm/collective/collective_builder.hpp"
    #include "cutlass/epilogue/collective/collective_builder.hpp"
    #include "cutlass/detail/sm100_blockscaled_layout.hpp"
    #include "cutlass/gemm/device/gemm_universal_adapter.h"
    #include "cutlass/gemm/kernel/gemm_universal.hpp"
    #include "cutlass/util/packed_stride.hpp"

    namespace nvfp4_dual_gemm_ext {
    using namespace cute;

    // -----------------------------
    // CUTLASS GEMM config (NVFP4 block-scaled -> FP32)
    // -----------------------------
    using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
    using ElementB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
    using ElementAccumulator = float;

    using ArchTag = cutlass::arch::Sm100;
    using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;

    // Task data layout: A is [M,K] K-major => RowMajor.
    // For NVFP4 kernels, CUTLASS requires B to be ColumnMajor.
    using LayoutATag = cutlass::layout::RowMajor;
    using LayoutBTag = cutlass::layout::ColumnMajor;
    constexpr int AlignmentA = 32;
    constexpr int AlignmentB = 32;

    // No C, write D as FP32
    using ElementC = void;
    using LayoutCTag = cutlass::layout::RowMajor;
    constexpr int AlignmentC = 1;

    using ElementD = float;
    using LayoutDTag = cutlass::layout::RowMajor;
    constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value; // 4

    // Tile: keep SMEM low; start from the known-good 1SM config used in CUTLASS examples.
    using MmaTileShape = Shape<_128, _128, _256>;
    using ClusterShape = Shape<_1, _1, _1>;

    using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
        ArchTag, OperatorClass,
        MmaTileShape, ClusterShape,
        cutlass::epilogue::collective::EpilogueTileAuto,
        ElementAccumulator, ElementAccumulator,
        ElementC, LayoutCTag, AlignmentC,
        ElementD, LayoutDTag, AlignmentD,
        cutlass::epilogue::collective::EpilogueScheduleAuto
      >::CollectiveOp;

    using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
        ArchTag, OperatorClass,
        ElementA, LayoutATag, AlignmentA,
        ElementB, LayoutBTag, AlignmentB,
        ElementAccumulator,
        MmaTileShape, ClusterShape,
        cutlass::gemm::collective::StageCountAutoCarveout<
          static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
        cutlass::gemm::collective::KernelScheduleAuto
      >::CollectiveOp;

    using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
        Shape<int,int,int,int>, // (M,N,K,L)
        CollectiveMainloop,
        CollectiveEpilogue,
        void>;

    using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;

    // -----------------------------
    // scale packing (fallback path):
    // sf_permuted [32,4,rest_mn,4,rest_k,L] (possibly contiguous) ->
    // packed [L, rest_mn*rest_k*32*4*4] in the order that CUTLASS expects.
    // -----------------------------
    __global__ void pack_scale_kernel(
        const uint8_t* __restrict__ in,
        uint8_t* __restrict__ out,
        int64_t s0, int64_t s1, int64_t s2, int64_t s3, int64_t s4, int64_t s5,
        int rest_mn, int rest_k, int L) {
      int64_t idx = (int64_t)blockIdx.x * blockDim.x + threadIdx.x;
      int64_t per_l = (int64_t)rest_mn * rest_k * 32 * 4 * 4;
      int64_t total = per_l * (int64_t)L;
      if (idx >= total) return;

      int l = (int)(idx / per_l);
      int64_t t = idx - (int64_t)l * per_l;

      int kk4 = (int)(t & 3); t >>= 2;          // fastest
      int mm4 = (int)(t & 3); t >>= 2;
      int mm32 = (int)(t & 31); t >>= 5;
      int rk = (int)(t % rest_k); t /= rest_k;
      int rm = (int)t;

      int64_t in_off = (int64_t)mm32 * s0 + (int64_t)mm4 * s1 + (int64_t)rm * s2 +
                       (int64_t)kk4 * s3 + (int64_t)rk * s4 + (int64_t)l * s5;
      out[idx] = in[in_off];
    }

    __global__ void gated_silu_mul_to_fp16(
        const float* __restrict__ x,
        const float* __restrict__ y,
        at::Half* __restrict__ out,
        int64_t total) {
      int64_t idx = (int64_t)blockIdx.x * blockDim.x + threadIdx.x;
      if (idx >= total) return;
      float a = x[idx];
      float b = y[idx];
      // SiLU(x) = x * sigmoid(x)
      float sig = 1.0f / (1.0f + __expf(-a));
      float v = (a * sig) * b;
      out[idx] = (at::Half)v;
    }

    inline void check_cuda(cudaError_t e, const char* what) {
      if (e == cudaSuccess) return;
      throw std::runtime_error(std::string(what) + ": " + cudaGetErrorString(e));
    }

    } // namespace nvfp4_dual_gemm_ext

    torch::Tensor pack_scale(torch::Tensor sf_permuted) {
      using namespace nvfp4_dual_gemm_ext;
      if (!sf_permuted.is_cuda()) {
        throw std::runtime_error("pack_scale: CUDA tensor required");
      }
      if (sf_permuted.dim() != 6) {
        throw std::runtime_error("pack_scale: expect 6D tensor [32,4,rest_mn,4,rest_k,L]");
      }
      if (sf_permuted.element_size() != 1) {
        throw std::runtime_error("pack_scale: expect 1-byte dtype (fp8)");
      }

      auto sizes = sf_permuted.sizes();
      if (sizes[0] != 32 || sizes[1] != 4 || sizes[3] != 4) {
        throw std::runtime_error("pack_scale: unexpected shape; expected [32,4,rest_mn,4,rest_k,L]");
      }
      int rest_mn = (int)sizes[2];
      int rest_k = (int)sizes[4];
      int L = (int)sizes[5];

      int64_t per_l = (int64_t)rest_mn * rest_k * 32 * 4 * 4;
      auto out = torch::empty({L, per_l}, torch::TensorOptions().device(sf_permuted.device()).dtype(at::kByte));

      auto st = sf_permuted.strides();
      int64_t s0 = st[0], s1 = st[1], s2 = st[2], s3 = st[3], s4 = st[4], s5 = st[5];

      const uint8_t* in = reinterpret_cast<const uint8_t*>(sf_permuted.data_ptr());
      uint8_t* outp = out.data_ptr<uint8_t>();

      int threads = 256;
      int64_t total = per_l * (int64_t)L;
      int blocks = (int)((total + threads - 1) / threads);

      pack_scale_kernel<<<blocks, threads>>>(in, outp, s0, s1, s2, s3, s4, s5, rest_mn, rest_k, L);
      check_cuda(cudaGetLastError(), "pack_scale_kernel launch");
      return out;
    }

    int64_t workspace_size_bytes(int64_t m, int64_t n, int64_t k, int64_t l) {
      using namespace nvfp4_dual_gemm_ext;
      // tmp1/tmp2: fp32 outputs of two GEMMs
      int64_t mn = m * n * l;
      int64_t tmp_bytes = 2 * mn * (int64_t)sizeof(float);
      (void)l;

      using StrideA = typename Gemm::GemmKernel::StrideA;
      using StrideB = typename Gemm::GemmKernel::StrideB;
      using StrideC = typename Gemm::GemmKernel::StrideC;
      using StrideD = typename Gemm::GemmKernel::StrideD;
      auto stride_A = cutlass::make_cute_packed_stride(StrideA{}, {(int)m, (int)k, 1});
      auto stride_B = cutlass::make_cute_packed_stride(StrideB{}, {(int)n, (int)k, 1});
      auto stride_C = cutlass::make_cute_packed_stride(StrideC{}, {(int)m, (int)n, 1});
      auto stride_D = cutlass::make_cute_packed_stride(StrideD{}, {(int)m, (int)n, 1});

      using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
      auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape((int)m, (int)n, (int)k, 1));
      auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape((int)m, (int)n, (int)k, 1));

      typename Gemm::Arguments args{
        cutlass::gemm::GemmUniversalMode::kGemm,
        {(int)m, (int)n, (int)k, 1},
        {nullptr, stride_A, nullptr, stride_B, nullptr, layout_SFA, nullptr, layout_SFB},
        {{}, nullptr, stride_C, nullptr, stride_D}
      };

      size_t cutlass_ws = Gemm::get_workspace_size(args);
      // small alignment slack
      return tmp_bytes + (int64_t)cutlass_ws + 256;
    }

    void run(
        torch::Tensor a,
        torch::Tensor b1,
        torch::Tensor b2,
        torch::Tensor sfa_packed,
        torch::Tensor sfb1_packed,
        torch::Tensor sfb2_packed,
        torch::Tensor out_fp16,
        torch::Tensor workspace_u8) {
      using namespace nvfp4_dual_gemm_ext;
      if (!a.is_cuda() || !b1.is_cuda() || !b2.is_cuda() || !out_fp16.is_cuda()) {
        throw std::runtime_error("run: CUDA tensors required");
      }
      if (out_fp16.scalar_type() != at::kHalf) {
        throw std::runtime_error("run: out must be fp16");
      }
      if (workspace_u8.scalar_type() != at::kByte) {
        throw std::runtime_error("run: workspace must be uint8");
      }
      if (a.dim() != 3 || b1.dim() != 3 || b2.dim() != 3 || out_fp16.dim() != 3) {
        throw std::runtime_error("run: expect a/b/out are 3D");
      }

      int m = (int)out_fp16.size(0);
      int n = (int)out_fp16.size(1);
      int L = (int)out_fp16.size(2);

      // a/b are torch.float4_e2m1fn_x2: their K dim is K/2 in tensor shape.
      int k2 = (int)a.size(1);
      int k = k2 * 2;

      if ((int)b1.size(0) != n || (int)b2.size(0) != n || (int)b1.size(1) != k2 || (int)b2.size(1) != k2) {
        throw std::runtime_error("run: b shape mismatch");
      }
      if ((int)a.size(0) != m || (int)b1.size(2) != L || (int)b2.size(2) != L || (int)a.size(2) != L) {
        throw std::runtime_error("run: batch (L) mismatch");
      }

      // packed SF is [L, m*n?] with dtype uint8
      if (!sfa_packed.is_cuda() || !sfb1_packed.is_cuda() || !sfb2_packed.is_cuda()) {
        throw std::runtime_error("run: packed scale factors must be CUDA tensors");
      }
      if (sfa_packed.dim() != 2 || sfb1_packed.dim() != 2 || sfb2_packed.dim() != 2) {
        throw std::runtime_error("run: packed scale factors must be 2D [L, ...]");
      }
      if ((int)sfa_packed.size(0) != L || (int)sfb1_packed.size(0) != L || (int)sfb2_packed.size(0) != L) {
        throw std::runtime_error("run: packed scale factors L mismatch");
      }
      if (sfa_packed.element_size() != 1 || sfb1_packed.element_size() != 1 || sfb2_packed.element_size() != 1) {
        throw std::runtime_error("run: packed scale factors must be 1-byte dtype");
      }
      if (sfa_packed.stride(1) != 1 || sfb1_packed.stride(1) != 1 || sfb2_packed.stride(1) != 1) {
        throw std::runtime_error("run: packed scale factors must be contiguous on dim=1");
      }

      // Workspace layout:
      // [tmp1 fp32 | tmp2 fp32 | cutlass workspace]
      int64_t mn_total = (int64_t)m * n * L;
      int64_t tmp_bytes = 2 * mn_total * (int64_t)sizeof(float);
      uint8_t* ws = workspace_u8.data_ptr<uint8_t>();
      float* tmp1 = reinterpret_cast<float*>(ws);
      float* tmp2 = reinterpret_cast<float*>(ws + mn_total * (int64_t)sizeof(float));
      void* cutlass_ws = (void*)(ws + tmp_bytes);

      // Strides (CUTLASS expects logical K, not K/2)
      using StrideA = typename Gemm::GemmKernel::StrideA;
      using StrideB = typename Gemm::GemmKernel::StrideB;
      using StrideC = typename Gemm::GemmKernel::StrideC;
      using StrideD = typename Gemm::GemmKernel::StrideD;
      auto stride_A = cutlass::make_cute_packed_stride(StrideA{}, {m, k, 1});
      auto stride_B = cutlass::make_cute_packed_stride(StrideB{}, {n, k, 1});
      auto stride_C = cutlass::make_cute_packed_stride(StrideC{}, {m, n, 1});
      auto stride_D = cutlass::make_cute_packed_stride(StrideD{}, {m, n, 1});

      using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
      auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(m, n, k, 1));
      auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(m, n, k, 1));

      // Base pointers + batch strides (in bytes, because torch uses 1-byte storage for fp4x2/fp8/uint8)
      const uint8_t* a_base = reinterpret_cast<const uint8_t*>(a.data_ptr());
      const uint8_t* b1_base = reinterpret_cast<const uint8_t*>(b1.data_ptr());
      const uint8_t* b2_base = reinterpret_cast<const uint8_t*>(b2.data_ptr());
      const int64_t a_l_stride_bytes = (int64_t)a.stride(2) * a.element_size();
      const int64_t b_l_stride_bytes = (int64_t)b1.stride(2) * b1.element_size();
      const int64_t out_l_stride_elems = out_fp16.stride(2); // fp16 elems

      const uint8_t* sfa_base = reinterpret_cast<const uint8_t*>(sfa_packed.data_ptr());
      const uint8_t* sfb1_base = reinterpret_cast<const uint8_t*>(sfb1_packed.data_ptr());
      const uint8_t* sfb2_base = reinterpret_cast<const uint8_t*>(sfb2_packed.data_ptr());
      const int64_t sfa_l_stride = (int64_t)sfa_packed.stride(0) * sfa_packed.element_size();
      const int64_t sfb_l_stride = (int64_t)sfb1_packed.stride(0) * sfb1_packed.element_size();

      Gemm gemm;

      // Loop batches; L is small (bench/tests use L=1)
      for (int l_idx = 0; l_idx < L; ++l_idx) {
        auto A = reinterpret_cast<const ElementA::DataType*>(a_base + (int64_t)l_idx * a_l_stride_bytes);
        auto B1 = reinterpret_cast<const ElementB::DataType*>(b1_base + (int64_t)l_idx * b_l_stride_bytes);
        auto B2 = reinterpret_cast<const ElementB::DataType*>(b2_base + (int64_t)l_idx * b_l_stride_bytes);

        auto SFA = reinterpret_cast<const ElementA::ScaleFactorType*>(sfa_base + (int64_t)l_idx * sfa_l_stride);
        auto SFB1 = reinterpret_cast<const ElementB::ScaleFactorType*>(sfb1_base + (int64_t)l_idx * sfb_l_stride);
        auto SFB2 = reinterpret_cast<const ElementB::ScaleFactorType*>(sfb2_base + (int64_t)l_idx * sfb_l_stride);

        float* D1 = tmp1 + (int64_t)l_idx * (int64_t)m * n;
        float* D2 = tmp2 + (int64_t)l_idx * (int64_t)m * n;
        at::Half* out = out_fp16.data_ptr<at::Half>() + (int64_t)l_idx * out_l_stride_elems;

        typename Gemm::Arguments args1{
          cutlass::gemm::GemmUniversalMode::kGemm,
          {m, n, k, 1},
          {A, stride_A, B1, stride_B, SFA, layout_SFA, SFB1, layout_SFB},
          {{}, nullptr, stride_C, D1, stride_D}
        };

        typename Gemm::Arguments args2{
          cutlass::gemm::GemmUniversalMode::kGemm,
          {m, n, k, 1},
          {A, stride_A, B2, stride_B, SFA, layout_SFA, SFB2, layout_SFB},
          {{}, nullptr, stride_C, D2, stride_D}
        };

        size_t cutlass_ws_needed = Gemm::get_workspace_size(args1);
        if ((int64_t)workspace_u8.numel() < tmp_bytes + (int64_t)cutlass_ws_needed) {
          throw std::runtime_error("run: workspace too small for CUTLASS");
        }

        cutlass::Status st1 = gemm(args1, cutlass_ws);
        if (st1 != cutlass::Status::kSuccess) {
          throw std::runtime_error("CUTLASS GEMM1 failed");
        }
        cutlass::Status st2 = gemm(args2, cutlass_ws);
        if (st2 != cutlass::Status::kSuccess) {
          throw std::runtime_error("CUTLASS GEMM2 failed");
        }
        // Fused activation+mul
        int threads = 256;
        int64_t elems = (int64_t)m * n;
        int blocks = (int)((elems + threads - 1) / threads);
        gated_silu_mul_to_fp16<<<blocks, threads>>>(D1, D2, out, elems);
        check_cuda(cudaGetLastError(), "gated_silu_mul_to_fp16 launch");
      }
    }

    """
)

_CPP_SRC = textwrap.dedent(
    r"""
    #include <torch/extension.h>
    #include <cstdint>

    // Implemented in the CUDA translation unit
    torch::Tensor pack_scale(torch::Tensor sf_permuted);
    int64_t workspace_size_bytes(int64_t m, int64_t n, int64_t k, int64_t l);
    void run(torch::Tensor a,
             torch::Tensor b1,
             torch::Tensor b2,
             torch::Tensor sfa_packed,
             torch::Tensor sfb1_packed,
             torch::Tensor sfb2_packed,
             torch::Tensor out_fp16,
             torch::Tensor workspace_u8);
    """
)
scrolls · 536 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON