KernelBench hard · RTX PRO 6000

FP8 GEMM Kimi K3 (256k)

cleandid not score

manually audited: clean

Genuine hand-written CUDA C++ tcgen05 (5th-gen tensor core) fp8 GEMM via torch load_inline: persistent warp-specialized pipeline (TMA loader warp + MMA umpire warp + 4-warp epilogue), inline-PTX tcgen05.mma kind::f8f6f4 with fp32 accumulation in tensor memory, per-channel dequant scale applied in an SMEM-staged epilogue, TMA bulk-tensor stores of bf16 output. The session started from a Gluon (triton.experimental.gluon) prototype and converged to raw CUDA+PTX. No torch._scaled_mm / aten._scaled_mm anywhere in solution.py. The off-alignment K=4127 shape is handled by a real zero-padding kernel to K0=4128, re-run from live input data on every forward. Host-side caching is limited to compiled-artifact state (CUtensorMap descriptors keyed on base ptr + dims, padded WEIGHT keyed on weight._version) — never outputs. Per-shape fractions 0.4489/0.3441/ 0.0259/0.6091 (geomean 0.2222); shape 2 (M=32 decode) is the known bandwidth/launch-bound shape where every model lands ~0.03.

harnesskinetic-claude (Claude-Code-routed, containerized, live CUDA, B200)
Kernel source (redacted)
"""FP8 e4m3 GEMM for B200 (SM100): y = (x @ w.T) * weight_scale, bf16 out.

Hand-written CUDA C++ tcgen05 (5th-gen tensor core) kernel: persistent
warp-specialized pipeline (TMA loader warp + MMA umpire warp + epilogue
warpgroup), fp8 x fp8 MMA with fp32 accumulation in tensor memory, per-channel
dequant scale in an SMEM-staged epilogue. No vendor GEMM calls anywhere.
"""
import torch
import torch.nn as nn
from torch.utils.cpp_extension import load_inline

OP_TYPE = "gemm"
SUPPORTED_PRECISIONS = ["fp8_e4m3"]
HARDWARE_REQUIRED = ["B200"]
E4M3_MAX = 448.0

_CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <cstdint>
#include <unordered_map>

#define DEV __device__ __forceinline__

DEV static uint32_t elect_one_sync() {
  uint32_t pred = 0;
  asm volatile(
      "{\n"
      ".reg .pred P1;\n"
      "elect.sync _|P1, %1;\n"
      "selp.b32 %0, 1, 0, P1;\n"
      "}\n"
      : "=r"(pred)
      : "r"(0xffffffff));
  return pred;
}

DEV static void mbar_init(uint64_t* bar, uint32_t count) {
  asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;\n"
               :
               : "r"((uint32_t)__cvta_generic_to_shared(bar)), "r"(count));
}

DEV static void mbar_expect(uint64_t* bar, uint32_t bytes) {
  asm volatile("mbarrier.arrive.expect_tx.shared::cta.b64 _, [%0], %1;\n"
               :
               : "r"((uint32_t)__cvta_generic_to_shared(bar)), "r"(bytes));
}

DEV static void mbar_arrive(uint64_t* bar) {
  asm volatile("mbarrier.arrive.shared::cta.b64 _, [%0];\n"
               :
               : "r"((uint32_t)__cvta_generic_to_shared(bar)));
}

DEV static void mbar_wait(uint64_t* bar, uint32_t phase) {
  asm volatile(
      "{\n"
      ".reg .pred P1;\n"
      "WAITLOOP:\n"
      "mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1;\n"
      "@P1 bra DONE;\n"
      "bra WAITLOOP;\n"
      "DONE:\n"
      "}\n"
      :
      : "r"((uint32_t)__cvta_generic_to_shared(bar)), "r"(phase));
}

DEV static void tma_load_2d(const void* smem_dst, const CUtensorMap* map,
                            int32_t c0, int32_t c1, uint64_t* bar) {
  asm volatile(
      "cp.async.bulk.tensor.2d.shared::cluster.global.mbarrier::complete_tx::bytes"
      " [%0], [%1, {%3, %4}], [%2];\n"
      :
      : "r"((uint32_t)__cvta_generic_to_shared(smem_dst)),
        "l"((uint64_t)map),
        "r"((uint32_t)__cvta_generic_to_shared(bar)),
        "r"(c0), "r"(c1)
      : "memory");
}

