Skip to content
KernelIndex
Search⌘K

submission 97026

v0i0 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission-22-12-31-53.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-97026?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.9µs
#55 of 678
2025-11-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f78afe6e0182e3cfc52d84c6a75e490e78f8fc4ed49ec3aa5be18e41b6e951c9
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-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;
vector-width = uint4using vec_type = uint4;

Kernel source

submission-22-12-31-53.py537 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
#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)

#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__ 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));
}

__launch_bounds__(128 + 64, 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
) {
  // 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, ?))

  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 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;
  };

  if (is_load) {
    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) {
      #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);
        }
        #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
        #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
          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;
    #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;
        barrier_wait(&barrier_smem_full[read_stage], read_phase);

        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;
        }
        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");
      for (;;) {
        int value = 1;
        ASMV(ld.acquire.gpu.global.b32 %0, [%1];) CLO("=r"(value) : "l"(sem) : "memory");
        if (value == 0) break;
      }
    }
    __syncwarp();
  }

  __syncthreads();

  // 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);
    }
  }
  __syncthreads();
  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]);
      }
    }
  }
}

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

  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 = 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;
  attrs[1].id = cudaLaunchAttributeProgrammaticStreamSerialization;
  attrs[1].val.programmaticStreamSerializationAllowed = USE_PDL;
  launch_config.attrs = attrs;
  launch_config.numAttrs = 2;
  void* args[] = {
    (void*) &ptr_a,
    (void*) &ptr_sfa,
    (void*) &ptr_b,
    (void*) &ptr_sfb,
    (void*) &ptr_c,
  };
  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(),
    c10::cuda::getCurrentCUDAStream().stream() 
  );
}
#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"],
    )


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


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 · 537 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 88871.

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
- module = load_inline(cuda_sources=["""
+
+ 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
#include <c10/cuda/CUDAStream.h>
#endif
⋯ 9 unchanged lines
#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 STAGE_COUNT = 8;
const int BLOCK_K = 256;
template<int N>
⋯ 23 unchanged lines
}) 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
+ __half * __restrict__ ptr_c
) {
// 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.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;
+ int block_offset_m = blockIdx.x / SPLIT_K;
+ int block_offset_k = blockIdx.x % SPLIT_K;
+ int k = STATIC_K / SPLIT_K;
- ptr_c += threadIdx.x * stride_c_m;
+ 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, ?))
⋯ 29 unchanged lines
}
__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); };
+ 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);
};
⋯ 3 unchanged lines
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;
+ 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);
⋯ 9 unchanged lines
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) {
+ #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);
}
#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_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) {
⋯ 2 unchanged lines
// 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]);
+ 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]);
+ 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;
⋯ 6 unchanged lines
} else if (is_mma) {
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;
barrier_wait(&barrier_smem_full[read_stage], read_phase);
+
+ 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_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
+ 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;
⋯ 22 unchanged lines
@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");
+ for (;;) {
+ int value = 1;
+ ASMV(ld.acquire.gpu.global.b32 %0, [%1];) CLO("=r"(value) : "l"(sem) : "memory");
+ if (value == 0) break;
+ }
+ }
+ __syncwarp();
}
__syncthreads();
- asm volatile("griddepcontrol.launch_dependents;");
+ // asm volatile("griddepcontrol.launch_dependents;");
// load tmem
- float result[LOCAL_M];
+ __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] = ld_result;
+ result[block_m * BLOCK_M + threadIdx.x] = __float2half(ld_result);
}
}
__syncthreads();
⋯ 1 unchanged lines
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]);
+ 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]);
+ }
}
}
}
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 grid(SPLIT_K * m / BLOCK_M / LOCAL_M, l);
dim3 block(128 + 64);
- void* func = local_m == 2 ? (void*) kernel<2> : (void*) kernel<1>;
+ void* func = (void*) kernel;
cudaLaunchConfig_t launch_config;
launch_config.blockDim = block;
⋯ 1 unchanged lines
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;
+ #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;
attrs[1].id = cudaLaunchAttributeProgrammaticStreamSerialization;
- attrs[1].val.programmaticStreamSerializationAllowed = 1; // 1
+ attrs[1].val.programmaticStreamSerializationAllowed = USE_PDL;
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);
}
⋯ 3 unchanged lines
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(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");
+ 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,
- 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"],
- )
+ """
+ ],
+ 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"],
+ )
- # TODO
- PRECOMPILE_AND_TUNE_STAGES = [
- (16384, 1, 7168),
- (7168, 8, 4096),
- (2048, 4, 7168),
- ]
+ 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])])
+
+
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)
+ 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 · 554 diff lines total

Best evidence level for this revision: reported

JSON