Skip to content
KernelIndex
Search⌘K

submission 109661

v0i0 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission-27-20-58-58.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-109661?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.8µs
#53 of 678
2025-11-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4efb926fbfe1507103e5c99c6f9427bee23a8dc2e28c10bdb8d3605fb36cde98
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__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-27-20-58-58.py758 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));
}

#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

__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)];

  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();
    // start computing
  } 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;
    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;

        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
  __attribute__((aligned(16))) __shared__ __half result[LOCAL_M * BLOCK_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 * 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 · 758 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 97026.

⋯ 54 unchanged lines
def compile_kernel(key, tunables):
- print(key, tunables)
+ # print(key, tunables)
(
(a_shape, a_stride),
(_, b_stride),
⋯ 19 unchanged lines
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 <cuda.h>
#include <vector>
+ #include <cstdint>
#define ASM(x...) asm(#x
#define ASMV(x...) asm volatile(#x
⋯ 38 unchanged lines
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;
⋯ 4 unchanged lines
}) CLO(: "r"((unsigned int) __cvta_generic_to_shared(barrier)), "r"(barrier_wait_phase));
}
- __launch_bounds__(128 + 64, 1)
+ #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
+
+ __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
+ __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;
⋯ 17 unchanged lines
// 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(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)];
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;
- 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)));
⋯ 7 unchanged lines
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);
};
- 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;
+ 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");
- #pragma unroll 1
- for (int block_k = 0; block_k < k; block_k += BLOCK_K) {
+ 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_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]);
+ 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]);
}
- // 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
#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]);
}
- // 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
+ 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])));
⋯ 4 unchanged lines
}
}
}
+ ASMV(bar.arrive 1, 160;) CLO();
+ // start computing
} 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;
- #pragma unroll 1
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;
+
+ 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], 16 * 9, 16 * 9 * 2, 0, 0, 0);
- unsigned long long smem_desc_b_base = make_smem_desc(&smem_b[read_stage][0], 16, 0, 0, 0, 0);
- unsigned int insn_desc = (1u << 7) | (1u << 10) | (1u << 17) | (1u << 27); // e2m1, e2m1, N=8, M=128
- #pragma unroll
- for (int partition_k = 0; partition_k < BLOCK_K / 2; partition_k += 2 * sizeof(uint4)) {
- 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;
- unsigned long long smem_desc_sfa = smem_desc_sfa_base + mde((partition_k / 32) * 128 * 4);
- unsigned long long smem_desc_sfb = smem_desc_sfb_base + mde((partition_k / 32) * 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 / 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;
- unsigned long long smem_desc_a = smem_desc_a_base + mde(offset_a(0, partition_k));
- unsigned long long smem_desc_b = smem_desc_b_base + mde(partition_k);
- 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;
- }
+ 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;
- 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])));
+ .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) {
⋯ 16 unchanged lines
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
⋯ 8 unchanged lines
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;");
}
⋯ 10 unchanged lines
}
}
}
+ 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(
⋯ 3 unchanged lines
void * ptr_sfa,
void * ptr_b,
void * ptr_sfb,
- void * ptr_c,
- cudaStream_t stream
+ void * ptr_c
) {
dim3 grid(SPLIT_K * m / BLOCK_M / LOCAL_M, l);
- dim3 block(128 + 64);
+ dim3 block(256);
void* func = (void*) kernel;
- cudaLaunchConfig_t launch_config;
+ cudaLaunchConfig_t launch_config = {0};
launch_config.blockDim = block;
launch_config.gridDim = grid;
- launch_config.stream = stream;
- launch_config.dynamicSmemBytes = 0;
+ // launch_config.dynamicSmemBytes = 0;
cudaLaunchAttribute attrs[16];
attrs[0].id = cudaLaunchAttributeIgnore;
#ifdef CLUSTER_M
⋯ 2 unchanged lines
#endif
attrs[0].val.clusterDim.y = 1;
attrs[0].val.clusterDim.z = 1;
- attrs[1].id = cudaLaunchAttributeProgrammaticStreamSerialization;
- attrs[1].val.programmaticStreamSerializationAllowed = USE_PDL;
launch_config.attrs = attrs;
- launch_config.numAttrs = 2;
+ 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);
}
⋯ 42 unchanged lines
sfa.data_ptr(),
b.data_ptr(),
sfb.data_ptr(),
- c.data_ptr(),
- c10::cuda::getCurrentCUDAStream().stream()
+ c.data_ptr()
);
}
#endif
⋯ 12 unchanged lines
"-O3",
]
+ defines,
- extra_ldflags=["-LcublasLt"],
+ extra_ldflags=["-LcublasLt", "-lcuda"],
)
⋯ 24 unchanged lines
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(
scrolls · 586 diff lines total

Best evidence level for this revision: reported

JSON