DEV static void tma_store_2d(const CUtensorMap* map, const void* smem_src,
                             int32_t c0, int32_t c1) {
  asm volatile(
      "cp.async.bulk.tensor.2d.global.shared::cta.bulk_group [%0, {%2, %3}], [%1];\n"
      :
      : "l"((uint64_t)map),
        "r"((uint32_t)__cvta_generic_to_shared(smem_src)),
        "r"(c0), "r"(c1)
      : "memory");
}

DEV static void tma_store_commit() { asm volatile("cp.async.bulk.commit_group;"); }

template <int N>
DEV static void tma_store_wait() {
  asm volatile("cp.async.bulk.wait_group.read %0;" ::"n"(N));
}

DEV static void fence_async_shared() {
  asm volatile("fence.proxy.async.shared::cta;");
}

// smem matrix descriptor (UMMA). layout_type: 2=SW128, 4=SW64, 6=SW32.
DEV static uint64_t smem_desc(uint32_t saddr, uint32_t sbo, uint32_t swz_type) {
  uint64_t d = ((uint64_t)(saddr >> 4) & 0x3FFFULL);
  d |= ((uint64_t)(sbo >> 4)) << 32;
  d |= 1ULL << 46;  // version for blackwell
  d |= (uint64_t)swz_type << 61;
  return d;
}

DEV static void umma_f8(uint64_t da, uint64_t db, uint32_t tmem_d, uint32_t idesc,
                        uint32_t accum) {
  asm volatile(
      "{\n\t"
      ".reg .pred p;\n\t"
      "setp.ne.b32 p, %4, 0;\n\t"
      "tcgen05.mma.cta_group::1.kind::f8f6f4 [%0], %1, %2, %3, {%5, %6, %7, %8}, p; \n\t"
      "}\n"
      :
      : "r"(tmem_d), "l"(da), "l"(db), "r"(idesc), "r"(accum),
        "r"(0), "r"(0), "r"(0), "r"(0));
}

DEV static void umma_commit(uint64_t* bar) {
  asm volatile(
      "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];\n"
      :
      : "r"((uint32_t)__cvta_generic_to_shared(bar)));
}

DEV static void tmem_alloc(uint32_t* smem_out, uint32_t ncols) {
  asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;\n"
               :
               : "r"((uint32_t)__cvta_generic_to_shared(smem_out)), "r"(ncols));
}

DEV static void tmem_dealloc(uint32_t taddr, uint32_t ncols) {
  asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n" ::"r"(taddr),
               "r"(ncols));
  asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;");
}

DEV static void tmem_ld64(float* vals, uint32_t taddr) {
  uint32_t* v = (uint32_t*)vals;
  asm volatile("tcgen05.ld.sync.aligned.32x32b.x64.b32 "
               "{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,%10,%11,%12,%13,%14,%15,"
               "%16,%17,%18,%19,%20,%21,%22,%23,%24,%25,%26,%27,%28,%29,%30,%31,"
               "%32,%33,%34,%35,%36,%37,%38,%39,%40,%41,%42,%43,%44,%45,%46,%47,"
               "%48,%49,%50,%51,%52,%53,%54,%55,%56,%57,%58,%59,%60,%61,%62,%63}, [%64];\n"
               : "=r"(v[0]), "=r"(v[1]), "=r"(v[2]), "=r"(v[3]), "=r"(v[4]), "=r"(v[5]),
                 "=r"(v[6]), "=r"(v[7]), "=r"(v[8]), "=r"(v[9]), "=r"(v[10]), "=r"(v[11]),
                 "=r"(v[12]), "=r"(v[13]), "=r"(v[14]), "=r"(v[15]), "=r"(v[16]), "=r"(v[17]),
                 "=r"(v[18]), "=r"(v[19]), "=r"(v[20]), "=r"(v[21]), "=r"(v[22]), "=r"(v[23]),
                 "=r"(v[24]), "=r"(v[25]), "=r"(v[26]), "=r"(v[27]), "=r"(v[28]), "=r"(v[29]),
                 "=r"(v[30]), "=r"(v[31]), "=r"(v[32]), "=r"(v[33]), "=r"(v[34]), "=r"(v[35]),
                 "=r"(v[36]), "=r"(v[37]), "=r"(v[38]), "=r"(v[39]), "=r"(v[40]), "=r"(v[41]),
                 "=r"(v[42]), "=r"(v[43]), "=r"(v[44]), "=r"(v[45]), "=r"(v[46]), "=r"(v[47]),
                 "=r"(v[48]), "=r"(v[49]), "=r"(v[50]), "=r"(v[51]), "=r"(v[52]), "=r"(v[53]),
                 "=r"(v[54]), "=r"(v[55]), "=r"(v[56]), "=r"(v[57]), "=r"(v[58]), "=r"(v[59]),
                 "=r"(v[60]), "=r"(v[61]), "=r"(v[62]), "=r"(v[63])
               : "r"(taddr));
  asm volatile("tcgen05.wait::ld.sync.aligned;");
}

