Skip to content
KernelIndex
Search⌘K

submission 297876

symlon · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

inline.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-297876?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
56.1µs
#302 of 420
2026-01-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:26c39caab57dbc604d05a22222417fc07feb52760b44859ddf17b26b8e082742
license declaredunknown
license concludedunknown
authorssymlon
imported2026-08-26

Techniques

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

fp4PyTorch reference implementation of NVFP4 block-scaled dual GEMM with silu activation,
fp8using sfA_gl = gl<fp8e4m3, -1, -1, PACKED_PER_TILE * 32, 16, sfA_tile>;
shared-memoryextern __shared__ int __shm[];
stages = 2static constexpr int PIPELINE_STAGES = 2;
warp-specializationstatic constexpr int CONSUMER_WARPGROUPS = 1;

Kernel source

inline.py574 lines
#!POPCORN leaderboard nvfp4_dual_gemm
#!POPCORN gpu NVIDIA
import os
import subprocess
import shutil
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

thunderkittens_dir = "/home/runner/ThunderKittens"

base_path = "/home/runner"
dirs = [d for d in os.listdir(base_path) if os.path.isdir(os.path.join(base_path, d))]

print("Existing directories in /home/runner:")
for d in dirs:
    print(d)

if os.path.exists(thunderkittens_dir):
    shutil.rmtree(thunderkittens_dir)
    print(f"Deleted existing folder {thunderkittens_dir}.")

if not os.path.exists(thunderkittens_dir):
    result = subprocess.run([
        "git", "clone", "--recursive",
        "-b", "tk-v2",
        "https://github.com/symlons/ThunderKittens.git",
        thunderkittens_dir
    ])
    if result.returncode == 0:
        print(f"ThunderKittens cloned successfully into {thunderkittens_dir}.")
    else:
        print("Failed to clone ThunderKittens.")
else:
    print(f"{thunderkittens_dir} already exists. Switching to tk-v2 branch.")
    subprocess.run(["git", "-C", thunderkittens_dir, "fetch", "origin"])
    subprocess.run(["git", "-C", thunderkittens_dir, "checkout", "tk-v2"])
    subprocess.run(["git", "-C", thunderkittens_dir, "pull"], check=True)
    subprocess.run(["git", "-C", thunderkittens_dir, "reset", "--hard", "origin/tk-v2"])
    subprocess.run(["git", "-C", thunderkittens_dir, "log", "-1", "--format=%H"])

