Skip to content
KernelIndex
Search⌘K

submission 111027

v0i0 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission-28-12-20-32.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-111027?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 GEMVsuite of 3 cases
NVIDIA B200
22.5µs
#46 of 678
2025-11-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:638359569e141628b5b85b7001f44d19e8d2995dbe2ae221e4e45a8346d3d1f0
license declaredunknown
license concludedunknown
authorsv0i0
imported2026-08-15

Techniques

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

async-copy__device__ void cp_async_wait_group() {
fp4uint4 veca = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k)]); // this has 32 fp4 elems
fp8__half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);
mbarrieragain: mbarrier.try_wait.parity.shared::cta.b64 done, [%0], %1, 0x100000;
shared-memory__shared__ clock_t timing[3][128];
split-k"SPLIT_K": 1,
tcgen05asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], 512;" :: "r"((unsigned int) __cvta_generic_to_shared(&tmem_addr)));
tile-k = 256const int BLOCK_K = 256;
tile-m = 128const int BLOCK_M = 128;
tma__grid_constant__ const CUtensorMap desc_a
vector-width = uint4using vec_type = uint4;

Kernel source

submission-28-12-20-32.py1002 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
from triton.testing import do_bench
import itertools


def kernel_key(*args):
    return tuple([(tuple(t.shape), tuple(t.stride())) for t in args])


TUNABLES = {
    # "USE_PDL": [0, 1],
    # "STAGE_COUNT": [4, 6, 7, 8, 9],
    # "CLUSTER_M": [None, 1, 2, 4],
    # "LOCAL_M": [1, 2],
    # "SPLIT_K": [1, 2],
}


def heuristic(a, b, sfa, sfb, c):
    defines = {
        "LOCAL_M": 1,
        "USE_PDL": 1,
        "CLUSTER_M": None,
        "SPLIT_K": 1,
        "STAGE_COUNT": 6,
    }
    if a.shape[0] // 128 * a.shape[2] < 64 and a.shape[1] % 512 == 0:
        defines["SPLIT_K"] = 2
    if a.shape[0] // 128 * a.shape[2] > 128 and a.shape[0] % 256 == 0:
        defines["LOCAL_M"] = 2
    return defines


def tune_all(a, b, sfa, sfb, c):
    results = []
    baseline = heuristic(a, b, sfa, sfb, c)
    for idx, values in enumerate(itertools.product(*TUNABLES.values())):
        tunables = baseline | dict(zip(TUNABLES.keys(), values))
        try:
            kernel = compile_kernel(kernel_key(a, b, sfa, sfb, c), tunables)
            time = do_bench(lambda: kernel.run(a, b, sfa, sfb, c), warmup=5, rep=10)
            print(f"{tunables} {time * 1000:.2f}")
            results.append((time, idx, tunables))
        except Exception as e:
            print(tunables, e)
    return min(results)[-1]


def tune_coordinate_descent(a, b, sfa, sfb, c):
    results = {}
    current = heuristic(a, b, sfa, sfb, c)
    directions = list(TUNABLES.values())