DEV static constexpr uint32_t make_idesc(int M, int N) {
  return (1u << 4) | ((uint32_t)(N >> 3) << 17) | ((uint32_t)(M >> 4) << 24);
}

// ------------------------------------------------------------------ scheduling
struct TileScheduler {
  int start_pid, num_pid_m, num_pid_n, num_pid_in_group, num_pid, num_progs;
  int group_m;
  __device__ void init(int M, int N, int BM, int BN, int gm) {
    start_pid = blockIdx.x;
    num_progs = gridDim.x;
    num_pid_m = (M + BM - 1) / BM;
    num_pid_n = (N + BN - 1) / BN;
    num_pid_in_group = gm * num_pid_n;
    num_pid = num_pid_m * num_pid_n;
    group_m = gm;
  }
  __device__ int num_tiles() const { return (num_pid - start_pid + num_progs - 1) / num_progs; }
  __device__ void get_tile(int idx, int& pm, int& pn) const {
    int tile_id = start_pid + idx * num_progs;
    int group_id = tile_id / num_pid_in_group;
    int first_pid_m = group_id * group_m;
    int group_size = min(num_pid_m - first_pid_m, group_m);
    pm = first_pid_m + (tile_id % group_size);
    pn = (tile_id % num_pid_in_group) / group_size;
  }
};