nvpf4_src = r'''
#include "kittens.cuh"
#include "pyutils/torchutils.cuh"

// #include "ATen/ops/linalg_vector_norm.h"
// #include "ATen/ops/unsqueeze.h"
// todo: 2-cta pipeline
// todo: work-stealing
// todo: register allocation
// todo: warp index order?
// todo: tuning
// todo: ncu profile

using namespace kittens;

namespace nvfp4_gemm {
struct config {
  static constexpr int CLUSTER_SIZE = 1;

  static constexpr int NUM_BLOCKS = 148;
  static constexpr int STATIC_SHARED_MEMORY = 1024;
  static constexpr int DYNAMIC_SHARED_MEMORY = MAX_SHARED_MEMORY - STATIC_SHARED_MEMORY;

  static constexpr int CONSUMER_WARPGROUPS = 1;
  static constexpr int PRODUCER_WARPGROUPS = 1;
  static constexpr int NUM_WARPGROUPS = CONSUMER_WARPGROUPS + PRODUCER_WARPGROUPS;
  static constexpr int NUM_WARPS = NUM_WARPGROUPS * WARPGROUP_WARPS;
  static constexpr int NUM_THREADS = NUM_WARPS * WARP_THREADS;

  static constexpr int PRODUCER_REGISTERS = 256;
  static constexpr int CONSUMER_REGISTERS = 256;
};

struct globals {
  static constexpr int PIPELINE_STAGES = 2;
  static constexpr int TILE_SIZE = 128; // todo: up to 256
  static constexpr int PACKED_PER_TILE = 4;
  static constexpr int REDUCTION_BLOCK = PACKED_PER_TILE * 32;

  using A_fp4x2_tile = st_fp4e2m1_2<TILE_SIZE, TILE_SIZE>;
  using A_fp4x2_gl = gl<fp4e2m1_2, 1, 1, -1, -1, A_fp4x2_tile>;

  using B_fp4x2_tile = st_fp4e2m1_2<TILE_SIZE, TILE_SIZE>;
  using B_fp4x2_gl = gl<fp4e2m1_2, 1, 1, -1, -1, B_fp4x2_tile>;

  using C_tile = st_hf<TILE_SIZE, TILE_SIZE>;
  using C_gl = gl<half, 1, 1, -1, -1, C_tile>;

  A_fp4x2_gl A;
  B_fp4x2_gl B;
  B_fp4x2_gl B_2;
  C_gl C;

  using sfA_tile = st_fp8e4m3<PACKED_PER_TILE * 32, 16, false>;
  using sfA_gl = gl<fp8e4m3, -1, -1, PACKED_PER_TILE * 32, 16, sfA_tile>;

  using sfB_tile = st_fp8e4m3<PACKED_PER_TILE * 32, 16, false>;
  using sfB_gl = gl<fp8e4m3, -1, -1, PACKED_PER_TILE * 32, 16, sfB_tile>;

  sfA_gl sfA;
  sfB_gl sfB;
  sfB_gl sfB_2;

  // __host__ inline dim3 grid() const {
  //   return dim3(A.cols() / TILE_SIZE, A.rows() / TILE_SIZE);
  // }
  __host__ inline int dynamic_shared_memory() const { // todo: check this
    return 4 * TILE_SIZE * TILE_SIZE * sizeof(bf16) + 1024;
  }
};

struct pipeline_input_tiles {
  globals::A_fp4x2_tile A;
  globals::B_fp4x2_tile B;
  globals::B_fp4x2_tile B_2;
};

struct pipeline_input_scales {
  globals::sfA_tile sfA;
  globals::sfB_tile sfB;
  globals::sfB_tile sfB_2;
};

struct pipeline_outputs {
  globals::C_tile C;
};

__device__ inline void kernel(const globals &G) {
  extern __shared__ int __shm[];
  tma_swizzle_allocator sm_allocator((int *)&__shm[0]);
  pipeline_input_tiles(&input_tiles)[globals::PIPELINE_STAGES] = sm_allocator.allocate<pipeline_input_tiles, globals::PIPELINE_STAGES>();
  pipeline_input_scales(&input_scales)[globals::PIPELINE_STAGES] = sm_allocator.allocate<pipeline_input_scales, globals::PIPELINE_STAGES>();

  pipeline_outputs &output_tiles = sm_allocator.allocate<pipeline_outputs>();

  tensor_allocator<1, config::CLUSTER_SIZE> tm_allocator;
  auto out_tm = tm_allocator.allocate<full_tt_fl<128>>(0);
  auto out_tm_2 = tm_allocator.allocate<full_tt_fl<128>>(128);
  auto A_sc_tm = tm_allocator.allocate<full_tt_fp8e4m3<16 * globals::PACKED_PER_TILE * globals::PIPELINE_STAGES>>(256);
  auto B_sc_tm = tm_allocator.allocate<full_tt_fp8e4m3<16 * globals::PACKED_PER_TILE * globals::PIPELINE_STAGES>>(256 + 4 * globals::PACKED_PER_TILE * globals::PIPELINE_STAGES);
  auto B_sc_tm_2 = tm_allocator.allocate<full_tt_fp8e4m3<16 * globals::PACKED_PER_TILE * globals::PIPELINE_STAGES>>(256 + 8 * globals::PACKED_PER_TILE * globals::PIPELINE_STAGES);


  __shared__ semaphore inputs_smem_arrived[globals::PIPELINE_STAGES];
  __shared__ semaphore matmul_finished_1[globals::PIPELINE_STAGES];
  __shared__ semaphore matmul_finished_2[globals::PIPELINE_STAGES];
  __shared__ semaphore inputs_all_arrived[globals::PIPELINE_STAGES];
  __shared__ semaphore outputs_arrived;
  __shared__ semaphore tensor_finished;

  if (threadIdx.x == 32) {
    #pragma unroll
    for (int i = 0; i < globals::PIPELINE_STAGES; ++i) {
        init_semaphore(matmul_finished_1[i], 0, 1);
        init_semaphore(matmul_finished_2[i], 0, 1);
        init_semaphore(inputs_all_arrived[i], 0, 1);
        init_semaphore(inputs_smem_arrived[i], 0, 1);
    }
    init_semaphore(outputs_arrived, 0, 1);
    init_semaphore(tensor_finished, 0, 1);
  }
  everyone::tma::cluster::sync();

  int lane_id = warp::laneid();
  int warp_id = warpgroup::warpid();
  int warpgroup_id = warpgroup::groupid();
  int cta_id = cluster_ctarank();
  int cluster_id = clusterIdx().x;
  
  // todo: super blocking
  int num_blocks_per_col = G.C.cols() / globals::TILE_SIZE;
  int num_blocks_per_row = G.C.rows() / globals::TILE_SIZE;
  int num_blocks = num_blocks_per_row * num_blocks_per_col;
  int num_iters_per_block = G.A.cols() / globals::REDUCTION_BLOCK; // todo: fix

    uint32_t stage = 0;
    uint32_t phasebits = 0xFFFF0000;
    uint32_t phasebits_2 = 0xFFFF0000;
    uint32_t last_stage = globals::PIPELINE_STAGES;

    __syncthreads();
    if (warpgroup_id == config::NUM_WARPGROUPS - 1) { // Producer group
      // warpgroup::increase_registers<config::PRODUCER_REGISTERS>();
      // warpgroup::decrease_registers<config::PRODUCER_REGISTERS>();
      if (warp_id == 1 && lane_id == 0) {
        // printf("G.A.cols() = %d\n", G.A.cols());
        // printf("num_iters_per_block = %d\n", num_iters_per_block);
        for (int block_idx = cluster_id; block_idx < num_blocks; block_idx += gridDim.x / config::CLUSTER_SIZE) {
          int row_block_idx = block_idx / num_blocks_per_col;
          int col_block_idx = block_idx % num_blocks_per_col;

          for (int i = 0; i < num_iters_per_block; ++i) {
            wait(matmul_finished_1[stage], get_phasebit<1>(phasebits, stage));
            wait(matmul_finished_2[stage], get_phasebit<1>(phasebits_2, stage));
            update_phasebit<1>(phasebits, stage);
            update_phasebit<1>(phasebits_2, stage);

            if (stage == last_stage) {
                arrive(outputs_arrived);
                last_stage = globals::PIPELINE_STAGES;
            }

            tma::expect_bytes(inputs_smem_arrived[stage], sizeof(globals::A_fp4x2_tile) + sizeof(globals::sfA_tile) + (sizeof(globals::B_fp4x2_tile) + sizeof(globals::sfB_tile)) * 2);
            tma::load_async(input_tiles[stage].A, G.A,      {row_block_idx, i}, inputs_smem_arrived[stage]);
            tma::load_async(input_tiles[stage].B, G.B,      {col_block_idx, i}, inputs_smem_arrived[stage]);
            tma::load_async(input_tiles[stage].B_2, G.B_2,  {col_block_idx, i}, inputs_smem_arrived[stage]);
            // todo: consider sepearete loading warp for B2

            tma::load_async(input_scales[stage].sfA, G.sfA, {row_block_idx, i, 0, 0}, inputs_smem_arrived[stage]);
            tma::load_async(input_scales[stage].sfB, G.sfB, {col_block_idx, i, 0, 0}, inputs_smem_arrived[stage]);
            tma::load_async(input_scales[stage].sfB_2, G.sfB_2, {col_block_idx, i, 0, 0}, inputs_smem_arrived[stage]);
            if (i == num_iters_per_block - 1) {
              last_stage = stage;
            }
            stage = (stage + 1) % globals::PIPELINE_STAGES;
        }
      }
      if (last_stage < globals::PIPELINE_STAGES) {
        wait(matmul_finished_1[last_stage], get_phasebit<1>(phasebits, last_stage));
        wait(matmul_finished_2[last_stage], get_phasebit<1>(phasebits_2, last_stage));
        arrive(outputs_arrived);
      }
    } else if (warp_id == 0 && lane_id == 0) {
      for (int block_idx = cluster_id; block_idx < num_blocks; block_idx += gridDim.x / config::CLUSTER_SIZE) {
        for (int i = 0; i < num_iters_per_block; i++) {
                wait(inputs_smem_arrived[stage], get_phasebit<0>(phasebits, stage));
                wait(inputs_smem_arrived[stage], get_phasebit<0>(phasebits_2, stage));
                update_phasebit<0>(phasebits, stage);
                update_phasebit<0>(phasebits_2, stage);

                #pragma unroll
                for (int ii = 0; ii < globals::PACKED_PER_TILE; ii++) {
                    auto A_sc_tm_subtile    = A_sc_tm.subtile<full_tt_fp8e4m3<16>>(stage * globals::PACKED_PER_TILE * 16 + ii * 16);
                    auto B_sc_tm_subtile    = B_sc_tm.subtile<full_tt_fp8e4m3<16>>(stage * globals::PACKED_PER_TILE * 16 + ii * 16);
                    auto &A_sc_sm_subtile   = *reinterpret_cast<st_fp8e4m3<32, 16, false> *>(reinterpret_cast<uint64_t>(&input_scales[stage].sfA.data[0]) + 16 * 32 * ii);
                    auto &B_sc_sm_subtile   = *reinterpret_cast<st_fp8e4m3<32, 16, false> *>(reinterpret_cast<uint64_t>(&input_scales[stage].sfB.data[0]) + 16 * 32 * ii);
                    load_mxnv_scale_async(A_sc_tm_subtile,   A_sc_sm_subtile);
                    load_mxnv_scale_async(B_sc_tm_subtile, B_sc_sm_subtile);

                    auto B_sc_tm_subtile_2  = B_sc_tm_2.subtile<full_tt_fp8e4m3<16>>(stage * globals::PACKED_PER_TILE * 16 + ii * 16);
                    auto &B_sc_sm_subtile_2   = *reinterpret_cast<st_fp8e4m3<32, 16, false> *>(reinterpret_cast<uint64_t>(&input_scales[stage].sfB_2.data[0]) + 16 * 32 * ii);
                    load_mxnv_scale_async(B_sc_tm_subtile_2, B_sc_sm_subtile_2);
                }

                // asm volatile(
                //     "{tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];}"
                // ::  "l"(__cvta_generic_to_shared(&inputs_all_arrived[stage]))); // todo: isn't there a wrapper?

                kittens::detail::tcgen05::commit<1>(inputs_all_arrived[stage]);
                stage = (stage + 1) % globals::PIPELINE_STAGES;
        }
      }
    } else if (warp_id == 2 && lane_id == 0) {
      for (int block_idx = cluster_id; block_idx < num_blocks; block_idx += gridDim.x / config::CLUSTER_SIZE) {
        wait(tensor_finished, get_phasebit<1>(phasebits_2, globals::PIPELINE_STAGES));
        update_phasebit<1>(phasebits_2, globals::PIPELINE_STAGES);
            for (int i = 0; i < num_iters_per_block; i++) {
              wait(inputs_all_arrived[stage], get_phasebit<0>(phasebits_2, stage));
              update_phasebit<0>(phasebits_2, stage);
              if (i == 0) mm_ABt(out_tm_2, input_tiles[stage].A, input_tiles[stage].B_2, 
                      A_sc_tm.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16), 
                      B_sc_tm_2.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16), 
                      matmul_finished_2[stage]);
              else mma_ABt(out_tm_2, input_tiles[stage].A, input_tiles[stage].B_2, 
                      A_sc_tm.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16), 
                      B_sc_tm_2.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16), 
                      matmul_finished_2[stage]);
              if (i == num_iters_per_block - 1) {
                last_stage = stage;
              }
              stage = (stage + 1) % globals::PIPELINE_STAGES;
        }
      }
    } else if (warp_id == 3 && lane_id == 0) {
      for (int block_idx = cluster_id; block_idx < num_blocks; block_idx += gridDim.x / config::CLUSTER_SIZE) {
        wait(tensor_finished, get_phasebit<1>(phasebits, globals::PIPELINE_STAGES));
        update_phasebit<1>(phasebits, globals::PIPELINE_STAGES);
        for (int i = 0; i < num_iters_per_block; i++) {
          wait(inputs_all_arrived[stage], get_phasebit<0>(phasebits, stage));
          update_phasebit<0>(phasebits, stage);

          if (i == 0) mm_ABt(out_tm, input_tiles[stage].A, input_tiles[stage].B,
              A_sc_tm.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16), 
              B_sc_tm.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16),
              matmul_finished_1[stage]);
          else mma_ABt(out_tm, input_tiles[stage].A, input_tiles[stage].B,
              A_sc_tm.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16), 
              B_sc_tm.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16),
              matmul_finished_1[stage]);
          stage = (stage + 1) % globals::PIPELINE_STAGES;
        }
      }
    }
  } 
  else {
    // warpgroup::increase_registers<256>();
    for (int block_idx = cluster_id; block_idx < num_blocks;
         block_idx += gridDim.x / config::CLUSTER_SIZE) {
      int row_block_idx = block_idx / num_blocks_per_col;
      int col_block_idx = block_idx % num_blocks_per_col;

      wait(outputs_arrived,
           get_phasebit<0>(phasebits, globals::PIPELINE_STAGES));
      // wait(outputs_arrived, get_phasebit<0>(phasebits_2,
      // globals::PIPELINE_STAGES));
      update_phasebit<0>(phasebits, globals::PIPELINE_STAGES);
      update_phasebit<0>(phasebits_2, globals::PIPELINE_STAGES);

      rt_fl<globals::TILE_SIZE / 4, globals::TILE_SIZE> C_reg_2;
      rt_fl<globals::TILE_SIZE / 4, globals::TILE_SIZE> C_reg_1;

      warpgroup::load_async(C_reg_1, out_tm);
      tensor_load_wait(); // are those syncs necessary?
      warpgroup::sync(1);

      warpgroup::load_async(C_reg_2, out_tm_2);
      tensor_load_wait();
      warpgroup::sync(1);

      // todo: explicit unrolling to guide compiler
      #pragma unroll // TODO: software emulate exp on CUDA Cores
      for (int i = 0; i < C_reg_1.height; i++) {
        #pragma unroll
        for (int j = 0; j < C_reg_1.width; j++) {
          #pragma unroll
          for (int k = 0; k < C_reg_1.packed_per_tile; k++) {
            float x = C_reg_1.tiles[i][j].data[k].x;
            float y = C_reg_1.tiles[i][j].data[k].y;
      
            float sx = (x >= 0.0f)
              ? 1.0f / (1.0f + expf(-x))
              : expf(x) / (1.0f + expf(x)); // todo: packed instructions?
      
            float sy = (y >= 0.0f)
              ? 1.0f / (1.0f + expf(-y))
              : expf(y) / (1.0f + expf(y));
      
            float silu_x = x * sx;
            float silu_y = y * sy;
      
            float gate_x = C_reg_2.tiles[i][j].data[k].x;
            float gate_y = C_reg_2.tiles[i][j].data[k].y;
      
            C_reg_1.tiles[i][j].data[k].x = silu_x * gate_x;
            C_reg_1.tiles[i][j].data[k].y = silu_y * gate_y;
          }
        }
      }

      if (warpgroup::laneid() == 0) arrive(tensor_finished, 1); // todo: can this be done earlier?

      warpgroup::store(output_tiles.C, C_reg_1);
      warpgroup::sync(1);
      warpgroup::tma::store_async(G.C, output_tiles.C, {row_block_idx, col_block_idx});
      warpgroup::tma::store_async_read_wait();
    }
   }
  }

void entrypoint(const at::Tensor &A, const at::Tensor &B,
                const at::Tensor &B_2, at::Tensor &C,
                const at::Tensor &sfA, const at::Tensor &sfB,
                const at::Tensor &sfB_2) {


    auto sfA_swizzled = sfA.view({sfA.size(0) / 128, 128, sfA.size(1) / 4, 4})
                            .transpose(1,2)
                            .reshape({sfA.size(0)/128, sfA.size(1)/4, 4, 32, 4})
                            .transpose(-2,-3)
                            .reshape({sfA.size(0)/128, sfA.size(1)/4, 32, 16})
                            .reshape({sfA.size(0)/128, (sfA.size(1)/4)/4, 128, 16});

    auto sfB_swizzled = sfB.view({sfB.size(0) / 128, 128, sfB.size(1) / 4, 4})
                            .transpose(1,2)
                            .reshape({sfB.size(0)/128, sfB.size(1)/4, 4, 32, 4})
                            .transpose(-2,-3)
                            .reshape({sfB.size(0)/128, sfB.size(1)/4, 32, 16})
                            .reshape({sfB.size(0)/128, (sfB.size(1)/4)/4, 128, 16});

    auto sfB2_swizzled = sfB_2.view({sfB_2.size(0) / 128, 128, sfB_2.size(1) / 4, 4})
                            .transpose(1,2)
                            .reshape({sfB_2.size(0)/128, sfB_2.size(1)/4, 4, 32, 4})
                            .transpose(-2,-3)
                            .reshape({sfB_2.size(0)/128, sfB_2.size(1)/4, 32, 16})
                            .reshape({sfB_2.size(0)/128, (sfB_2.size(1)/4)/4, 128, 16});


    globals G{
        .A       = kittens::py::tensor_to_gl<globals::A_fp4x2_gl>(A),
        .B       = kittens::py::tensor_to_gl<globals::B_fp4x2_gl>(B),
        .B_2     = kittens::py::tensor_to_gl<globals::B_fp4x2_gl>(B_2),
        .C       = kittens::py::tensor_to_gl<globals::C_gl>(C),
        .sfA     = kittens::py::tensor_to_gl<globals::sfA_gl>(sfA_swizzled),
        .sfB     = kittens::py::tensor_to_gl<globals::sfB_gl>(sfB_swizzled),
        .sfB_2   = kittens::py::tensor_to_gl<globals::sfB_gl>(sfB2_swizzled)
    };
    kittens::py::launch_kernel<config, globals, kernel>(G);

}

} // namespace nvfp4_gemm

PYBIND11_MODULE(_C, m) {
  m.doc() = "nvpf4 python module";
  m.def("kernel", &nvfp4_gemm::entrypoint);
}

'''