def compile_kernel(key, tunables):
    # print(key, tunables)
    (
        (a_shape, a_stride),
        (_, b_stride),
        (_, sfa_stride),
        (_, sfb_stride),
        (_, c_stride),
    ) = key
    defines = {
        "STATIC_K": 2 * a_shape[1],
        "STRIDE_A_M": a_stride[0],
        "STRIDE_A_L": a_stride[2],
        "STRIDE_B_L": b_stride[2],
        "STRIDE_SFA_M": sfa_stride[2],
        "STRIDE_SFA_K": sfa_stride[4],
        "STRIDE_SFA_L": sfa_stride[5],
        "STRIDE_SFB_K": sfb_stride[4],
        "STRIDE_SFB_L": sfb_stride[5],
        "STRIDE_C_L": c_stride[2],
    } | tunables
    defines = [f"-D{k}={v}" for k, v in defines.items() if v is not None]

    return load_inline(
        cuda_sources=[
            """
#ifdef BUILD_PYTORCH
#endif

#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cuda.h>
#include <vector>
#include <cstdint>

#define ASM(x...) asm(#x
#define ASMV(x...) asm volatile(#x
#define CLO(x...) : x)

#ifndef STATIC_K
#define STATIC_K (2*1024)
#define STRIDE_A_M 1024
#define STRIDE_A_L (1024*1024)
#define STRIDE_B_L (1024*1024)
#define STRIDE_C_L (1024*1024)
#define STRIDE_SFA_M 16384
#define STRIDE_SFA_K (4*16384)
#define STRIDE_SFA_L (8*16384)
#define STRIDE_SFB_K (4*16384)
#define STRIDE_SFB_L (8*16384)
#define LOCAL_M 1
#define SPLIT_K 1
#define STAGE_COUNT 6
#define USE_PDL 1
#endif

__device__ int semaphore[128][32] = {0};

const int BLOCK_M = 128;
const int BLOCK_K = 256;

template<int N>
__device__ void cp_async_wait_group() {
  asm("cp.async.wait_group %0;" :: "n"(N) : "memory");
}

__device__ void copy_a_vec(void* dst, const void* src) {
  asm("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"((unsigned int)__cvta_generic_to_shared(dst)), "l"(src) : "memory");
};

__device__ void copy_b_vec(void* dst, const void* src) {
  asm("cp.async.ca.shared.global [%0], [%1], 16;" :: "r"((unsigned int)__cvta_generic_to_shared(dst)), "l"(src) : "memory");
};

__device__ void copy_sf_vec(void* dst, const void* src) {
  asm("cp.async.ca.shared.global [%0], [%1], 4;" :: "r"((unsigned int)__cvta_generic_to_shared(dst)), "l"(src) : "memory");
};

__device__ unsigned int make_warp_uniform(unsigned int v) {
  return __shfl_sync(0xffffffff, v, 0);
}

__device__ void barrier_wait(unsigned long long* barrier, int barrier_wait_phase) {
  ASMV({
    .reg .pred done;
    again: mbarrier.try_wait.parity.shared::cta.b64 done, [%0], %1, 0x100000;
    @done bra end;
    bra again;
    end:
  }) CLO(: "r"((unsigned int) __cvta_generic_to_shared(barrier)), "r"(barrier_wait_phase));
}

__device__ float e2m1x2_to_float(__nv_fp4x2_e2m1 value, int subbyte_idx) {
  __half2_raw values = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<const __nv_fp4x2_storage_t&>(value), __NV_E2M1);
  return __half2float(reinterpret_cast<const __half&>(subbyte_idx == 0 ? values.x : values.y));
}

__device__ __half2 e2m1x2_to_half2(__nv_fp4x2_e2m1 value) {
  __half2_raw values = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<const __nv_fp4x2_storage_t&>(value), __NV_E2M1);
  return reinterpret_cast<const __half2&>(values);
}

__device__ __half2 e2m1x2_to_half2(unsigned char value) {
  __half2_raw values = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<const __nv_fp4x2_storage_t&>(value), __NV_E2M1);
  return reinterpret_cast<const __half2&>(values);
}

__device__ float e4m3_to_float(unsigned char value) {
  __half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);
  return __half2float(reinterpret_cast<const __half&>(half_value));
}

__device__ float e4m3_to_float(__nv_fp8_e4m3 value) {
  __half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);
  return __half2float(reinterpret_cast<const __half&>(half_value));
}

__device__ __half e4m3_to_half(__nv_fp8_e4m3 value) {
  __half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);
  return reinterpret_cast<const __half&>(half_value);
}

__device__ __half e4m3_to_half(unsigned char value) {
  return e4m3_to_half(reinterpret_cast<const __nv_fp8_e4m3&>(value));
}

#ifdef TIMING
#define TIME_EVENT(event_id) do { \
  if (is_timer) { \
    timing[time_slot][current_slot] = clock(); \
    tag[time_slot][current_slot] = event_id; \
    current_slot += 1; \
  } \
} while (0)
#else
#define TIME_EVENT(x) do {} while(0)
#endif

struct TileIterator {
  int start_m, start_k;
  int num_tiles;
  int total_k_tiles;
  int total_m_tiles;
  int total_sms;
  int index;
  // simple, we have N tiles
  // we need to be able to decode them
  // locally, we need to know this tile and the next tile
  // we need to be able to chunk
  //
};

__device__ int get_chunked_local_pos(int num, int idx, int total) {
  int count = total / num;
  int rest = total % num;
  int pos = min(idx, rest) + idx * count;
  return pos;
}

__launch_bounds__(256, 1)
__global__ void kernel(
  const __nv_fp4x2_e2m1 * __restrict__ ptr_a,
  const __nv_fp8_e4m3 * __restrict__ ptr_sfa,
  const __nv_fp4x2_e2m1 * __restrict__ ptr_b,
  const __nv_fp8_e4m3 * __restrict__ ptr_sfb,
  __half * __restrict__ ptr_c,
  __grid_constant__ const CUtensorMap desc_a
) {
#ifdef TIMING
  __shared__ clock_t timing[3][128];
  __shared__ int tag[3][128];
  bool is_timer = threadIdx.x == 0 || threadIdx.x == 128 || threadIdx.x == (128+32);
  int time_slot = threadIdx.x == 0 ? 0 : threadIdx.x == 128 ? 1 : 2;
  int current_slot = 0;
  TIME_EVENT(0);
#endif


  // let's do the "real" blocking: 128x64 in M and K, block along M and L
  ptr_a += blockIdx.y * STRIDE_A_L;
  ptr_b += blockIdx.y * STRIDE_B_L;
  ptr_sfa += blockIdx.y * STRIDE_SFA_L;
  ptr_sfb += blockIdx.y * STRIDE_SFB_L;
  ptr_c += blockIdx.y * STRIDE_C_L;

  int block_offset_m = blockIdx.x / SPLIT_K;
  int block_offset_k = blockIdx.x % SPLIT_K;
  int k = STATIC_K / SPLIT_K;

  ptr_a += block_offset_m * BLOCK_M * STRIDE_A_M * LOCAL_M;
  ptr_sfa += block_offset_m * (BLOCK_M / 128) * STRIDE_SFA_M * LOCAL_M;
  ptr_c += block_offset_m * BLOCK_M * LOCAL_M;

  ptr_a += block_offset_k * (k / 2);
  ptr_b += block_offset_k * (k / 2);
  ptr_sfa += block_offset_k * (k / 16 / 4) * STRIDE_SFA_K;
  ptr_sfb += block_offset_k * (k / 16 / 4) * STRIDE_SFB_K;

  // have 128 threads, each computing one output element
  // SF: MN, K -> addr is ((32, 4, RM), (16, 4, RK)):((16, 4, ?), (0, 1, ?))

  __attribute__((aligned(128))) __shared__ __nv_fp8_e4m3 smem_sfb[STAGE_COUNT][4 * BLOCK_K / 16];
  __attribute__((aligned(1024))) __shared__ __nv_fp4x2_e2m1 smem_b[STAGE_COUNT][BLOCK_K / 2];
  __attribute__((aligned(128))) __shared__ __nv_fp8_e4m3 smem_sfa[STAGE_COUNT][BLOCK_M * (BLOCK_K / 16)];
  // 9 * 16, 2 * 9 * 17
  __attribute__((aligned(1024))) __shared__ __nv_fp4x2_e2m1 smem_a[STAGE_COUNT][BLOCK_M * (BLOCK_K / 2)];
  __attribute__((aligned(16))) __shared__ __half result[LOCAL_M * BLOCK_M];

  int warp_id = threadIdx.x / 32;
  int wg_id = warp_id / 4;
  bool is_load = wg_id == 0;
  bool is_mma = warp_id == 4;
  bool is_reset = warp_id == 5;
  bool is_tma = warp_id == 6;
  int lane_id = threadIdx.x % 32;
  __shared__ int tmem_addr;
  __shared__ unsigned long long barrier;
  if (warp_id == 0 && lane_id == 0) {
    asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"((unsigned int) __cvta_generic_to_shared(&barrier)));
  }
  int barrier_wait_phase = 0;
  __shared__ unsigned long long barrier_smem_empty[STAGE_COUNT];
  if (warp_id == 1 && lane_id < STAGE_COUNT) {
    asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_empty[lane_id])));
  }
  __shared__ unsigned long long barrier_smem_full[STAGE_COUNT];
  if (warp_id == 2 && lane_id < STAGE_COUNT) {
    asm volatile("mbarrier.init.shared::cta.b64 [%0], 128;" :: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[lane_id])));
  }
  if (warp_id == 3) {
    ASMV(prefetch.tensormap [%0];) CLO(: "l"(&desc_a));
  }
  TIME_EVENT(1);
  __syncthreads();
  TIME_EVENT(2);

  auto mde = [](unsigned int x) { return (unsigned long long) ((x & 0x3FFFF) >> 4); };
  auto make_smem_desc = [mde](void* smem_ptr, unsigned long long leading_dimension, unsigned long long stride_dimension, unsigned long long matrix_base_offset, unsigned long long is_leading_byte_address_absolute, unsigned long long swizzling_mode) {
    return mde(__cvta_generic_to_shared(smem_ptr)) | (mde(leading_dimension) << 16) | (mde(stride_dimension) << 32) | (1ull << 46) | (matrix_base_offset << 49) | (is_leading_byte_address_absolute << 52) | (swizzling_mode << 61);
  };


  auto offset_a = [&](int m, int k) {
    return (k % 16) + 16 * (((k / 16) % 8) ^ (m % 8)) + 128 * m;
  };
  using sf_vec_type = unsigned int;

  int cutover_k = (STAGE_COUNT / LOCAL_M) * BLOCK_K;
  if (cutover_k > k) cutover_k = k;
  // cutover_k = 0;

  if (is_load) {

    const int vec_size_elems = 32;
    const int threads_in_k = BLOCK_K / vec_size_elems;
    const int threads_in_m = BLOCK_M / threads_in_k;
    using vec_type = uint4;
    int ptr_a_m_offset = threadIdx.x / threads_in_k;
    int ptr_a_k_offset = (threadIdx.x % threads_in_k) * sizeof(vec_type);
    ptr_a += ptr_a_k_offset + ptr_a_m_offset * STRIDE_A_M;

    int ptr_b_n_offset = threadIdx.x / threads_in_k;
    int ptr_b_k_offset = (threadIdx.x % threads_in_k) * sizeof(vec_type);

    auto smem_a_offs = reinterpret_cast<__nv_fp4x2_e2m1 (*)[BLOCK_M*BLOCK_K/2]>(&smem_a[0][offset_a(ptr_a_m_offset, ptr_a_k_offset)]);
    // ptr_sfa += threadIdx.x * 16;
    // auto smem_sfa_offs = reinterpret_cast<__nv_fp8_e4m3 (*)[BLOCK_M*BLOCK_K/16]>(&smem_a[0][(threadIdx.x % 32) * 16 + 128 * 4 * (threadIdx.x / 32)]);

    int has_sfb = threadIdx.x * sizeof(sf_vec_type) < BLOCK_K / 16;

    ptr_b += ptr_b_k_offset;

    int load_stage = 0;
    int load_phase = 2;  // modify to not wait initially?
    asm volatile("griddepcontrol.wait;" ::: "memory");
    for (int block_k = 0; block_k < cutover_k; block_k += BLOCK_K) {
      #pragma unroll
      for (int block_m = 0; block_m < LOCAL_M; block_m++) {
        TIME_EVENT(3);
        if (load_phase != 2) {
          barrier_wait(&barrier_smem_empty[load_stage], load_phase);
        }
        TIME_EVENT(4);
        #pragma unroll
        for (int partition_m = 0; partition_m < BLOCK_M; partition_m += threads_in_m) {
          copy_a_vec(&smem_a_offs[load_stage][partition_m * 128], &ptr_a[block_k / 2 + partition_m * STRIDE_A_M + block_m * BLOCK_M * STRIDE_A_M]);
        }
        if (ptr_b_n_offset == 0) {
          copy_b_vec(&smem_b[load_stage][ptr_b_k_offset], &ptr_b[block_k / 2]);
        }
        #pragma unroll
        for (int rest_k = threadIdx.x / 32; rest_k < BLOCK_K / 16 / 4; rest_k += BLOCK_M / 32) {
          copy_a_vec(&smem_sfa[load_stage][(threadIdx.x % 32) * 16 + 128 * 4 * rest_k], &ptr_sfa[(threadIdx.x % 32) * 16 + (rest_k + block_k / 64) * STRIDE_SFA_K + block_m * (BLOCK_M / 128) * STRIDE_SFA_M]);
        }
        if (has_sfb) {
          int partition_k = threadIdx.x * sizeof(sf_vec_type);
          copy_sf_vec(&smem_sfb[load_stage][4 * partition_k], &ptr_sfb[((block_k / 16 + partition_k) % 4) + (block_k / 64 + partition_k / (64 / 16)) * STRIDE_SFB_K]);
        }
        ASMV(cp.async.mbarrier.arrive.noinc.shared::cta.b64 [%0];) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])));
        load_stage += 1;
        if (load_stage == STAGE_COUNT) {
          load_stage = 0;
          load_phase = load_phase ? 0 : 1;
        }
      }
    }
    ASMV(bar.arrive 1, 160;) CLO();
#if 0
    // start computing
    int read_stage = 0;
    int read_phase = 0;
    int local_m = threadIdx.x;
    float acc[LOCAL_M][8] = {0};
    int idx = 0;
    for (int block_k = 0; block_k < k; block_k += BLOCK_K) {
      #pragma unroll
      for (int block_m = 0; block_m < LOCAL_M; block_m++) {

        idx += 1;
        if (idx % 2 == 0) {
          read_stage += 1;
          if (read_stage == STAGE_COUNT) {
            read_stage = 0;
            read_phase ^= 1;
          }
          continue;
        }

        barrier_wait(&barrier_smem_full[read_stage], read_phase);

        // each thread reads one column of K values from A and B
        // that means 256 / 32 = 8 elems, or 4 bytes.
        // seems tunable, too.
        // if i wanted 16 bytes, i'd have 8 threads along K
        // so four along M, so B is only reused 8 times. that might be fine.

        // load B: 32 elems, dequantize to fp16
        // load sfB: 2 elem, dequantize to fp16
        // load A: 32 elems in K, 8 elems in M, dequantize to fp16
        // load sfB: 2 elem, dequantize to fp16
        // update each accumulator

        // do cutover: take every other one to tc or to local
        int partition_k = (lane_id % 8) * 16;

        // each thread loads 16 bytes == 32 elements
        // for blocking 256, there are 8 (= 256 / 32) threads in the k dimension then
        // this means there are 4 threads in the m dimension
        // this means each thread then has 32 / 4 = 8 elements in the m dimension
        // let them be contiguous
        uint4 vecb = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k]);
        __half2 vecbv[sizeof(vecb)];
        #pragma unroll
        for (int i = 0; i < sizeof(vecb); i++) {
          vecbv[i] = e2m1x2_to_half2(reinterpret_cast<unsigned char*>(&vecb)[i]);
        }

        uchar2 vecsfb = *reinterpret_cast<uchar2*>(&smem_sfb[read_stage][16 * (partition_k / 32) + (partition_k / 8) % 4]);

        __half value_sfb[2] = {e4m3_to_half(vecsfb.x), e4m3_to_half(vecsfb.y)};

        #pragma unroll
        for (int the_m = 0; the_m < 8; the_m += 1) {
          int local_m = warp_id * 32 + the_m  + (lane_id / 8) * 8;
          uint4 veca = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k)]); // this has 32 fp4 elems
          __half2 vecav[sizeof(veca)];
          #pragma unroll
          for (int i = 0; i < sizeof(veca); i++) {
            vecav[i] = e2m1x2_to_half2(reinterpret_cast<unsigned char*>(&veca)[i]);
          }

          uchar2 vecsfa = *reinterpret_cast<uchar2*>(&smem_sfa[read_stage][(partition_k / 8 / 4) * 128 * 4 + ((partition_k / 8) % 4) + (local_m % 32) * 16 + (local_m / 32) * 4]);
          __half value_sfa[2] = {e4m3_to_half(vecsfa.x), e4m3_to_half(vecsfa.y)};
         
          __half2 local_acc[2];
          local_acc[0] = __hmul2(vecav[0], vecbv[0]);
          local_acc[1] = __hmul2(vecav[8], vecbv[8]);
          #pragma unroll
          for (int i = 1; i < 8; i++) {
            local_acc[0] = __hfma2(vecav[i+0], vecbv[i+0], local_acc[0]);
            local_acc[1] = __hfma2(vecav[i+8], vecbv[i+8], local_acc[1]);
          }
          acc[block_m][the_m] += __half2float(__hmul(__hadd(local_acc[0].x, local_acc[0].y), __hmul(value_sfa[0], value_sfb[0])));
          acc[block_m][the_m] += __half2float(__hmul(__hadd(local_acc[1].x, local_acc[1].y), __hmul(value_sfa[1], value_sfb[1])));
        }
        

#if 0

        // better maybe: each warp still has 32 elements, but each thread in the warp does, too, and they are split along K
        #pragma unroll
        for (int partition_k = 0; partition_k < BLOCK_K / 2; partition_k += 2 * sizeof(uint4)) {
          uint4 veca = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k)]);
          uint4 vecb = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k]);
          uchar4 vecsfa = *reinterpret_cast<uchar4*>(&smem_sfa[read_stage][(partition_k / 8 / 4) * 128 * 4 + (local_m % 32) * 16 + (local_m / 32) * 4]);
          uchar4 vecsfb = *reinterpret_cast<uchar4*>(&smem_sfb[read_stage][4 * partition_k / 8]);

          uint4 vecc = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k + sizeof(uint4))]);
          uint4 vecd = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k + sizeof(uint4)]);

          auto dot_sf = [](unsigned char sfa, unsigned char sfb, unsigned int a0, unsigned int a1, unsigned int b0, unsigned int b1) {
            uchar4 veca = reinterpret_cast<const uchar4&>(a0);
            uchar4 vecb = reinterpret_cast<const uchar4&>(b0);
            __half2 value_a = e2m1x2_to_half2(veca.x);
            __half2 value_b = e2m1x2_to_half2(vecb.x);
            __half2 acc = __hmul2(value_a, value_b);
            value_a = e2m1x2_to_half2(veca.y);
            value_b = e2m1x2_to_half2(vecb.y);
            acc = __hfma2(value_a, value_b, acc);
            value_a = e2m1x2_to_half2(veca.z);
            value_b = e2m1x2_to_half2(vecb.z);
            acc = __hfma2(value_a, value_b, acc);
            value_a = e2m1x2_to_half2(veca.w);
            value_b = e2m1x2_to_half2(vecb.w);
            acc = __hfma2(value_a, value_b, acc);
            veca = reinterpret_cast<const uchar4&>(a1);
            vecb = reinterpret_cast<const uchar4&>(b1);
            value_a = e2m1x2_to_half2(veca.x);
            value_b = e2m1x2_to_half2(vecb.x);
            acc = __hfma2(value_a, value_b, acc);
            value_a = e2m1x2_to_half2(veca.y);
            value_b = e2m1x2_to_half2(vecb.y);
            acc = __hfma2(value_a, value_b, acc);
            value_a = e2m1x2_to_half2(veca.z);
            value_b = e2m1x2_to_half2(vecb.z);
            acc = __hfma2(value_a, value_b, acc);
            value_a = e2m1x2_to_half2(veca.w);
            value_b = e2m1x2_to_half2(vecb.w);
            acc = __hfma2(value_a, value_b, acc);

            __half sum = __hadd(acc.x, acc.y);
            float value_sfa = e4m3_to_float(sfa);
            float value_sfb = e4m3_to_float(sfb);
            return __half2float(sum) * value_sfa * value_sfb;
          };

          acc[block_m] += dot_sf(vecsfa.x, vecsfb.x, veca.x, veca.y, vecb.x, vecb.y);
          acc[block_m] += dot_sf(vecsfa.y, vecsfb.y, veca.z, veca.w, vecb.z, vecb.w);
          acc[block_m] += dot_sf(vecsfa.z, vecsfb.z, vecc.x, vecc.y, vecd.x, vecd.y);
          acc[block_m] += dot_sf(vecsfa.w, vecsfb.w, vecc.z, vecc.w, vecd.z, vecd.w);
        }
#endif

        ASMV(bar.sync 2, 128;) CLO();
        if (threadIdx.x == 0) {
          ASMV(mbarrier.arrive.shared::cta.b64 _, [%0], 1;) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_empty[read_stage])));
        }

        read_stage += 1;
        if (read_stage == STAGE_COUNT) {
          read_stage = 0;
          read_phase ^= 1;
        }
      }
    }

    for (int block_m = 0; block_m < LOCAL_M; block_m++) {
      for (int i = 0; i < 8; i++) {
        // reduce the accumulator
        float reduced_acc = acc[block_m][i];
        reduced_acc += __shfl_xor_sync((unsigned int) -1, reduced_acc, 1);
        reduced_acc += __shfl_xor_sync((unsigned int) -1, reduced_acc, 2);
        reduced_acc += __shfl_xor_sync((unsigned int) -1, reduced_acc, 4);
        if (threadIdx.x % 8 == 0) {
          result[block_m * BLOCK_M + 32 * warp_id + i + 8 * (lane_id / 8)] = __float2half(reduced_acc);
        }
      }
    }
    // __syncthreads();
    // for (int i = threadIdx.x * 8; i < LOCAL_M * BLOCK_M; i += 128 * 8) {
    //   if (SPLIT_K > 1) {
    //     ASM({
    //       .reg .v4 .f16x2 reg;
    //       ld.shared.v4.b32 reg, [%0];
    //       red.relaxed.gpu.global.add.noftz.v4.f16x2 [%1], reg;
    //     }) CLO(: "r"((unsigned int) __cvta_generic_to_shared(result + i)), "l"(ptr_c + i));
    //   } else {
    //     reinterpret_cast<uint4&>(ptr_c[i]) = reinterpret_cast<uint4&>(result[i]);
    //   }
    // }
#endif
  } else if (is_mma) {
    asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32  [%0], 512;" :: "r"((unsigned int) __cvta_generic_to_shared(&tmem_addr)));
    int read_stage = 0;
    int read_phase = 0;
    int idx = 0;
    int has_zerod[LOCAL_M] = {0};
    for (int block_k = 0; block_k < k; block_k += BLOCK_K) {
      #pragma unroll
      for (int block_m = 0; block_m < LOCAL_M; block_m++) {

        // idx += 1;
        // if (idx % 2 == 1) {
        //   read_stage += 1;
        //   if (read_stage == STAGE_COUNT) {
        //     read_stage = 0;
        //     read_phase ^= 1;
        //   }
        //   continue;
        // }

        int enable_input_d = has_zerod[block_m];
        has_zerod[block_m] = 1;
        // int enable_input_d = block_k != 0;

        TIME_EVENT(5);
        barrier_wait(&barrier_smem_full[read_stage], read_phase);
        TIME_EVENT(6);

        unsigned long long smem_desc_sfa_base = make_smem_desc(&smem_sfa[read_stage][0], 16, 8 * 16, 0, 0, 0);
        unsigned long long smem_desc_sfb_base = make_smem_desc(&smem_sfb[read_stage][0], 16, 8 * 16, 0, 0, 0);
        unsigned long long smem_desc_a_base = make_smem_desc(&smem_a[read_stage][0], 1, 128 * 8, 0, 0, 2);
        unsigned long long smem_desc_b_base = make_smem_desc(&smem_b[read_stage][0], 1, 128 * 8, read_stage, 0, 2);

        unsigned int smem_desc_sfa_hi = make_warp_uniform((unsigned int) (smem_desc_sfa_base >> 32));
        unsigned int smem_desc_sfa_lo = make_warp_uniform((unsigned int) (smem_desc_sfa_base >> 0));
        unsigned int smem_desc_sfb_hi = make_warp_uniform((unsigned int) (smem_desc_sfb_base >> 32));
        unsigned int smem_desc_sfb_lo = make_warp_uniform((unsigned int) (smem_desc_sfb_base >> 0));

        unsigned int smem_desc_a_hi = make_warp_uniform((unsigned int) (smem_desc_a_base >> 32));
        unsigned int smem_desc_a_lo = make_warp_uniform((unsigned int) (smem_desc_a_base >> 0));
        unsigned int smem_desc_b_hi = make_warp_uniform((unsigned int) (smem_desc_b_base >> 32));
        unsigned int smem_desc_b_lo = make_warp_uniform((unsigned int) (smem_desc_b_base >> 0));

        // unsigned int insn_desc = (1u << 7) | (1u << 10) | (1u << 17) | (1u << 27);  // e2m1, e2m1, N=8, M=128
        int tmem_d = make_warp_uniform(block_m * 16);
        int tmem_sfa_base = make_warp_uniform(32 + read_stage * 8 * (BLOCK_K / 2 / 2 / sizeof(uint4)));

        static_assert(BLOCK_K == 256);

        ASMV({
          .reg .pred elect;
          .reg .pred pred;
          .reg .b32 tmem_d, tmem_sfa_base, enable_input_d, barrier;
          .reg .b32 sfa_lo_start, sfb_lo_start, a_lo_start, b_lo_start;
          .reg .b32 sfa_hi, sfb_hi, a_hi, b_hi;
          .reg .b32 sfa_lo, sfb_lo, a_lo, b_lo;
          .reg .b32 idesc;
          .reg .b64 sfa, sfb, a, b;
          mov.b32 tmem_d, %0;
          mov.b32 tmem_sfa_base, %1;
          mov.b32 enable_input_d, %2;
          mov.b32 barrier, %3;
          mov.b32 sfa_lo_start, %4;
          mov.b32 sfb_lo_start, %5;
          mov.b32 a_lo_start, %6;
          mov.b32 b_lo_start, %7;
          mov.b32 sfa_hi, %8;
          mov.b32 sfb_hi, %9;
          mov.b32 a_hi, %10;
          mov.b32 b_hi, %11;
          elect.sync _|elect, -1;

          add.u32 sfa_lo, sfa_lo_start, 0;
          mov.b64 sfa, {sfa_lo, sfa_hi};
          add.u32 sfb_lo, sfb_lo_start, 0;
          mov.b64 sfb, {sfb_lo, sfb_hi};
          @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+ 0], sfa;
          @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+ 4], sfb;

          add.u32 sfa_lo, sfa_lo_start, 32;
          mov.b64 sfa, {sfa_lo, sfa_hi};
          add.u32 sfb_lo, sfb_lo_start, 1;
          mov.b64 sfb, {sfb_lo, sfb_hi};
          @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+ 8], sfa;
          @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+12], sfb;

          add.u32 sfa_lo, sfa_lo_start, 64;
          mov.b64 sfa, {sfa_lo, sfa_hi};
          add.u32 sfb_lo, sfb_lo_start, 2;
          mov.b64 sfb, {sfb_lo, sfb_hi};
          @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+16], sfa;
          @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+20], sfb;

          add.u32 sfa_lo, sfa_lo_start, 96;
          mov.b64 sfa, {sfa_lo, sfa_hi};
          add.u32 sfb_lo, sfb_lo_start, 3;
          mov.b64 sfb, {sfb_lo, sfb_hi};
          @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+24], sfa;
          @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+28], sfb;

          setp.ne.b32 pred, enable_input_d, 0;
          mov.b32 idesc, 0x8020480;

          add.u32 a_lo, a_lo_start, 0;
          mov.b64 a, {a_lo, a_hi};
          add.u32 b_lo, b_lo_start, 0;
          mov.b64 b, {b_lo, b_hi};
          @elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [tmem_d], a, b, idesc, [tmem_sfa_base+ 0], [tmem_sfa_base+ 4], pred;

          add.u32 a_lo, a_lo_start, 2;
          mov.b64 a, {a_lo, a_hi};
          add.u32 b_lo, b_lo_start, 2;
          mov.b64 b, {b_lo, b_hi};
          @elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [tmem_d], a, b, idesc, [tmem_sfa_base+ 8], [tmem_sfa_base+12], 1;

          add.u32 a_lo, a_lo_start, 4;
          mov.b64 a, {a_lo, a_hi};
          add.u32 b_lo, b_lo_start, 4;
          mov.b64 b, {b_lo, b_hi};
          @elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [tmem_d], a, b, idesc, [tmem_sfa_base+16], [tmem_sfa_base+20], 1;

          add.u32 a_lo, a_lo_start, 6;
          mov.b64 a, {a_lo, a_hi};
          add.u32 b_lo, b_lo_start, 6;
          mov.b64 b, {b_lo, b_hi};
          @elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [tmem_d], a, b, idesc, [tmem_sfa_base+24], [tmem_sfa_base+28], 1;

          @elect tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%3];
        }) CLO(:
          "r"(tmem_d), "r"(tmem_sfa_base), "r"(enable_input_d), "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_empty[read_stage])),
          "r"(smem_desc_sfa_lo), "r"(smem_desc_sfb_lo), "r"(smem_desc_a_lo), "r"(smem_desc_b_lo), 
          "r"(smem_desc_sfa_hi), "r"(smem_desc_sfb_hi), "r"(smem_desc_a_hi), "r"(smem_desc_b_hi)
        );
        // #pragma unroll
        // for (int partition_k = 0; partition_k < BLOCK_K / 64; partition_k += 1) {
        //   int tmem_sfa = tmem_sfa_base + 8 * partition_k;
        //   int tmem_sfb = tmem_sfa + 4;
        //   unsigned long long smem_desc_sfa = smem_desc_sfa_base + partition_k * (128 * 4 / 16);
        //   unsigned long long smem_desc_sfb = smem_desc_sfb_base + partition_k * (16 / 16);
        //   ASMV({
        //     .reg .pred elect;
        //     elect.sync _|elect, 0xFFFFFFFF;
        //     @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;
        //     @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [%2], %3;
        //   }) CLO(: "r"(tmem_sfa), "l"(smem_desc_sfa), "r"(tmem_sfb), "l"(smem_desc_sfb));
        // }
        // #pragma unroll
        // for (int partition_k = 0; partition_k < BLOCK_K / 64; partition_k += 1) {
        //   int tmem_sfa = tmem_sfa_base + 8 * partition_k;
        //   int tmem_sfb = tmem_sfa + 4;
        //   unsigned long long smem_desc_a = smem_desc_a_base + partition_k * (32 / 16);
        //   unsigned long long smem_desc_b = smem_desc_b_base + partition_k * (32 / 16);
        //   ASMV({
        //     .reg .pred pred;
        //     setp.ne.b32 pred, %6, 0;
        //     .reg .pred elect;
        //     elect.sync _|elect, 0xFFFFFFFF;
        //     @elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], pred;
        //   }) CLO(: "r"(tmem_d), "l"(smem_desc_a), "l"(smem_desc_b), "r"(insn_desc), "r"(tmem_sfa), "r"(tmem_sfb), "r"(enable_input_d));
        //   enable_input_d = 1;
        // }
        // ASMV({
        //   .reg .pred elect;
        //   elect.sync _|elect, 0xFFFFFFFF;
        //   @elect tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];
        // }) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_empty[read_stage])));
  
        read_stage += 1;
        if (read_stage == STAGE_COUNT) {
          read_stage = 0;
          read_phase ^= 1;
        }
      }
    }
    ASMV({
     .reg .pred elect;
     elect.sync _|elect, 0xFFFFFFFF;
     @elect tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];
    }) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier)));
    barrier_wait(&barrier, barrier_wait_phase);
  } else if (is_reset && SPLIT_K > 1) {
    for (int i = 2 * lane_id; i < BLOCK_M; i += 64) {
      reinterpret_cast<int&>(ptr_c[i]) = 0;
    }
    __syncwarp();
    if (lane_id == 0) {
      int* sem = &semaphore[gridDim.x / SPLIT_K * blockIdx.y + blockIdx.x / SPLIT_K][0];
      ASMV(red.release.gpu.global.inc.u32 [%0], %1;) CLO(: "l"(sem), "n"(SPLIT_K - 1) : "memory");
      TIME_EVENT(7);
      for (;;) {
        int value = 1;
        ASMV(ld.acquire.gpu.global.b32 %0, [%1];) CLO("=r"(value) : "l"(sem) : "memory");
        if (value == 0) break;
      }
      TIME_EVENT(8);
    }
    __syncwarp();
  } else if (is_tma) {
    int load_stage = 0;
    int load_phase = 2;  // modify to not wait initially?
    for (int block_k = 0; block_k < cutover_k; block_k += BLOCK_K) {
      for (int block_m = 0; block_m < LOCAL_M; block_m++) {
        load_stage += 1;
        if (load_stage == STAGE_COUNT) {
          load_stage = 0;
          load_phase = load_phase ? 0 : 1;
        }
      }
    }
    ASMV(bar.sync 1, 160;) CLO();
    for (int block_k = cutover_k; block_k < k; block_k += BLOCK_K) {
      #pragma unroll
      for (int block_m = 0; block_m < LOCAL_M; block_m++) {
        if (load_phase != 2) {
          barrier_wait(&barrier_smem_empty[load_stage], load_phase);
        }
        ASMV({
          .reg .pred elect;
          elect.sync _|elect, -1;
          @elect cp.async.bulk.tensor.3d.shared::cta.global.tile.mbarrier::complete_tx::bytes [%1], [%0, {%4, %5, %6}], [%2];
          @elect mbarrier.expect_tx.shared::cta.b64 [%2], %3;
        }) CLO(: "l"(&desc_a), "r"((unsigned int) __cvta_generic_to_shared(&smem_a[load_stage][0])), "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])), "n"(BLOCK_M * BLOCK_K / 2), "r"((block_offset_k * k + block_k) / 4 / 2), "r"(block_offset_m * BLOCK_M * LOCAL_M + block_m * BLOCK_M), "r"(blockIdx.y));
        ASMV({
          .reg .pred elect;
          elect.sync _|elect, -1;
          @elect cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes [%1], [%0], %3, [%2];
          @elect mbarrier.expect_tx.shared::cta.b64 [%2], %3;
        }) CLO(: "l"(&ptr_b[block_k / 2]), "r"((unsigned int) __cvta_generic_to_shared(&smem_b[load_stage][0])), "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])), "n"(BLOCK_K / 2));

        ASMV({
          .reg .pred elect;
          elect.sync _|elect, -1;
          @elect cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes [%1], [%0], %3, [%2];
          @elect mbarrier.arrive.expect_tx.shared::cta.b64 _, [%2], %3;
        }) CLO(: "l"(&ptr_sfa[block_k / 64 * STRIDE_SFA_K + block_m * (BLOCK_M / 128) * STRIDE_SFA_M]), "r"((unsigned int) __cvta_generic_to_shared(&smem_sfa[load_stage][0])), "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])), "n"(BLOCK_M * BLOCK_K / 16));
        // if (has_sfb) {
        //   int partition_k = threadIdx.x * sizeof(sf_vec_type);
        //   // 128b loads
        // }
        if (lane_id < BLOCK_K / 16 / 4) {
          copy_sf_vec(&smem_sfb[load_stage][4 * sizeof(sf_vec_type) * lane_id], &ptr_sfb[((block_k / 16 + lane_id * sizeof(sf_vec_type)) % 4) + (block_k / 64 + lane_id * sizeof(sf_vec_type) / (64 / 16)) * STRIDE_SFB_K]);
          ASMV(cp.async.mbarrier.arrive.noinc.shared::cta.b64 [%0];) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])));
        }

        ASMV({
          .reg .pred elect;
          elect.sync _|elect, -1;
          @elect mbarrier.arrive.shared::cta.b64 _, [%0], %1;
        }) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])), "n"(128 - BLOCK_K / 16 / 4 - 1));

        load_stage += 1;
        if (load_stage == STAGE_COUNT) {
          load_stage = 0;
          load_phase = load_phase ? 0 : 1;
        }
      }
    }
  }

  TIME_EVENT(9);
  __syncthreads();
  TIME_EVENT(10);

  // asm volatile("griddepcontrol.launch_dependents;");
  // load tmem
  if (is_load) {

    #pragma unroll
    for (int block_m = 0; block_m < LOCAL_M; block_m++) {
      int tmem_d = block_m * 16;
      float ld_result;
      ASMV(tcgen05.ld.sync.aligned.32x32b.x1.b32 {%0}, [%1];) CLO("=f"(ld_result) : "r"(tmem_d));
      // result[block_m * BLOCK_M + threadIdx.x] = __hadd(result[block_m * BLOCK_M + threadIdx.x], __float2half(ld_result));
      result[block_m * BLOCK_M + threadIdx.x] = __float2half(ld_result);
    }
  }
  TIME_EVENT(11);
  __syncthreads();
  TIME_EVENT(12);
  if (is_mma) {
    asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 0, 512;");
  }
  if (is_load) {
    for (int i = threadIdx.x * 8; i < LOCAL_M * BLOCK_M; i += 128 * 8) {
      if (SPLIT_K > 1) {
        ASM({
          .reg .v4 .f16x2 reg;
          ld.shared.v4.b32 reg, [%0];
          red.relaxed.gpu.global.add.noftz.v4.f16x2 [%1], reg;
        }) CLO(: "r"((unsigned int) __cvta_generic_to_shared(result + i)), "l"(ptr_c + i));
      } else {
        reinterpret_cast<uint4&>(ptr_c[i]) = reinterpret_cast<uint4&>(result[i]);
      }
    }
  }
  TIME_EVENT(13);
#ifdef TIMING
  __syncthreads();
  if (blockIdx.x == 0 && blockIdx.y == 0 && threadIdx.x == 0) {
    for (int i = 0; i < 3; i++) {
      for (int j = 0; j < 128; j += 1) {
        printf("%d %d %d %ld\\n", i, j, tag[i][j], timing[i][j]);
      }
    }
  }
#endif
}

void launch(
  int m,
  int l,
  void * ptr_a,
  void * ptr_sfa,
  void * ptr_b,
  void * ptr_sfb,
  void * ptr_c
) {
  dim3 grid(SPLIT_K * m / BLOCK_M / LOCAL_M, l);
  dim3 block(256);
  void* func = (void*) kernel;

  cudaLaunchConfig_t launch_config = {0};
  launch_config.blockDim = block;
  launch_config.gridDim = grid;
  // launch_config.dynamicSmemBytes = 0;
  cudaLaunchAttribute attrs[16];
  attrs[0].id = cudaLaunchAttributeIgnore;
#ifdef CLUSTER_M
  attrs[0].id = cudaLaunchAttributeClusterDimension;
  attrs[0].val.clusterDim.x = CLUSTER_M;
#endif
  attrs[0].val.clusterDim.y = 1;
  attrs[0].val.clusterDim.z = 1;
  launch_config.attrs = attrs;
  launch_config.numAttrs = 1;
  CUtensorMap desc_a;
  uint64_t global_dims[3] = {STATIC_K / 2 / 4, (uint64_t) m, (uint64_t) l};
  uint64_t global_strides[2] = {STRIDE_A_M, STRIDE_A_L};
  uint32_t box_dim[3] = {BLOCK_K / 2 / 4, BLOCK_M, 1};
  uint32_t element_strides[3] = {1, 1, 1};
  // todo try different L2 PROMOTION
  cuTensorMapEncodeTiled(&desc_a, CU_TENSOR_MAP_DATA_TYPE_INT32, 3, ptr_a, global_dims,  global_strides, box_dim, element_strides, CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B, CU_TENSOR_MAP_L2_PROMOTION_NONE, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
  void* args[] = {
    (void*) &ptr_a,
    (void*) &ptr_sfa,
    (void*) &ptr_b,
    (void*) &ptr_sfb,
    (void*) &ptr_c,
    (void*) &desc_a
  };
  cudaLaunchKernelExC(&launch_config, func, args);
}

#ifdef BUILD_PYTORCH
void run(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c) {
  int m = a.size(0);
  int k = 2 * a.size(1);
  int l = a.size(2);

  TORCH_CHECK(k == STATIC_K, "K must match static shape");
  TORCH_CHECK(m % (BLOCK_M * LOCAL_M) == 0, "m must divide blocking evenly");
  TORCH_CHECK(k % (BLOCK_K * SPLIT_K) == 0, "k must divide blocking evenly");

  TORCH_CHECK(a.dim() == 3, "A must be 3d");
  TORCH_CHECK(a.stride(0) == STRIDE_A_M, "M stride A must match static shape");
  TORCH_CHECK(a.stride(1) == 1, "K stride in A must be 1");
  TORCH_CHECK(a.stride(2) == STRIDE_A_L, "L stride A must match static shape");

  TORCH_CHECK(b.dim() == 3, "B must be 3d");
  TORCH_CHECK(b.stride(1) == 1, "K stride in B must be 1");
  TORCH_CHECK(b.stride(2) == STRIDE_B_L, "L stride B must match static shape");
  
  TORCH_CHECK(sfa.dim() == 6, "SFA must be 6d");
  TORCH_CHECK(sfa.stride(0) == 16, "M0 stride in SFA must be 16");
  TORCH_CHECK(sfa.stride(1) == 4, "M1 stride in SFA must be 4");
  TORCH_CHECK(sfa.stride(2) == STRIDE_SFA_M, "M2 stride in SFA must match static shape");
  TORCH_CHECK(sfa.stride(3) == 1, "K0 stride in SFA must be 1");
  TORCH_CHECK(sfa.stride(4) == STRIDE_SFA_K, "K1 stride in SFA must match static shape");
  TORCH_CHECK(sfa.stride(5) == STRIDE_SFA_L, "L stride in SFA must match static shape");

  TORCH_CHECK(sfb.dim() == 6, "SFB must be 6d");
  TORCH_CHECK(sfb.stride(0) == 16, "M0 stride in SFB must be 16");
  TORCH_CHECK(sfb.stride(1) == 4, "M1 stride in SFB must be 4");
  TORCH_CHECK(sfb.stride(3) == 1, "K0 stride in SFA must be 1");
  TORCH_CHECK(sfb.stride(4) == STRIDE_SFB_K, "K1 stride in SFA must match static shape");
  TORCH_CHECK(sfb.stride(5) == STRIDE_SFB_L, "L stride in SFA must match static shape");

  TORCH_CHECK(c.dim() == 3, "C must be 3d");
  TORCH_CHECK(c.stride(0) == 1, "M stride in C must be 1");
  TORCH_CHECK(c.stride(2) == STRIDE_C_L, "L stride in C must match static shape");
  launch(
    m,
    l,
    a.data_ptr(),
    sfa.data_ptr(),
    b.data_ptr(),
    sfb.data_ptr(),
    c.data_ptr()
  );
}
#endif
        """
        ],
        cpp_sources=[
            "void run(torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor);"
        ],
        name="inline_module_" + str(abs(hash(key))),
        functions=["run"],
        extra_cflags=["-DBUILD_PYTORCH", "-O3"] + defines,
        extra_cuda_cflags=[
            "-DBUILD_PYTORCH",
            "--resource-usage",
            "-gencode=arch=compute_100a,code=sm_100a",
            "-O3",
        ]
        + defines,
        extra_ldflags=["-LcublasLt", "-lcuda"],
    )