// ------------------------------------------------------------------ kernel
template <int BM, int BN, int BK, int STAGES, int GM>
__global__ void __launch_bounds__(192, 1)
k_gemm(const __grid_constant__ CUtensorMap a_map,
       const __grid_constant__ CUtensorMap b_map,
       const __grid_constant__ CUtensorMap c_map,
       const float* __restrict__ scale, int M, int N, int K) {
  constexpr int A_BYTES = BM * BK;
  constexpr int B_BYTES = BN * BK;
  constexpr int C_COLS = 64;               // bf16 columns per c box
  constexpr int C_ELEMS = BM * C_COLS;     // bf16 elements per c buffer
  constexpr int N_CHUNKS = BN / C_COLS;
  constexpr int EPI_WARPS = 4;
  constexpr int EPI_THREADS = EPI_WARPS * 32;
  // m64 accumulator packs rows into 16-lane groups; the two acc buffers are
  // interleaved in lane halves: buf b at lane offset 16*b (same columns).
  // m128 accumulator: buf b at column offset b * BN.
  constexpr int SWZ_TYPE = (BK == 128) ? 2 : 4;  // 128B or 64B swizzle

  extern __shared__ __align__(1024) uint8_t smem[];
  uint8_t* a_smem = smem;                       // STAGES * A_BYTES
  uint8_t* b_smem = a_smem + STAGES * A_BYTES;  // STAGES * B_BYTES
  __nv_bfloat16* c_smem = (__nv_bfloat16*)(b_smem + STAGES * B_BYTES);  // 2 * C_ELEMS
  float* s_smem = (float*)(c_smem + 2 * C_ELEMS);                       // BN floats

  __shared__ __align__(8) uint64_t full_bars[STAGES], empty_bars[STAGES];
  __shared__ __align__(8) uint64_t acc_full[2], acc_empty[2];
  __shared__ uint32_t tmem_addr_sh;

  const int warp = threadIdx.x >> 5;
  const int lane = threadIdx.x & 31;

  if (warp == EPI_WARPS) {
    tmem_alloc(&tmem_addr_sh, 512);
  } else if (warp == EPI_WARPS + 1 && lane == 0) {
    for (int i = 0; i < STAGES; ++i) {
      mbar_init(&full_bars[i], 1);
      mbar_init(&empty_bars[i], 1);
    }
    mbar_init(&acc_full[0], 1);
    mbar_init(&acc_full[1], 1);
    mbar_init(&acc_empty[0], EPI_WARPS);
    mbar_init(&acc_empty[1], EPI_WARPS);
  }
  asm volatile("barrier.sync 1, 192;");
  const uint32_t taddr_base = tmem_addr_sh;
  const uint32_t idesc = make_idesc(BM, BN);
  constexpr uint32_t ACC1_OFF = (BM == 64) ? (16u << 16) : (uint32_t)BN;

  TileScheduler sched;
  sched.init(M, N, BM, BN, GM);
  const int k_iters = (K + BK - 1) / BK;

  if (warp == EPI_WARPS) {
    // ---------------------------------------------------------- loader
    if (elect_one_sync()) {
      uint32_t e_phase = 1;  // first STAGES waits bypass
      int stage = 0;
      for (int t = 0; t < sched.num_tiles(); ++t) {
        int pm, pn;
        sched.get_tile(t, pm, pn);
        const int off_m = pm * BM, off_n = pn * BN;
        for (int k = 0; k < k_iters; ++k) {
          const int s = stage;
          mbar_wait(&empty_bars[s], e_phase);
          mbar_expect(&full_bars[s], A_BYTES + B_BYTES);
          tma_load_2d(a_smem + s * A_BYTES, &a_map, k * BK, off_m, &full_bars[s]);
          tma_load_2d(b_smem + s * B_BYTES, &b_map, k * BK, off_n, &full_bars[s]);
          if (++stage == STAGES) {
            stage = 0;
            e_phase ^= 1;
          }
        }
      }
    }
  } else if (warp == EPI_WARPS + 1) {
    // ---------------------------------------------------------- mma umpire
    if (elect_one_sync()) {
      uint32_t f_phase = 0, ae_phase = 1;
      int stage = 0, abuf = 0;
      for (int t = 0; t < sched.num_tiles(); ++t) {
        mbar_wait(&acc_empty[abuf], ae_phase);
        const uint32_t taddr = taddr_base + abuf * ACC1_OFF;
        for (int k = 0; k < k_iters; ++k) {
          const int s = stage;
          mbar_wait(&full_bars[s], f_phase);
          uint64_t da = smem_desc((uint32_t)__cvta_generic_to_shared(a_smem + s * A_BYTES),
                                  8 * BK, SWZ_TYPE);
          uint64_t db = smem_desc((uint32_t)__cvta_generic_to_shared(b_smem + s * B_BYTES),
                                  8 * BK, SWZ_TYPE);
          for (int kk = 0; kk < BK / 32; ++kk) {
            umma_f8(da + kk * 2, db + kk * 2, taddr, idesc, (k > 0 || kk > 0) ? 1u : 0u);
          }
          umma_commit(&empty_bars[s]);
          if (++stage == STAGES) {
            stage = 0;
            f_phase ^= 1;
          }
        }
        umma_commit(&acc_full[abuf]);
        if (++abuf == 2) {
          abuf = 0;
          ae_phase ^= 1;
        }
      }
    }
  } else if (warp < EPI_WARPS) {
    // ---------------------------------------------------------- epilogue
    // m128: each warp owns 32 rows. m64: lanes hold 16 rows per warp in the
    // 16-lane half selected by the accumulator buffer index.
    int row = 0;
    bool row_ok = true;
    if (BM == 128) {
      row = warp * 32 + lane;
    } else {
      row_ok = (lane / 16) == 0;  // placeholder; abuf-dependent, set in loop
    }
    uint32_t af_phase = 0;
    int abuf = 0;
    for (int t = 0; t < sched.num_tiles(); ++t) {
      int pm, pn;
      sched.get_tile(t, pm, pn);
      const int off_m = pm * BM, off_n = pn * BN;
      // preload scales for this tile (BN floats) into smem
      for (int j = threadIdx.x; j < BN; j += EPI_THREADS) {
        const int n = off_n + j;
        s_smem[j] = (n < N) ? scale[n] : 0.f;
      }
      mbar_wait(&acc_full[abuf], af_phase);
      const uint32_t taddr = taddr_base + abuf * ACC1_OFF;
      for (int c = 0; c < N_CHUNKS; ++c) {
        float vals[C_COLS];
        tmem_ld64(vals, taddr + c * C_COLS);
        if (c == N_CHUNKS - 1 && lane == 0) {
          // all tmem reads of this accumulator are done
          mbar_arrive(&acc_empty[abuf]);
        }
        __nv_bfloat16 outv[C_COLS];
        const float* sv = s_smem + c * C_COLS;
        for (int j = 0; j < C_COLS; ++j) {
          outv[j] = __float2bfloat16(vals[j] * sv[j]);
        }
        // wait for the tma store that used this buffer 2 chunks ago
        if (c >= 2 && warp == 0 && elect_one_sync()) {
          tma_store_wait<1>();
        }
        asm volatile("barrier.sync 2, %0;" ::"n"(EPI_THREADS));
        __nv_bfloat16* dst = c_smem + (c & 1) * C_ELEMS;
        if (BM == 128) {
          uint4* d4 = (uint4*)(dst + row * C_COLS);
          const uint4* s4 = (const uint4*)outv;
          for (int j = 0; j < C_COLS * 2 / 16; ++j) d4[j] = s4[j];
        } else {
          // m64: lanes 16*abuf..16*abuf+16 of this warp hold rows
          // warp*16 + (lane - 16*abuf).
          const int l16 = lane - 16 * abuf;
          if (l16 >= 0 && l16 < 16) {
            uint4* d4 = (uint4*)(dst + (warp * 16 + l16) * C_COLS);
            const uint4* s4 = (const uint4*)outv;
            for (int j = 0; j < C_COLS * 2 / 16; ++j) d4[j] = s4[j];
          }
        }
        asm volatile("barrier.sync 2, %0;" ::"n"(EPI_THREADS));
        if (warp == 0 && elect_one_sync()) {
          fence_async_shared();
          tma_store_2d(&c_map, c_smem + (c & 1) * C_ELEMS, off_n + c * C_COLS, off_m);
          tma_store_commit();
        }
      }
      if (++abuf == 2) {
        abuf = 0;
        af_phase ^= 1;
      }
    }
    if (warp == 0 && elect_one_sync()) {
      tma_store_wait<0>();
    }
  }
  asm volatile("barrier.sync 3, 192;");
  if (warp == EPI_WARPS + 1) {
    tmem_dealloc(taddr_base, 512);
  }
}