include_flags = [
    f'-I{thunderkittens_dir}/include',
    f'-I{thunderkittens_dir}/prototype',
]

extra_cflags = ['-O3','-std=c++20','-DNDEBUG','-fPIE','-fopenmp','-Wno-psabi','-fno-strict-aliasing','-fPIC'] + include_flags
extra_cuda_cflags = [
    '-O3','-std=c++20','-DNDEBUG',
    '--expt-extended-lambda','--expt-relaxed-constexpr',
    '-Xcompiler=-fPIE','-Xcompiler=-fopenmp',
    '-Xcompiler=-Wno-psabi','-Xcompiler=-fno-strict-aliasing',
    '--use_fast_math','-forward-unknown-to-host-compiler',
    '-Xcompiler=-fPIC',
    # '-DKITTENS_HOPPER',
    '-DKITTENS_BLACKWELL',
    '-D__CUDA_NO_HALF_OPERATORS__',
    '-D__CUDA_NO_HALF_CONVERSIONS__',
    '-D__CUDA_NO_BFLOAT16_CONVERSIONS__',
    '-D__CUDA_NO_HALF2_OPERATORS__',
    '-gencode=arch=compute_100a,code=sm_100a'
] + include_flags

extra_ldflags = [
    '-L/usr/local/cuda/lib64',
    '-lcuda',
    '-lcudart',
    '-L/home/runner/competition/cuda/lib/python3.12/site-packages/torch/lib',
    '-lc10', '-lc10_cuda', '-ltorch_cpu', '-ltorch_cuda', '-ltorch', '-ltorch_python',
    '-Wl,-rpath,/usr/local/cuda/lib64',
]

