Skip to content
KernelIndex
Search⌘K

submission 88871

v0i0 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission-19-06-48-45.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-88871?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
27.6µs
#131 of 678
2025-11-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b720e4fd2112f0fce84acdbae3f0e35eb906adb6bd654701ea7ac6c0887485fe
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() {
fp8const __nv_fp8_e4m3 * __restrict__ ptr_sfa,
mbarrieragain: mbarrier.try_wait.parity.shared::cta.b64 done, [%0], %1, 0x100000;
shared-memory__attribute__((aligned(128))) __shared__ __nv_fp8_e4m3 smem_sfb[STAGE_COUNT][4 * BLOCK_K / 16];
split-kint split_k
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;
vector-width = uint4using vec_type = uint4;

Kernel source

submission-19-06-48-45.py403 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline

module = load_inline(cuda_sources=["""
#ifdef BUILD_PYTORCH
#include <c10/cuda/CUDAStream.h>
#endif

#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <iostream>
#include <vector>

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

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

const int BLOCK_M = 128;
const int STAGE_COUNT = 8;
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__ 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));
}

template<int LOCAL_M>
__launch_bounds__(128 + 64, 1)
__global__ void kernel(
  const int m,
  const int k,
  const int l,
  const __nv_fp4x2_e2m1 * __restrict__ ptr_a,
  const int stride_a_m,
  const int stride_a_l,
  const __nv_fp8_e4m3 * __restrict__ ptr_sfa,
  const int stride_sfa_m,
  const int stride_sfa_k,
  const int stride_sfa_l,
  const __nv_fp4x2_e2m1 * __restrict__ ptr_b,
  const int stride_b_l,
  const __nv_fp8_e4m3 * __restrict__ ptr_sfb,
  const int stride_sfb_k,
  const int stride_sfb_l,
  __half * __restrict__ ptr_c,
  const int stride_c_m,
  const int stride_c_l,
  int split_k
) {
  // 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;

  ptr_a += blockIdx.x * BLOCK_M * stride_a_m * LOCAL_M;
  ptr_sfa += blockIdx.x * (BLOCK_M / 128) * stride_sfa_m * LOCAL_M;
  ptr_c += blockIdx.x * BLOCK_M * stride_c_m * LOCAL_M;

  ptr_c += threadIdx.x * stride_c_m;

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

  const int stride_k_rest = 2 * 9 * 16 * (BLOCK_M / 8 + 1);
  __attribute__((aligned(128))) __shared__ __nv_fp8_e4m3 smem_sfb[STAGE_COUNT][4 * BLOCK_K / 16];
  __attribute__((aligned(128))) __shared__ __nv_fp4x2_e2m1 smem_b[STAGE_COUNT][BLOCK_K / 2];
  // 9 * 16, 2 * 9 * 17
  __attribute__((aligned(128))) __shared__ __nv_fp4x2_e2m1 smem_a[STAGE_COUNT][stride_k_rest * (BLOCK_K / 2 / 32)];
  __attribute__((aligned(128))) __shared__ __nv_fp8_e4m3 smem_sfa[STAGE_COUNT][BLOCK_M * (BLOCK_K / 16)];

  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;
  int lane_id = threadIdx.x % 32;
  __shared__ int tmem_addr;
  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)));
  }
  __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])));
  }
  __syncthreads();

  auto make_smem_desc = [](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) {
    auto mde = [](unsigned int x) { return (unsigned long long) ((x & 0x3FFFF) >> 4); };
    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);
  };

  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);
  ptr_b += ptr_b_k_offset;

  using sf_vec_type = unsigned int;

  auto offset_a = [&](int m, int k) {
    return k % 16 + (m % 8) * 16 + ((k / 16) % 2) * 9 * 16 + (m / 8) * 2 * 9 * 16 + (k / 32) * stride_k_rest;
  };

  if (is_load) {
    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 < k; block_k += BLOCK_K) {
      for (int block_m = 0; block_m < LOCAL_M; block_m++) {
        if (load_phase != 2) {
          barrier_wait(&barrier_smem_empty[load_stage], load_phase);
        }
        #pragma unroll
        for (int partition_m = 0; partition_m < BLOCK_M; partition_m += threads_in_m) {
          copy_b_vec(&smem_a[load_stage][offset_a(partition_m + ptr_a_m_offset, ptr_a_k_offset)], &ptr_a[block_k / 2 + partition_m * stride_a_m + block_m * BLOCK_M * stride_a_m]);
        }
        // load b
        if (ptr_b_n_offset == 0) {
          copy_b_vec(&smem_b[load_stage][ptr_b_k_offset], &ptr_b[block_k / 2]);
        }
        // load sf a
        // new storage format is 4K x 32M4 x 4M1 x RK
        // 4 * 32 * 4 / 16 = 32, RK = BLOCK_K / 16 / 4
        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]);
        }
        // load sf b
        #pragma unroll
        for (int partition_k = threadIdx.x * sizeof(sf_vec_type); partition_k < BLOCK_K / 16; partition_k += BLOCK_M * sizeof(sf_vec_type)) {
          // 128b loads
          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;
        }
      }
    }
  } else if (is_mma) {
    int read_stage = 0;
    int read_phase = 0;
    for (int block_k = 0; block_k < k; block_k += BLOCK_K) {
      for (int block_m = 0; block_m < LOCAL_M; block_m++) {
        int enable_input_d = block_k != 0;
        barrier_wait(&barrier_smem_full[read_stage], read_phase);
        #pragma unroll
        for (int partition_k = 0; partition_k < BLOCK_K / 2; partition_k += 2 * sizeof(uint4)) {
          int tmem_d = block_m * 16;
          int tmem_sfa = 32 + 8 * (partition_k / 2 / sizeof(uint4)) + read_stage * 8 * (BLOCK_K / 2 / 2 / sizeof(uint4));
          int tmem_sfb = tmem_sfa + 4;
          barrier_wait(&barrier_smem_full[read_stage], read_phase);
          unsigned long long smem_desc_sfa = make_smem_desc(&smem_sfa[read_stage][(partition_k / 32) * 128 * 4], 16, 8 * 16, 0, 0, 0);
          unsigned long long smem_desc_sfb = make_smem_desc(&smem_sfb[read_stage][(partition_k / 32) * 16], 16, 8 * 16, 0, 0, 0);
          unsigned long long smem_desc_a = make_smem_desc(&smem_a[read_stage][offset_a(0, partition_k)], 16 * 9, 16 * 9 * 2, 0, 0, 0);
          unsigned long long smem_desc_b = make_smem_desc(&smem_b[read_stage][partition_k], 16, 0, 0, 0, 0);
          unsigned int insn_desc = (1u << 7) | (1u << 10) | (1u << 17) | (1u << 27);  // e2m1, e2m1, N=8, M=128
          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));
          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);
  }

  __syncthreads();

  asm volatile("griddepcontrol.launch_dependents;");
  // load tmem
  float result[LOCAL_M];
  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] = ld_result;
    }
  }
  __syncthreads();
  if (is_mma) {
    asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 0, 512;");
  }
  if (is_load) {
    #pragma unroll
    for (int block_m = 0; block_m < LOCAL_M; block_m++) {
      ptr_c[block_m * BLOCK_M * stride_c_m] = __float2half(result[block_m]);
    }
  }
}

void launch(
  int m,
  int k,
  int l,
  void * ptr_a,
  int stride_a_m,
  int stride_a_l,
  void * ptr_sfa,
  int stride_sfa_m,
  int stride_sfa_k,
  int stride_sfa_l,
  void * ptr_b,
  int stride_b_l,
  void * ptr_sfb,
  int stride_sfb_k,
  int stride_sfb_l,
  void * ptr_c,
  int stride_c_m,
  int stride_c_l,
  cudaStream_t stream
) {
  int total_blocks = (m / BLOCK_M) * l;
  int local_m = 1;
  int split_k = 1;
  if (total_blocks > 128 && (m % (BLOCK_M * 2)) == 0) {
    local_m = 2;
  }
  // if (total_blocks < 64) {
  //   split_k = 2;
  // }
  dim3 grid(split_k * m / BLOCK_M / local_m, l);
  dim3 block(128 + 64);
  void* func = local_m == 2 ? (void*) kernel<2> : (void*) kernel<1>;

  cudaLaunchConfig_t launch_config;
  launch_config.blockDim = block;
  launch_config.gridDim = grid;
  launch_config.stream = stream;
  launch_config.dynamicSmemBytes = 0;
  cudaLaunchAttribute attrs[16];
  attrs[0].id = cudaLaunchAttributeClusterDimension;
  attrs[0].id = cudaLaunchAttributeIgnore;
  attrs[0].val.clusterDim.x = 1;
  attrs[0].val.clusterDim.y = 1;
  attrs[0].val.clusterDim.z = 1;
  attrs[1].id = cudaLaunchAttributeProgrammaticStreamSerialization;
  attrs[1].val.programmaticStreamSerializationAllowed = 1;  // 1
  launch_config.attrs = attrs;
  launch_config.numAttrs = 2;
  void* args[] = {
    (void*) &m,
    (void*) &k,
    (void*) &l,
    (void*) &ptr_a,
    (void*) &stride_a_m,
    (void*) &stride_a_l,
    (void*) &ptr_sfa,
    (void*) &stride_sfa_m,
    (void*) &stride_sfa_k,
    (void*) &stride_sfa_l,
    (void*) &ptr_b,
    (void*) &stride_b_l,
    (void*) &ptr_sfb,
    (void*) &stride_sfb_k,
    (void*) &stride_sfb_l,
    (void*) &ptr_c,
    (void*) &stride_c_m,
    (void*) &stride_c_l,
    (void*) &split_k
  };
  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(a.dim() == 3, "A must be 3d");
  TORCH_CHECK(a.stride(1) == 1, "K stride in A must be 1");
  TORCH_CHECK(b.dim() == 3, "B must be 3d");
  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(3) == 1, "K0 stride in SFA must be 1");
  TORCH_CHECK(sfb.dim() == 6, "SFB must be 6d");
  TORCH_CHECK(sfb.stride(3) == 1, "K0 stride in SFA must be 1");
  TORCH_CHECK(c.dim() == 3, "C must be 3d");
  TORCH_CHECK(b.stride(1) == 1, "K stride in B must be 1");
  // 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");
  launch(
    m,
    k,
    l,
    a.data_ptr(),
    a.stride(0), 
    a.stride(2), 
    sfa.data_ptr(),
    sfa.stride(2),
    sfa.stride(4),
    sfa.stride(5),
    b.data_ptr(),
    b.stride(2), 
    sfb.data_ptr(),
    // sfb.stride(2),
    sfb.stride(4),
    sfb.stride(5),
    c.data_ptr(),
    c.stride(0),
    c.stride(2),
    c10::cuda::getCurrentCUDAStream().stream() 
  );
}
#endif
    """],
    cpp_sources=["void run(torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor);"],
    name="inline_module",
    functions=["run"],
    extra_cflags=["-DBUILD_PYTORCH", "-O3"],
    extra_cuda_cflags=["-DBUILD_PYTORCH", "--resource-usage", "-gencode=arch=compute_100a,code=sm_100a", "-O3"],
    extra_ldflags=["-LcublasLt"],
)

# TODO
PRECOMPILE_AND_TUNE_STAGES = [
  (16384, 1, 7168),
  (7168, 8, 4096),
  (2048, 4, 7168),
]

def custom_kernel(
    data: input_t,
) -> output_t:
    a, b, _, _, sfa_permuted, sfb_permuted, c = data
    module.run(a, b, sfa_permuted, sfb_permuted, c)
    return c
scrolls · 403 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 74602.

⋯ 11 unchanged lines
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <iostream>
+ #include <vector>
+ #define ASM(x...) asm(#x
+ #define ASMV(x...) asm volatile(#x
+ #define CLO(x...) : x)
+
+ __device__ int semaphore[128][32] = {0};
+
const int BLOCK_M = 128;
const int STAGE_COUNT = 8;
const int BLOCK_K = 256;
- __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));
+ template<int N>
+ __device__ void cp_async_wait_group() {
+ asm("cp.async.wait_group %0;" :: "n"(N) : "memory");
}
- __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__ 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__ __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__ 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__ 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__ 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__ 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__ 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__ __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);
- }
-
- template<class V, class D, class S>
- __device__ void copy(D* dst, const S* src) {
- *reinterpret_cast<V*>(dst) = *reinterpret_cast<const V*>(src);
- }
-
- template<int N>
- __device__ void cp_async_wait_group() {
- #define CAWG_COND(x) \
- if constexpr (N == x) { \
- asm("cp.async.wait_group " #x ";" ::: "memory"); \
- }
- CAWG_COND(16) CAWG_COND(15) CAWG_COND(14) CAWG_COND(13)
- CAWG_COND(12) CAWG_COND(11) CAWG_COND(10) CAWG_COND( 9)
- CAWG_COND( 8) CAWG_COND( 7) CAWG_COND( 6) CAWG_COND( 5)
- CAWG_COND( 4) CAWG_COND( 3) CAWG_COND( 2) CAWG_COND( 1)
- CAWG_COND( 0)
- #undef CAWG_COND
- }
-
- __launch_bounds__(BLOCK_M, 1)
+ template<int LOCAL_M>
+ __launch_bounds__(128 + 64, 1)
__global__ void kernel(
const int m,
const int k,
⋯ 12 unchanged lines
const int stride_sfb_l,
__half * __restrict__ ptr_c,
const int stride_c_m,
- const int stride_c_l
+ const int stride_c_l,
+ int split_k
) {
// let's do the "real" blocking: 128x64 in M and K, block along M and L
ptr_a += blockIdx.y * stride_a_l;
⋯ 2 unchanged lines
ptr_sfb += blockIdx.y * stride_sfb_l;
ptr_c += blockIdx.y * stride_c_l;
- ptr_a += blockIdx.x * BLOCK_M * stride_a_m;
- ptr_sfa += blockIdx.x * (BLOCK_M / 128) * stride_sfa_m;
- ptr_c += blockIdx.x * BLOCK_M * stride_c_m;
+ ptr_a += blockIdx.x * BLOCK_M * stride_a_m * LOCAL_M;
+ ptr_sfa += blockIdx.x * (BLOCK_M / 128) * stride_sfa_m * LOCAL_M;
+ ptr_c += blockIdx.x * BLOCK_M * stride_c_m * LOCAL_M;
+ ptr_c += threadIdx.x * stride_c_m;
+
// have 128 threads, each computing one output element
- int local_m = threadIdx.x;
// SF: MN, K -> addr is ((32, 4, RM), (16, 4, RK)):((16, 4, ?), (0, 1, ?))
- ptr_c += local_m * stride_c_m;
- __attribute__((aligned(16))) __shared__ __nv_fp4x2_e2m1 smem_a[STAGE_COUNT][BLOCK_M][BLOCK_K / 2 + 16];
- __attribute__((aligned(16))) __shared__ __nv_fp8_e4m3 smem_sfa[STAGE_COUNT][BLOCK_M][BLOCK_K / 16];
- __attribute__((aligned(16))) __shared__ __nv_fp4x2_e2m1 smem_b[STAGE_COUNT][BLOCK_K / 2 + 16];
- __attribute__((aligned(16))) __shared__ __nv_fp8_e4m3 smem_sfb[STAGE_COUNT][BLOCK_K / 16];
+ const int stride_k_rest = 2 * 9 * 16 * (BLOCK_M / 8 + 1);
+ __attribute__((aligned(128))) __shared__ __nv_fp8_e4m3 smem_sfb[STAGE_COUNT][4 * BLOCK_K / 16];
+ __attribute__((aligned(128))) __shared__ __nv_fp4x2_e2m1 smem_b[STAGE_COUNT][BLOCK_K / 2];
+ // 9 * 16, 2 * 9 * 17
+ __attribute__((aligned(128))) __shared__ __nv_fp4x2_e2m1 smem_a[STAGE_COUNT][stride_k_rest * (BLOCK_K / 2 / 32)];
+ __attribute__((aligned(128))) __shared__ __nv_fp8_e4m3 smem_sfa[STAGE_COUNT][BLOCK_M * (BLOCK_K / 16)];
+ 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;
+ int lane_id = threadIdx.x % 32;
+ __shared__ int tmem_addr;
+ 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)));
+ }
+ __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])));
+ }
+ __syncthreads();
+
+ auto make_smem_desc = [](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) {
+ auto mde = [](unsigned int x) { return (unsigned long long) ((x & 0x3FFFF) >> 4); };
+ 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);
+ };
+
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;
⋯ 8 unchanged lines
using sf_vec_type = unsigned int;
- float acc_out = 0;
- int load_stage = 0;
- int load_block_k = 0;
- int read_stage = 0;
- auto load_step = [&]() {
- if (load_block_k < k) {
- auto copy_a_vec = [](auto dst, auto src) {
- asm("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"((unsigned int)__cvta_generic_to_shared(dst)), "l"(src) : "memory");
- };
+ auto offset_a = [&](int m, int k) {
+ return k % 16 + (m % 8) * 16 + ((k / 16) % 2) * 9 * 16 + (m / 8) * 2 * 9 * 16 + (k / 32) * stride_k_rest;
+ };
- auto copy_b_vec = [](auto dst, auto src) {
- asm("cp.async.ca.shared.global [%0], [%1], 16;" :: "r"((unsigned int)__cvta_generic_to_shared(dst)), "l"(src) : "memory");
- };
-
- auto copy_sf_vec = [](auto dst, auto src) {
- asm("cp.async.ca.shared.global [%0], [%1], 4;" :: "r"((unsigned int)__cvta_generic_to_shared(dst)), "l"(src) : "memory");
- };
-
- // load a
- #pragma unroll
- for (int partition_m = 0; partition_m < BLOCK_M; partition_m += threads_in_m) {
- copy_a_vec(&smem_a[load_stage][partition_m + ptr_a_m_offset][ptr_a_k_offset], &ptr_a[load_block_k / 2 + partition_m * stride_a_m]);
- }
- // load b
- if (ptr_b_n_offset == 0) {
- copy_b_vec(&smem_b[load_stage][ptr_b_k_offset], &ptr_b[load_block_k / 2]);
- }
- // load sf a
- for (int partition_m = threadIdx.x; partition_m < BLOCK_M; partition_m += BLOCK_M) {
+ if (is_load) {
+ 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 < k; block_k += BLOCK_K) {
+ for (int block_m = 0; block_m < LOCAL_M; block_m++) {
+ if (load_phase != 2) {
+ barrier_wait(&barrier_smem_empty[load_stage], load_phase);
+ }
#pragma unroll
- for (int partition_k = 0; partition_k < BLOCK_K / 16; partition_k += sizeof(sf_vec_type)) {
+ for (int partition_m = 0; partition_m < BLOCK_M; partition_m += threads_in_m) {
+ copy_b_vec(&smem_a[load_stage][offset_a(partition_m + ptr_a_m_offset, ptr_a_k_offset)], &ptr_a[block_k / 2 + partition_m * stride_a_m + block_m * BLOCK_M * stride_a_m]);
+ }
+ // load b
+ if (ptr_b_n_offset == 0) {
+ copy_b_vec(&smem_b[load_stage][ptr_b_k_offset], &ptr_b[block_k / 2]);
+ }
+ // load sf a
+ // new storage format is 4K x 32M4 x 4M1 x RK
+ // 4 * 32 * 4 / 16 = 32, RK = BLOCK_K / 16 / 4
+ 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]);
+ }
+ // load sf b
+ #pragma unroll
+ for (int partition_k = threadIdx.x * sizeof(sf_vec_type); partition_k < BLOCK_K / 16; partition_k += BLOCK_M * sizeof(sf_vec_type)) {
// 128b loads
- // 256 elems = 16 sfs, 4 contigous, then 32, then 4, 4 sfs = 4B
- copy_sf_vec(&smem_sfa[load_stage][partition_m][partition_k], &ptr_sfa[((load_block_k / 16 + partition_k) % 4) + (load_block_k / 64 + partition_k / (64 / 16)) * stride_sfa_k + (local_m % 32) * 16 + ((local_m / 32) % 4) * 4 + (local_m / 128) * stride_sfa_m]);
+ 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;
+ }
}
- // load sf b
- // 256 elems = 16 scale factors = 1 ldgsts
- #pragma unroll
- for (int partition_k = threadIdx.x * sizeof(sf_vec_type); partition_k < BLOCK_K / 16; partition_k += BLOCK_M * sizeof(sf_vec_type)) {
- // 128b loads
- copy_sf_vec(&smem_sfb[load_stage][partition_k], &ptr_sfb[((load_block_k / 16 + partition_k) % 4) + (load_block_k / 64 + partition_k / (64 / 16)) * stride_sfb_k]);
+ }
+ } else if (is_mma) {
+ int read_stage = 0;
+ int read_phase = 0;
+ for (int block_k = 0; block_k < k; block_k += BLOCK_K) {
+ for (int block_m = 0; block_m < LOCAL_M; block_m++) {
+ int enable_input_d = block_k != 0;
+ barrier_wait(&barrier_smem_full[read_stage], read_phase);
+ #pragma unroll
+ for (int partition_k = 0; partition_k < BLOCK_K / 2; partition_k += 2 * sizeof(uint4)) {
+ int tmem_d = block_m * 16;
+ int tmem_sfa = 32 + 8 * (partition_k / 2 / sizeof(uint4)) + read_stage * 8 * (BLOCK_K / 2 / 2 / sizeof(uint4));
+ int tmem_sfb = tmem_sfa + 4;
+ barrier_wait(&barrier_smem_full[read_stage], read_phase);
+ unsigned long long smem_desc_sfa = make_smem_desc(&smem_sfa[read_stage][(partition_k / 32) * 128 * 4], 16, 8 * 16, 0, 0, 0);
+ unsigned long long smem_desc_sfb = make_smem_desc(&smem_sfb[read_stage][(partition_k / 32) * 16], 16, 8 * 16, 0, 0, 0);
+ unsigned long long smem_desc_a = make_smem_desc(&smem_a[read_stage][offset_a(0, partition_k)], 16 * 9, 16 * 9 * 2, 0, 0, 0);
+ unsigned long long smem_desc_b = make_smem_desc(&smem_b[read_stage][partition_k], 16, 0, 0, 0, 0);
+ unsigned int insn_desc = (1u << 7) | (1u << 10) | (1u << 17) | (1u << 27); // e2m1, e2m1, N=8, M=128
+ 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));
+ 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;
+ }
}
- load_block_k += BLOCK_K;
- load_stage += 1;
- load_stage %= STAGE_COUNT;
}
- asm("cp.async.commit_group;" ::: "memory");
- };
-
- asm volatile("griddepcontrol.wait;" ::: "memory");
- for (int i = 0; i < STAGE_COUNT - 1; i++) {
- load_step();
+ 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);
}
- for (int block_k = 0; block_k < k; block_k += BLOCK_K) {
- // if constexpr (STAGE_COUNT == 4) {
- // asm("cp.async.wait_group 2;" ::: "memory");
- // }
- cp_async_wait_group<STAGE_COUNT-2>();
- __syncthreads();
- load_step();
- // __syncthreads();
- #pragma unroll
- for (int partition_k = 0; partition_k < BLOCK_K / 2; partition_k += sizeof(uint4)) {
- // load at 64b from smem
- uint4 veca = *reinterpret_cast<uint4*>(&smem_a[read_stage][local_m][partition_k]);
- uint4 vecb = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k]);
- uchar2 vecsfa = *reinterpret_cast<uchar2*>(&smem_sfa[read_stage][local_m][partition_k / 8]);
- uchar2 vecsfb = *reinterpret_cast<uchar2*>(&smem_sfb[read_stage][partition_k / 8]);
+ __syncthreads();
- 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);
+ asm volatile("griddepcontrol.launch_dependents;");
+ // load tmem
+ float result[LOCAL_M];
+ if (is_load) {
- __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_out += dot_sf(vecsfa.x, vecsfb.x, veca.x, veca.y, vecb.x, vecb.y);
- acc_out += dot_sf(vecsfa.y, vecsfb.y, veca.z, veca.w, vecb.z, vecb.w);
+ #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] = ld_result;
}
-
- read_stage += 1;
- read_stage %= STAGE_COUNT;
}
- asm volatile("griddepcontrol.launch_dependents;" ::: "memory");
- *ptr_c = __float2half(acc_out);
+ __syncthreads();
+ if (is_mma) {
+ asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 0, 512;");
+ }
+ if (is_load) {
+ #pragma unroll
+ for (int block_m = 0; block_m < LOCAL_M; block_m++) {
+ ptr_c[block_m * BLOCK_M * stride_c_m] = __float2half(result[block_m]);
+ }
+ }
}
void launch(
⋯ 17 unchanged lines
int stride_c_l,
cudaStream_t stream
) {
- dim3 grid(m / BLOCK_M, l);
- dim3 block(BLOCK_M);
- void* func = (void*) kernel;
- int dyn_smem = 0; //200*1024;
- // cudaFuncSetAttribute(func, cudaFuncAttributeNonPortableClusterSizeAllowed, 1);
- // if (dyn_smem > (48 * 1024)) {
- // cudaFuncSetAttribute(func, cudaFuncAttributeMaxDynamicSharedMemorySize, dyn_smem);
+ int total_blocks = (m / BLOCK_M) * l;
+ int local_m = 1;
+ int split_k = 1;
+ if (total_blocks > 128 && (m % (BLOCK_M * 2)) == 0) {
+ local_m = 2;
+ }
+ // if (total_blocks < 64) {
+ // split_k = 2;
// }
+ dim3 grid(split_k * m / BLOCK_M / local_m, l);
+ dim3 block(128 + 64);
+ void* func = local_m == 2 ? (void*) kernel<2> : (void*) kernel<1>;
cudaLaunchConfig_t launch_config;
launch_config.blockDim = block;
launch_config.gridDim = grid;
launch_config.stream = stream;
- launch_config.dynamicSmemBytes = dyn_smem;
+ launch_config.dynamicSmemBytes = 0;
cudaLaunchAttribute attrs[16];
attrs[0].id = cudaLaunchAttributeClusterDimension;
attrs[0].id = cudaLaunchAttributeIgnore;
⋯ 22 unchanged lines
(void*) &stride_sfb_l,
(void*) &ptr_c,
(void*) &stride_c_m,
- (void*) &stride_c_l
+ (void*) &stride_c_l,
+ (void*) &split_k
};
cudaLaunchKernelExC(&launch_config, func, args);
}
⋯ 49 unchanged lines
extra_ldflags=["-LcublasLt"],
)
+ # TODO
+ PRECOMPILE_AND_TUNE_STAGES = [
+ (16384, 1, 7168),
+ (7168, 8, 4096),
+ (2048, 4, 7168),
+ ]
+
def custom_kernel(
data: input_t,
) -> output_t:
scrolls · 450 diff lines total

Best evidence level for this revision: reported

JSON