// ------------------------------------------------------------------ host
__global__ void k_pad2_funnel(const uint32_t* __restrict__ src_a, uint32_t* __restrict__ dst_a,
                              const uint32_t* __restrict__ src_b, uint32_t* __restrict__ dst_b,
                              int K, int dst_words, int src_words_a, int src_words_b, int rows_b) {
  const int row = blockIdx.x;
  const bool is_a = row < rows_b;
  const uint32_t* src = is_a ? src_a : src_b;
  uint32_t* dst = is_a ? dst_a : dst_b;
  const int src_words = is_a ? src_words_a : src_words_b;
  const long long sr = (long long)(is_a ? row : row - rows_b) * K;
  const int w0 = (int)(sr >> 2);
  const int sh = (int)(sr & 3) * 8;
  for (int j = threadIdx.x; j < dst_words; j += blockDim.x) {
    const int sw = w0 + j;
    const unsigned a = (sw < src_words) ? src[sw] : 0u;
    const unsigned b = (sw + 1 < src_words) ? src[sw + 1] : 0u;
    unsigned w = (a >> sh);
    if (sh != 0) w |= (b << (32 - sh));
    const int byte0 = j * 4;
    int vb = K - byte0;
    vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb);
    const unsigned keep = (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1);
    dst[(size_t)(is_a ? row : row - rows_b) * dst_words + j] = w & keep;
  }
}