nvpf4 = load_inline(
    name='_C',
    cpp_sources='',
    cuda_sources=nvpf4_src,
    extra_cflags=extra_cflags,
    extra_cuda_cflags=extra_cuda_cflags,
    extra_ldflags=extra_ldflags,
    verbose=True,
    with_cuda=True,
)

def scale_swizzle(
    V_sc_unswizzled: torch.Tensor, # (M, N // 16) fp8e4m3
    packed_per_tile: int           # number of packed K=64 blocks per scale tile
) -> torch.Tensor:
    # https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
    assert len(V_sc_unswizzled.shape) == 2
    assert V_sc_unswizzled.dtype == torch.float8_e4m3fn
    assert V_sc_unswizzled.shape[0] % 128 == 0
    assert (V_sc_unswizzled.shape[1] * 16) % 64 == 0

    M_BLOCK = 128
    N_BLOCK = 4 # 64 / 16

    M, N_16 = V_sc_unswizzled.shape

    V_sc = V_sc_unswizzled        # (M, N_16)
    V_sc = V_sc.reshape(          # (M / 128, 128, N_16 / 4, 4)
        M // M_BLOCK, M_BLOCK,
        N_16 // N_BLOCK, N_BLOCK
    )
    V_sc = V_sc.transpose(1, 2)   # (M / 128, N_16 / 4, 128, 4) --> last 2 dims are all we need per MM
    V_sc = V_sc.reshape(          # (M / 128, N_16 / 4, 4, 32, 4)
        M // M_BLOCK, N_16 // N_BLOCK,
        4, M_BLOCK // 4, N_BLOCK
    )
    V_sc = V_sc.transpose(-2, -3) # (M / 128, N_16 / 4, 32, 4, 4)
    V_sc = V_sc.reshape(          # (M / 128, N_16 / 4, 32, 16)
        M // M_BLOCK,
        N_16 // N_BLOCK,
        M_BLOCK // 4, N_BLOCK * 4
    )
    V_sc = V_sc.reshape(          # (M / 128, N_16 / 4, 32, 16)
        M // M_BLOCK,             # Pack the scale tiles. This is purely for GEMM efficiency
        N_16 // N_BLOCK // packed_per_tile,
        packed_per_tile * M_BLOCK // 4, N_BLOCK * 4
    )

    return V_sc