PREPOPULATE = {
    (
        ((7168, 8192, 1), (8192, 1, 58720256)),
        ((128, 8192, 1), (8192, 1, 1048576)),
        ((32, 4, 56, 4, 256, 1), (16, 4, 131072, 1, 512, 7340032)),
        ((32, 4, 1, 4, 256, 1), (16, 4, 131072, 1, 512, 131072)),
        ((7168, 1, 1), (1, 1, 7168)),
    ): {"LOCAL_M": 1, "USE_PDL": 1, "CLUSTER_M": None, "SPLIT_K": 2, "STAGE_COUNT": 6},
    (
        ((4096, 3584, 8), (3584, 1, 14680064)),
        ((128, 3584, 8), (3584, 1, 458752)),
        ((32, 4, 32, 4, 112, 8), (16, 4, 57344, 1, 512, 1835008)),
        ((32, 4, 1, 4, 112, 8), (16, 4, 57344, 1, 512, 57344)),
        ((4096, 1, 8), (1, 1, 4096)),
    ): {"LOCAL_M": 2, "USE_PDL": 1, "CLUSTER_M": None, "SPLIT_K": 1, "STAGE_COUNT": 6},
    (
        ((7168, 1024, 4), (1024, 1, 7340032)),
        ((128, 1024, 4), (1024, 1, 131072)),
        ((32, 4, 56, 4, 32, 4), (16, 4, 16384, 1, 512, 917504)),
        ((32, 4, 1, 4, 32, 4), (16, 4, 16384, 1, 512, 16384)),
        ((7168, 1, 4), (1, 1, 7168)),
    ): {"LOCAL_M": 2, "USE_PDL": 1, "CLUSTER_M": None, "SPLIT_K": 1, "STAGE_COUNT": 6},
}
kernels = {}
for k, tunable in PREPOPULATE.items():
    kernels[k] = compile_kernel(k, tunable)
    kernels[k].run(*[torch.empty_strided(size, stride, dtype=dtype, device='cuda') for (size, stride), dtype in zip(k, [torch.float4_e2m1fn_x2, torch.float4_e2m1fn_x2, torch.float8_e4m3fnuz, torch.float8_e4m3fnuz, torch.float16])])