// pad kernel: unpack odd-K fp8 rows into u32-aligned rows via funnel shifts
// (128-bit stores + shared adjacent-word loads)
__global__ void k_pad_funnel(const uint32_t* __restrict__ src, uint4* __restrict__ dst,
                             int K, int dst_words, int src_words) {
  const int row = blockIdx.x;
  const long long sr = (long long)row * K;
  const int w0 = (int)(sr >> 2);
  const int sh = (int)(sr & 3) * 8;
  const int n_v = dst_words >> 2;
  for (int j = threadIdx.x; j < n_v; j += blockDim.x) {
    const int sw = w0 + (j << 2);
    const unsigned s0 = (sw < src_words) ? src[sw] : 0u;
    const unsigned s1 = (sw + 1 < src_words) ? src[sw + 1] : 0u;
    const unsigned s2 = (sw + 2 < src_words) ? src[sw + 2] : 0u;
    const unsigned s3 = (sw + 3 < src_words) ? src[sw + 3] : 0u;
    const unsigned s4 = (sw + 4 < src_words) ? src[sw + 4] : 0u;
    uint4 v;
    if (sh == 0) {
      v.x = s0; v.y = s1; v.z = s2; v.w = s3;
    } else {
      const int shr = 32 - sh;
      v.x = (s0 >> sh) | (s1 << shr);
      v.y = (s1 >> sh) | (s2 << shr);
      v.z = (s2 >> sh) | (s3 << shr);
      v.w = (s3 >> sh) | (s4 << shr);
    }
    const int base = (j << 2) * 4;
    int vb;
    vb = K - base;         vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb);
    v.x &= (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1);
    vb = K - (base + 4);   vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb);
    v.y &= (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1);
    vb = K - (base + 8);   vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb);
    v.z &= (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1);
    vb = K - (base + 12);  vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb);
    v.w &= (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1);
    dst[(size_t)row * n_v + j] = v;
  }
}