def ceil_div(a, b):
    return (a + b - 1) // b

def to_blocked(input_matrix):
    rows, cols = input_matrix.shape

    n_row_blocks = ceil_div(rows, 128)
    n_col_blocks = ceil_div(cols, 4)

    padded = input_matrix
    blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
    rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)

    return rearranged.flatten()

def ref_kernel(
    data: input_t,
) -> output_t:
    """
    PyTorch reference implementation of NVFP4 block-scaled dual GEMM with silu activation,
    C = silu(A @ B1) * (A @ B2).
    """
    a_ref, b1_ref, b2_ref, sfa_ref_cpu, sfb1_ref_cpu, sfb2_ref_cpu, _, _, _, c_ref = data
    
    # Get dimensions from MxNxL layout
    m, n, l = c_ref.shape

    # Call torch._scaled_mm to compute the GEMV result
    ref1 = torch.empty(
        (l, m, n),
        dtype=torch.float32,
        device="cuda",
    ).permute(1, 2, 0)
    ref2 = torch.empty(
        (l, m, n),
        dtype=torch.float32,
        device="cuda",
    ).permute(1, 2, 0)
    for l_idx in range(l):
        # Convert the scale factor tensor to blocked format
        scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx])
        scale_b1 = to_blocked(sfb1_ref_cpu[:, :, l_idx])
        scale_b2 = to_blocked(sfb2_ref_cpu[:, :, l_idx])
        # (m, k) @ (n, k).T -> (m, n)
        res1 = torch._scaled_mm(
            a_ref[:, :, l_idx],
            b1_ref[:, :, l_idx].transpose(0, 1),
            scale_a.cuda(),
            scale_b1.cuda(),
            # torch.ones_like(scale_a, device="cuda"),
            # torch.ones_like(scale_b1, device="cuda"),
            bias=None,
            out_dtype=torch.float32,
        )
        ref1[:, :, l_idx] = res1

        res2 = torch._scaled_mm(
            a_ref[:, :, l_idx],
            b2_ref[:, :, l_idx].transpose(0, 1),
            scale_a.cuda(),
            scale_b2.cuda(),
            # torch.ones_like(scale_a, device="cuda"),
            # torch.ones_like(scale_b2, device="cuda"),
            bias=None,
            out_dtype=torch.float32,
        )
        ref2[:, :, l_idx] = res2
    # Do silu on the first GEMM result and multiply with the second GEMM result
    c_ref = (torch.nn.functional.silu(ref1) * ref2).to(torch.float16)
    # c_ref = ref1
    return c_ref

def custom_kernel(data: input_t) -> output_t:

    a_ref, b1_ref, b2_ref, sfa_ref_cpu, sfb1_ref_cpu, sfb2_ref_cpu, sfa_cute, sfb1_cute, sfb2_cute, c_ref = data

    if a_ref.shape[1] < 256:
        out = ref_kernel(data)
        return out
    else:
        nvpf4.kernel(a_ref, b1_ref, b2_ref, c_ref, sfa_ref_cpu, sfb1_ref_cpu, sfb2_ref_cpu)
        return c_ref
scrolls · 574 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