torch.cuda.synchronize()


def custom_kernel(
    data: input_t,
) -> output_t:
    a, b, _, _, sfa_permuted, sfb_permuted, c = data
    args = a, b, sfa_permuted, sfb_permuted, c
    key = kernel_key(*args)
    if key not in kernels:
        kernels[key] = compile_kernel(
            key, heuristic(a, b, sfa_permuted, sfb_permuted, c)
        )
    kernels[key].run(*args)
    return c
scrolls · 1002 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 109661.

⋯ 147 unchanged lines
}) CLO(: "r"((unsigned int) __cvta_generic_to_shared(barrier)), "r"(barrier_wait_phase));
}
+ __device__ float e2m1x2_to_float(__nv_fp4x2_e2m1 value, int subbyte_idx) {
+ __half2_raw values = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<const __nv_fp4x2_storage_t&>(value), __NV_E2M1);
+ return __half2float(reinterpret_cast<const __half&>(subbyte_idx == 0 ? values.x : values.y));
+ }
+
+ __device__ __half2 e2m1x2_to_half2(__nv_fp4x2_e2m1 value) {
+ __half2_raw values = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<const __nv_fp4x2_storage_t&>(value), __NV_E2M1);
+ return reinterpret_cast<const __half2&>(values);
+ }
+
+ __device__ __half2 e2m1x2_to_half2(unsigned char value) {
+ __half2_raw values = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<const __nv_fp4x2_storage_t&>(value), __NV_E2M1);
+ return reinterpret_cast<const __half2&>(values);
+ }
+
+ __device__ float e4m3_to_float(unsigned char value) {
+ __half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);
+ return __half2float(reinterpret_cast<const __half&>(half_value));
+ }
+
+ __device__ float e4m3_to_float(__nv_fp8_e4m3 value) {
+ __half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);
+ return __half2float(reinterpret_cast<const __half&>(half_value));
+ }
+
+ __device__ __half e4m3_to_half(__nv_fp8_e4m3 value) {
+ __half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);
+ return reinterpret_cast<const __half&>(half_value);
+ }
+
+ __device__ __half e4m3_to_half(unsigned char value) {
+ return e4m3_to_half(reinterpret_cast<const __nv_fp8_e4m3&>(value));
+ }
+
#ifdef TIMING
#define TIME_EVENT(event_id) do { \
if (is_timer) { \
⋯ 6 unchanged lines
#define TIME_EVENT(x) do {} while(0)
#endif
+ struct TileIterator {
+ int start_m, start_k;
+ int num_tiles;
+ int total_k_tiles;
+ int total_m_tiles;
+ int total_sms;
+ int index;
+ // simple, we have N tiles
+ // we need to be able to decode them
+ // locally, we need to know this tile and the next tile
+ // we need to be able to chunk
+ //
+ };
+
+ __device__ int get_chunked_local_pos(int num, int idx, int total) {
+ int count = total / num;
+ int rest = total % num;
+ int pos = min(idx, rest) + idx * count;
+ return pos;
+ }
+
__launch_bounds__(256, 1)
__global__ void kernel(
const __nv_fp4x2_e2m1 * __restrict__ ptr_a,
⋯ 41 unchanged lines
__attribute__((aligned(128))) __shared__ __nv_fp8_e4m3 smem_sfa[STAGE_COUNT][BLOCK_M * (BLOCK_K / 16)];
// 9 * 16, 2 * 9 * 17
__attribute__((aligned(1024))) __shared__ __nv_fp4x2_e2m1 smem_a[STAGE_COUNT][BLOCK_M * (BLOCK_K / 2)];
+ __attribute__((aligned(16))) __shared__ __half result[LOCAL_M * BLOCK_M];
int warp_id = threadIdx.x / 32;
int wg_id = warp_id / 4;
⋯ 36 unchanged lines
int cutover_k = (STAGE_COUNT / LOCAL_M) * BLOCK_K;
if (cutover_k > k) cutover_k = k;
- cutover_k = 0;
+ // cutover_k = 0;
if (is_load) {
⋯ 51 unchanged lines
}
}
ASMV(bar.arrive 1, 160;) CLO();
+ #if 0
// start computing
+ int read_stage = 0;
+ int read_phase = 0;
+ int local_m = threadIdx.x;
+ float acc[LOCAL_M][8] = {0};
+ int idx = 0;
+ for (int block_k = 0; block_k < k; block_k += BLOCK_K) {
+ #pragma unroll
+ for (int block_m = 0; block_m < LOCAL_M; block_m++) {
+
+ idx += 1;
+ if (idx % 2 == 0) {
+ read_stage += 1;
+ if (read_stage == STAGE_COUNT) {
+ read_stage = 0;
+ read_phase ^= 1;
+ }
+ continue;
+ }
+
+ barrier_wait(&barrier_smem_full[read_stage], read_phase);
+
+ // each thread reads one column of K values from A and B
+ // that means 256 / 32 = 8 elems, or 4 bytes.
+ // seems tunable, too.
+ // if i wanted 16 bytes, i'd have 8 threads along K
+ // so four along M, so B is only reused 8 times. that might be fine.
+
+ // load B: 32 elems, dequantize to fp16
+ // load sfB: 2 elem, dequantize to fp16
+ // load A: 32 elems in K, 8 elems in M, dequantize to fp16
+ // load sfB: 2 elem, dequantize to fp16
+ // update each accumulator
+
+ // do cutover: take every other one to tc or to local
+ int partition_k = (lane_id % 8) * 16;
+
+ // each thread loads 16 bytes == 32 elements
+ // for blocking 256, there are 8 (= 256 / 32) threads in the k dimension then
+ // this means there are 4 threads in the m dimension
+ // this means each thread then has 32 / 4 = 8 elements in the m dimension
+ // let them be contiguous
+ uint4 vecb = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k]);
+ __half2 vecbv[sizeof(vecb)];
+ #pragma unroll
+ for (int i = 0; i < sizeof(vecb); i++) {
+ vecbv[i] = e2m1x2_to_half2(reinterpret_cast<unsigned char*>(&vecb)[i]);
+ }
+
+ uchar2 vecsfb = *reinterpret_cast<uchar2*>(&smem_sfb[read_stage][16 * (partition_k / 32) + (partition_k / 8) % 4]);
+
+ __half value_sfb[2] = {e4m3_to_half(vecsfb.x), e4m3_to_half(vecsfb.y)};
+
+ #pragma unroll
+ for (int the_m = 0; the_m < 8; the_m += 1) {
+ int local_m = warp_id * 32 + the_m + (lane_id / 8) * 8;
+ uint4 veca = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k)]); // this has 32 fp4 elems
+ __half2 vecav[sizeof(veca)];
+ #pragma unroll
+ for (int i = 0; i < sizeof(veca); i++) {
+ vecav[i] = e2m1x2_to_half2(reinterpret_cast<unsigned char*>(&veca)[i]);
+ }
+
+ uchar2 vecsfa = *reinterpret_cast<uchar2*>(&smem_sfa[read_stage][(partition_k / 8 / 4) * 128 * 4 + ((partition_k / 8) % 4) + (local_m % 32) * 16 + (local_m / 32) * 4]);
+ __half value_sfa[2] = {e4m3_to_half(vecsfa.x), e4m3_to_half(vecsfa.y)};
+
+ __half2 local_acc[2];
+ local_acc[0] = __hmul2(vecav[0], vecbv[0]);
+ local_acc[1] = __hmul2(vecav[8], vecbv[8]);
+ #pragma unroll
+ for (int i = 1; i < 8; i++) {
+ local_acc[0] = __hfma2(vecav[i+0], vecbv[i+0], local_acc[0]);
+ local_acc[1] = __hfma2(vecav[i+8], vecbv[i+8], local_acc[1]);
+ }
+ acc[block_m][the_m] += __half2float(__hmul(__hadd(local_acc[0].x, local_acc[0].y), __hmul(value_sfa[0], value_sfb[0])));
+ acc[block_m][the_m] += __half2float(__hmul(__hadd(local_acc[1].x, local_acc[1].y), __hmul(value_sfa[1], value_sfb[1])));
+ }
+
+
+ #if 0
+
+ // better maybe: each warp still has 32 elements, but each thread in the warp does, too, and they are split along K
+ #pragma unroll
+ for (int partition_k = 0; partition_k < BLOCK_K / 2; partition_k += 2 * sizeof(uint4)) {
+ uint4 veca = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k)]);
+ uint4 vecb = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k]);
+ uchar4 vecsfa = *reinterpret_cast<uchar4*>(&smem_sfa[read_stage][(partition_k / 8 / 4) * 128 * 4 + (local_m % 32) * 16 + (local_m / 32) * 4]);
+ uchar4 vecsfb = *reinterpret_cast<uchar4*>(&smem_sfb[read_stage][4 * partition_k / 8]);
+
+ uint4 vecc = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k + sizeof(uint4))]);
+ uint4 vecd = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k + sizeof(uint4)]);
+
+ auto dot_sf = [](unsigned char sfa, unsigned char sfb, unsigned int a0, unsigned int a1, unsigned int b0, unsigned int b1) {
+ uchar4 veca = reinterpret_cast<const uchar4&>(a0);
+ uchar4 vecb = reinterpret_cast<const uchar4&>(b0);
+ __half2 value_a = e2m1x2_to_half2(veca.x);
+ __half2 value_b = e2m1x2_to_half2(vecb.x);
+ __half2 acc = __hmul2(value_a, value_b);
+ value_a = e2m1x2_to_half2(veca.y);
+ value_b = e2m1x2_to_half2(vecb.y);
+ acc = __hfma2(value_a, value_b, acc);
+ value_a = e2m1x2_to_half2(veca.z);
+ value_b = e2m1x2_to_half2(vecb.z);
+ acc = __hfma2(value_a, value_b, acc);
+ value_a = e2m1x2_to_half2(veca.w);
+ value_b = e2m1x2_to_half2(vecb.w);
+ acc = __hfma2(value_a, value_b, acc);
+ veca = reinterpret_cast<const uchar4&>(a1);
+ vecb = reinterpret_cast<const uchar4&>(b1);
+ value_a = e2m1x2_to_half2(veca.x);
+ value_b = e2m1x2_to_half2(vecb.x);
+ acc = __hfma2(value_a, value_b, acc);
+ value_a = e2m1x2_to_half2(veca.y);
+ value_b = e2m1x2_to_half2(vecb.y);
+ acc = __hfma2(value_a, value_b, acc);
+ value_a = e2m1x2_to_half2(veca.z);
+ value_b = e2m1x2_to_half2(vecb.z);
+ acc = __hfma2(value_a, value_b, acc);
+ value_a = e2m1x2_to_half2(veca.w);
+ value_b = e2m1x2_to_half2(vecb.w);
+ acc = __hfma2(value_a, value_b, acc);
+
+ __half sum = __hadd(acc.x, acc.y);
+ float value_sfa = e4m3_to_float(sfa);
+ float value_sfb = e4m3_to_float(sfb);
+ return __half2float(sum) * value_sfa * value_sfb;
+ };
+
+ acc[block_m] += dot_sf(vecsfa.x, vecsfb.x, veca.x, veca.y, vecb.x, vecb.y);
+ acc[block_m] += dot_sf(vecsfa.y, vecsfb.y, veca.z, veca.w, vecb.z, vecb.w);
+ acc[block_m] += dot_sf(vecsfa.z, vecsfb.z, vecc.x, vecc.y, vecd.x, vecd.y);
+ acc[block_m] += dot_sf(vecsfa.w, vecsfb.w, vecc.z, vecc.w, vecd.z, vecd.w);
+ }
+ #endif
+
+ ASMV(bar.sync 2, 128;) CLO();
+ if (threadIdx.x == 0) {
+ ASMV(mbarrier.arrive.shared::cta.b64 _, [%0], 1;) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_empty[read_stage])));
+ }
+
+ read_stage += 1;
+ if (read_stage == STAGE_COUNT) {
+ read_stage = 0;
+ read_phase ^= 1;
+ }
+ }
+ }
+
+ for (int block_m = 0; block_m < LOCAL_M; block_m++) {
+ for (int i = 0; i < 8; i++) {
+ // reduce the accumulator
+ float reduced_acc = acc[block_m][i];
+ reduced_acc += __shfl_xor_sync((unsigned int) -1, reduced_acc, 1);
+ reduced_acc += __shfl_xor_sync((unsigned int) -1, reduced_acc, 2);
+ reduced_acc += __shfl_xor_sync((unsigned int) -1, reduced_acc, 4);
+ if (threadIdx.x % 8 == 0) {
+ result[block_m * BLOCK_M + 32 * warp_id + i + 8 * (lane_id / 8)] = __float2half(reduced_acc);
+ }
+ }
+ }
+ // __syncthreads();
+ // for (int i = threadIdx.x * 8; i < LOCAL_M * BLOCK_M; i += 128 * 8) {
+ // if (SPLIT_K > 1) {
+ // ASM({
+ // .reg .v4 .f16x2 reg;
+ // ld.shared.v4.b32 reg, [%0];
+ // red.relaxed.gpu.global.add.noftz.v4.f16x2 [%1], reg;
+ // }) CLO(: "r"((unsigned int) __cvta_generic_to_shared(result + i)), "l"(ptr_c + i));
+ // } else {
+ // reinterpret_cast<uint4&>(ptr_c[i]) = reinterpret_cast<uint4&>(result[i]);
+ // }
+ // }
+ #endif
} else if (is_mma) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], 512;" :: "r"((unsigned int) __cvta_generic_to_shared(&tmem_addr)));
int read_stage = 0;
int read_phase = 0;
+ int idx = 0;
+ int has_zerod[LOCAL_M] = {0};
for (int block_k = 0; block_k < k; block_k += BLOCK_K) {
#pragma unroll
for (int block_m = 0; block_m < LOCAL_M; block_m++) {
- int enable_input_d = block_k != 0;
+ // idx += 1;
+ // if (idx % 2 == 1) {
+ // read_stage += 1;
+ // if (read_stage == STAGE_COUNT) {
+ // read_stage = 0;
+ // read_phase ^= 1;
+ // }
+ // continue;
+ // }
+
+ int enable_input_d = has_zerod[block_m];
+ has_zerod[block_m] = 1;
+ // int enable_input_d = block_k != 0;
+
TIME_EVENT(5);
barrier_wait(&barrier_smem_full[read_stage], read_phase);
TIME_EVENT(6);
⋯ 235 unchanged lines
// asm volatile("griddepcontrol.launch_dependents;");
// load tmem
- __attribute__((aligned(16))) __shared__ __half result[LOCAL_M * BLOCK_M];
if (is_load) {
#pragma unroll
⋯ 1 unchanged lines
int tmem_d = block_m * 16;
float ld_result;
ASMV(tcgen05.ld.sync.aligned.32x32b.x1.b32 {%0}, [%1];) CLO("=f"(ld_result) : "r"(tmem_d));
+ // result[block_m * BLOCK_M + threadIdx.x] = __hadd(result[block_m * BLOCK_M + threadIdx.x], __float2half(ld_result));
result[block_m * BLOCK_M + threadIdx.x] = __float2half(ld_result);
}
}
scrolls · 308 diff lines total

Best evidence level for this revision: reported

JSON