namespace {

struct MapCache {
  std::unordered_map<uint64_t, CUtensorMap> maps;
  CUtensorMap get(const void* base, uint64_t inner, uint64_t outer, uint32_t box_inner,
                  uint32_t box_outer, CUtensorMapSwizzle swz, CUtensorMapDataType dt,
                  int elembytes) {
    uint64_t key = ((uint64_t)base) ^ (inner << 1) ^ (outer << 13) ^ ((uint64_t)box_inner << 40) ^
                   ((uint64_t)box_outer << 48) ^ ((uint64_t)swz << 56) ^ ((uint64_t)dt << 61);
    auto it = maps.find(key);
    if (it != maps.end()) return it->second;
    CUtensorMap map{};
    uint64_t dims[2] = {inner, outer};
    uint64_t strides[1] = {inner * (uint64_t)elembytes};
    uint32_t box[2] = {box_inner, box_outer};
    uint32_t estrides[2] = {1, 1};
    cuTensorMapEncodeTiled(&map, dt, 2, const_cast<void*>(base), dims, strides, box,
                           estrides, CU_TENSOR_MAP_INTERLEAVE_NONE, swz,
                           CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
                           CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
    return map;
  }
};

constexpr CUtensorMapSwizzle CuSw(int bk) {
  return bk == 128 ? CU_TENSOR_MAP_SWIZZLE_128B : CU_TENSOR_MAP_SWIZZLE_64B;
}

MapCache g_maps;
int g_num_sms = 0;

}  // namespace

template <int BM, int BN, int BK, int STAGES, int GM>
void launch_cfg(const CUtensorMap* a_map, const CUtensorMap* b_map, const CUtensorMap* c_map,
                const float* scale, int M, int N, int K, cudaStream_t stream) {
  constexpr int SMEM = STAGES * (BM * BK + BN * BK) + 2 * (BM * 64 * 2) + BN * 4 + 1024;
  auto kern = k_gemm<BM, BN, BK, STAGES, GM>;
  static bool attr_set = false;
  if (!attr_set) {
    cudaFuncSetAttribute(kern, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
    attr_set = true;
  }
  int tiles_m = (M + BM - 1) / BM, tiles_n = (N + BN - 1) / BN;
  int grid = std::min(g_num_sms, tiles_m * tiles_n);
  kern<<<grid, 192, SMEM, stream>>>(*a_map, *b_map, *c_map, scale, M, N, K);
}

void fp8_run(int64_t x_ptr, int64_t w_ptr, int64_t s_ptr, int64_t y_ptr,
             int64_t M, int64_t N, int64_t K, int64_t cfg, int64_t stream_ptr) {
  void* x = (void*)x_ptr;
  void* w = (void*)w_ptr;
  if (!g_num_sms) {
    cudaDeviceProp prop;
    cudaGetDeviceProperties(&prop, 0);
    g_num_sms = prop.multiProcessorCount;
  }
  cudaStream_t stream = (cudaStream_t)stream_ptr;
  if (cfg == 0) {
    CUtensorMap a_map = g_maps.get(x, K, M, 128, 64, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1);
    CUtensorMap b_map = g_maps.get(w, K, N, 128, 64, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1);
    CUtensorMap c_map = g_maps.get((void*)y_ptr, N, M, 64, 64, CU_TENSOR_MAP_SWIZZLE_NONE,
                                   CU_TENSOR_MAP_DATA_TYPE_UINT16, 2);
    launch_cfg<64, 64, 128, 12, 1>(&a_map, &b_map, &c_map, (const float*)s_ptr, (int)M, (int)N, (int)K, stream);
  } else if (cfg == 1) {
    CUtensorMap a_map = g_maps.get(x, K, M, 128, 128, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1);
    CUtensorMap b_map = g_maps.get(w, K, N, 128, 256, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1);
    CUtensorMap c_map = g_maps.get((void*)y_ptr, N, M, 64, 128, CU_TENSOR_MAP_SWIZZLE_NONE,
                                   CU_TENSOR_MAP_DATA_TYPE_UINT16, 2);
    launch_cfg<128, 256, 128, 4, 4>(&a_map, &b_map, &c_map, (const float*)s_ptr, (int)M, (int)N, (int)K, stream);
  } else {
    CUtensorMap a_map = g_maps.get(x, K, M, 128, 128, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1);
    CUtensorMap b_map = g_maps.get(w, K, N, 128, 256, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1);
    CUtensorMap c_map = g_maps.get((void*)y_ptr, N, M, 64, 128, CU_TENSOR_MAP_SWIZZLE_NONE,
                                   CU_TENSOR_MAP_DATA_TYPE_UINT16, 2);
    launch_cfg<128, 256, 128, 4, 16>(&a_map, &b_map, &c_map, (const float*)s_ptr, (int)M, (int)N, (int)K, stream);
  }
}

void fp8_pad2(int64_t xa_ptr, int64_t xb_ptr, int64_t wa_ptr, int64_t wb_ptr,
              int64_t Ma, int64_t Mb, int64_t K, int64_t K_pad, int64_t stream_ptr) {
  const int dst_words = (int)(K_pad / 4);
  const int sa = (int)(Ma * K / 4);
  const int sb = (int)(Mb * K / 4);
  k_pad2_funnel<<<(int)(Ma + Mb), 256, 0, (cudaStream_t)stream_ptr>>>(
      (const uint32_t*)xa_ptr, (uint32_t*)xb_ptr, (const uint32_t*)wa_ptr, (uint32_t*)wb_ptr,
      (int)K, dst_words, sa, sb, (int)Ma);
}

void fp8_pad(int64_t src_ptr, int64_t dst_ptr, int64_t M, int64_t K, int64_t K_pad,
             int64_t stream_ptr) {
  const int dst_words = (int)(K_pad / 4);
  const int src_words = (int)(M * K / 4);
  k_pad_funnel<<<(int)M, 256, 0, (cudaStream_t)stream_ptr>>>(
      (const uint32_t*)src_ptr, (uint4*)dst_ptr, (int)K, dst_words, src_words);
}

"""

_ext = load_inline(
    name="fp8_gemm_b200_v1",
    cpp_sources=[r"""
void fp8_run(int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t);
void fp8_pad(int64_t, int64_t, int64_t, int64_t, int64_t, int64_t);
void fp8_pad2(int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t);
"""],
    cuda_sources=[_CUDA_SRC],
    functions=["fp8_run", "fp8_pad", "fp8_pad2"],
    extra_cuda_cflags=["-O3", "-std=c++17", "--use_fast_math",
                        "-gencode=arch=compute_100a,code=sm_100a", "-lcuda"],
    extra_ldflags=["-lcuda"],
    verbose=False,
)


def _raw_stream():
    return torch._C._cuda_getCurrentRawStream(torch.cuda.current_device())


class Model(nn.Module):
    """y = ((x_fp8 @ w_fp8.T) * weight_scale).to(bf16) on tcgen05 tensor cores."""

    def __init__(self, M: int, N: int, K: int):
        super().__init__()
        self.M, self.N, self.K = M, N, K
        w = torch.empty(N, K, dtype=torch.bfloat16)
        nn.init.normal_(w, std=0.02)
        s = (w.float().abs().amax(dim=1, keepdim=True) / E4M3_MAX).clamp(min=1e-12)
        w_fp8 = (w.float() / s).to(torch.float8_e4m3fn)
        self.register_buffer("weight", w_fp8)                         # (N, K) fp8
        self.register_buffer("weight_scale", s.squeeze(1).to(torch.float32))  # (N,)
        self._run = _ext.fp8_run
        if M <= 64:
            self._cfg = 0
        elif N >= 8192:
            self._cfg = 2
        else:
            self._cfg = 1
        # pad state for unaligned K
        self._w_pad = None  # (version, buf)
        self._x_pad = None  # buf for x


    def _pad(self, t, buf):
        _ext.fp8_pad(t.data_ptr(), buf.data_ptr(), t.shape[0], t.shape[1], buf.shape[1],
                     _raw_stream())
        return buf


    def forward(self, x: torch.Tensor) -> torch.Tensor:
        M, K = x.shape
        y = torch.empty((M, self.N), dtype=torch.bfloat16, device=x.device)
        stream = torch._C._cuda_getCurrentRawStream(0)
        if K % 16 == 0:
            w = self.weight
            K0 = K
        else:
            K0 = (K + 15) // 16 * 16
            if self._x_pad is None:
                self._x_pad = torch.zeros((M, K0), dtype=x.dtype, device=x.device)
            w = self.weight
            ver = w._version
            need_w = self._w_pad is None or self._w_pad[0] != ver
            if need_w:
                wbuf = torch.zeros((self.N, K0), dtype=w.dtype, device=w.device)
                _ext.fp8_pad2(x.data_ptr(), self._x_pad.data_ptr(), w.data_ptr(), wbuf.data_ptr(),
                              M, self.N, K, K0, stream)
                self._w_pad = (ver, wbuf)
                w = wbuf
                x = self._x_pad
            else:
                x = self._pad(x, self._x_pad)
                w = self._w_pad[1]
        self._run(x.data_ptr(), w.data_ptr(), self.weight_scale.data_ptr(), y.data_ptr(),
              M, self.N, K0, self._cfg, stream)
        return y


Model.__call__ = Model.forward


M = 4096
N = 4096
K = 4096


def get_inputs():
    x = (torch.rand(M, K) * 8 - 4).to(torch.float8_e4m3fn)
    return [x]


def get_init_inputs():
    return [M, N, K]

20260715_220649_kinetic-claude_kinetic-0715_01_fp8_gemm