Skip to content
KernelIndex
Search⌘K

submission 487235

Joel🏴 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sub.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-487235?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 group GEMMsuite of 4 cases
NVIDIA B200
68.8µs
#79 of 145
2026-02-08

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:416f944e6fc5a4c3c4f20e87c3e0fb31bf3ef32f63e0b9fb45117fb4532eda04
license declaredunknown
license concludedunknown
authorsJoel🏴
imported2026-08-15

Techniques

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

cluster__global__ void __cluster_dims__(CLUSTER_M, CLUSTER_N_PARAM, CLUSTER_Z)
fused-epiloguealignas(8) uint64_t epilogue_mbar;
mbarrier__device__ inline void mbarrier_init(int mbar_addr, int count)
num-warps = 6constexpr int NUM_WARPS = 6;
shared-memory__device__ inline void tma_1d_gmem2smem_mcast(int dst, const void *tmap_ptr, int x, int mbar_addr, int16_t cta_mask, uint64_t cache_policy)
tcgen05asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" ::"r"(taddr), "l"(s_desc), "n"(CTA_GROUP));
tile-k = 256constexpr int BLOCK_K = 256;
tile-m = 128constexpr int BLOCK_M = 128;
tile-n = 128constexpr int BLOCK_N = 128;
tmaasm volatile("cp.async.bulk.prefetch.L2.global.L2::cache_hint [%0], %1, %2;" ::"l"(src), "r"(size), "l"(cache_policy) : "memory");
vector-width = half2half2 h2_tmp[4];

Kernel source

sub.py1425 lines
#!POPCORN leaderboard nvfp4_group_gemm
#!POPCORN gpu B200

import os

import torch
from torch.utils.cpp_extension import load_inline

CUDA_SRC_UTILS = r"""
#include <cuda.h>
#include <cudaTypedefs.h>
#include <cuda_fp16.h>
#include <cstdint>
#include <cstdio>
#include <cstdlib>

constexpr uint64_t EVICT_FIRST = 0x12F0000000000000;
constexpr uint64_t EVICT_LAST = 0x14F0000000000000;
__device__ inline constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; };
__device__
    uint32_t
    elect_sync()
{
  uint32_t pred = 0;
  asm volatile(
      "{\n\t"
      ".reg .pred %%px;\n\t"
      "elect.sync _|%%px, %1;\n\t"
      "@%%px mov.s32 %0, 1;\n\t"
      "}"
      : "+r"(pred)
      : "r"(0xFFFFFFFF));
  return pred;
}
__device__ inline void mbarrier_init(int mbar_addr, int count)
{
  asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" ::"r"(mbar_addr), "r"(count));
}
__device__ inline void mbarrier_arrive_expect_tx(int mbar_addr, int size)
{
  asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;" ::"r"(mbar_addr), "r"(size) : "memory");
}
__device__ void mbarrier_wait(int mbar_addr, int phase)
{
  uint32_t ticks = 0x989680;
  asm volatile(
      "{\n\t"
      ".reg .pred P1;\n\t"
      "LAB_WAIT:\n\t"
      "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t"
      "@!P1 bra.uni LAB_WAIT;\n\t"
      "}" ::"r"(mbar_addr),
      "r"(phase), "r"(ticks));
}
__device__ inline void mbarrier_arrive(int mbar_addr)
{
  asm volatile("mbarrier.arrive.release.cta.shared::cta.b64 _, [%0];" ::"r"(mbar_addr) : "memory");
}
__device__ inline void fence_mbarrier_init()
{
  asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
}
__device__ inline void prefetch_tensormap(const void *tmap_ptr)
{
  asm volatile("prefetch.tensormap [%0];" ::"l"(tmap_ptr) : "memory");
}
__device__ inline void tma_prefetch(const void *src, int size, uint64_t cache_policy)
{
  asm volatile("cp.async.bulk.prefetch.L2.global.L2::cache_hint [%0], %1, %2;" ::"l"(src), "r"(size), "l"(cache_policy) : "memory");
}
__device__ inline void tma_1d_prefetch(const void *tmap_ptr, int x, uint64_t cache_policy)
{
  asm volatile("cp.async.bulk.prefetch.tensor.1d.L2.global.L2::cache_hint [%0, {%1}], %2;" ::"l"(tmap_ptr), "r"(x), "l"(cache_policy) : "memory");
}
__device__ inline void tma_2d_prefetch(const void *tmap_ptr, int x, int y, uint64_t cache_policy)
{
  asm volatile("cp.async.bulk.prefetch.tensor.2d.L2.global.L2::cache_hint [%0, {%1, %2}], %3;" ::"l"(tmap_ptr), "r"(x), "r"(y), "l"(cache_policy) : "memory");
}
__device__ inline void tma_3d_prefetch(const void *tmap_ptr, int x, int y, int z, uint64_t cache_policy)
{
  asm volatile("cp.async.bulk.prefetch.tensor.3d.L2.global.L2::cache_hint [%0, {%1, %2, %3}], %4;" ::"l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "l"(cache_policy) : "memory");
}
__device__ inline void tma_gmem2smem(int dst, const void *src, int size, int mbar_addr, uint64_t cache_policy)
{
  asm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;" ::"r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy));
}
template <int CTA_GROUP = 1>
__device__ inline void tma_1d_gmem2smem(int dst, const void *tmap_ptr, int x, int mbar_addr, uint64_t cache_policy)
{
  asm volatile("cp.async.bulk.tensor.1d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::%5.L2::cache_hint "
               "[%0], [%1, {%2}], [%3], %4;" ::"r"(dst),
               "l"(tmap_ptr), "r"(x), "r"(mbar_addr), "l"(cache_policy), "n"(CTA_GROUP)
               : "memory");
}
template <int CTA_GROUP = 1>
__device__ inline void tma_1d_gmem2smem_mcast(int dst, const void *tmap_ptr, int x, int mbar_addr, int16_t cta_mask, uint64_t cache_policy)
{
  asm volatile("cp.async.bulk.tensor.1d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.cta_group::%6.L2::cache_hint "
               "[%0], [%1, {%2}], [%3], %4, %5;" ::"r"(dst),
               "l"(tmap_ptr), "r"(x), "r"(mbar_addr), "h"(cta_mask), "l"(cache_policy), "n"(CTA_GROUP)
               : "memory");
}
template <int CTA_GROUP = 1>
__device__ inline void tma_2d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int mbar_addr, uint64_t cache_policy)
{
  asm volatile("cp.async.bulk.tensor.2d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::%6.L2::cache_hint "
               "[%0], [%1, {%2, %3}], [%4], %5;" ::"r"(dst),
               "l"(tmap_ptr), "r"(x), "r"(y), "r"(mbar_addr), "l"(cache_policy), "n"(CTA_GROUP)
               : "memory");
}
template <int CTA_GROUP = 1>
__device__ inline void tma_2d_gmem2smem_mcast(int dst, const void *tmap_ptr, int x, int y, int mbar_addr, int16_t cta_mask, uint64_t cache_policy)
{
  asm volatile("cp.async.bulk.tensor.2d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.cta_group::%7.L2::cache_hint "
               "[%0], [%1, {%2, %3}], [%4], %5, %6;" ::"r"(dst),
               "l"(tmap_ptr), "r"(x), "r"(y), "r"(mbar_addr), "h"(cta_mask), "l"(cache_policy), "n"(CTA_GROUP)
               : "memory");
}
template <int CTA_GROUP = 1>
__device__ inline void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy)
{
  asm volatile("cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::%7.L2::cache_hint "
               "[%0], [%1, {%2, %3, %4}], [%5], %6;" ::"r"(dst),
               "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "l"(cache_policy), "n"(CTA_GROUP)
               : "memory");
}
template <int CTA_GROUP = 1>
__device__ inline void tma_3d_gmem2smem_mcast(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, int16_t cta_mask, uint64_t cache_policy)
{
  asm volatile("cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.cta_group::%8.L2::cache_hint "
               "[%0], [%1, {%2, %3, %4}], [%5], %6, %7;" ::"r"(dst),
               "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "h"(cta_mask), "l"(cache_policy), "n"(CTA_GROUP)
               : "memory");
}
template <int CTA_GROUP = 1>
__device__ inline void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc)
{
  asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" ::"r"(taddr), "l"(s_desc), "n"(CTA_GROUP));
}
template <int CTA_GROUP = 1>
__device__ inline void tcgen05_commit(int mbar_addr)
{
  asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];" ::"r"(mbar_addr), "n"(CTA_GROUP) : "memory");
}
template <int CTA_GROUP = 1>
__device__ inline void tcgen05_commit_mcast(int mbar_addr, uint16_t cta_mask)
{
  asm volatile("tcgen05.commit.cta_group::2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;" ::"r"(mbar_addr), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
}
__device__ inline void tcgen05_fence_after_thread_sync()
{
  asm volatile("tcgen05.fence::after_thread_sync;");
}
__device__ inline void tcgen05_fence_before_thread_sync()
{
  asm volatile("tcgen05.fence::before_thread_sync;");
}
struct COLLECTOR_USAGE
{
  static constexpr char NONE[] = "";
  static constexpr char A_FILL[] = ".collector::a::fill";
  static constexpr char A_USE[] = ".collector::a::use";
  static constexpr char A_LASTUSE[] = ".collector::a::lastuse";
  static constexpr char A_DISCARD[] = ".collector::a::discard";
};
template <int CTA_GROUP = 1, const char *collector_usage = COLLECTOR_USAGE::NONE>
__device__ inline void tcgen05_mma_nvfp4(
    int d_tmem,
    uint64_t a_desc,
    uint64_t b_desc,
    uint32_t i_desc,
    int scale_A_tmem,
    int scale_B_tmem,
    int enable_input_d)
{
  asm volatile(
      "{\n\t"
      ".reg .pred p;\n\t"
      "setp.ne.b32 p, %6, 0;\n\t"
      "tcgen05.mma.cta_group::%7.kind::mxf4nvf4.block_scale.block16%8 [%0], %1, %2, %3, [%4], [%5], p;\n\t"
      "}" ::"r"(d_tmem),
      "l"(a_desc), "l"(b_desc), "r"(i_desc),
      "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d),
      "n"(CTA_GROUP), "C"(collector_usage));
}
struct SHAPE
{
  static constexpr char _32x32b[] = ".32x32b";
  static constexpr char _16x128b[] = ".16x128b";
  static constexpr char _16x256b[] = ".16x256b";
};
template <int NUM_REGS, const char *SHAPE, int NUM>
__device__ inline void tcgen05_ld(float *tmp, uint32_t tmem_addr, int row, int col)
{
  int addr = (row << 16) + tmem_addr + col;
  if constexpr (NUM_REGS == 1)
    asm volatile("tcgen05.ld.sync.aligned%2.x%3.b32 {%0}, [%1];"
                 : "=f"(tmp[0]) : "r"(addr), "C"(SHAPE), "n"(NUM));
  if constexpr (NUM_REGS == 2)
    asm volatile("tcgen05.ld.sync.aligned%3.x%4.b32 {%0, %1}, [%2];"
                 : "=f"(tmp[0]), "=f"(tmp[1]) : "r"(addr), "C"(SHAPE), "n"(NUM));
  if constexpr (NUM_REGS == 4)
    asm volatile("tcgen05.ld.sync.aligned%5.x%6.b32 "
                 "{%0, %1, %2, %3}, [%4];"
                 : "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3])
                 : "r"(addr), "C"(SHAPE), "n"(NUM));
  if constexpr (NUM_REGS == 8)
    asm volatile("tcgen05.ld.sync.aligned%9.x%10.b32 "
                 "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7}, [%8];"
                 : "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3]), "=f"(tmp[4]), "=f"(tmp[5]), "=f"(tmp[6]), "=f"(tmp[7])
                 : "r"(addr), "C"(SHAPE), "n"(NUM));
  if constexpr (NUM_REGS == 16)
    asm volatile("tcgen05.ld.sync.aligned%17.x%18.b32 "
                 "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7, "
                 "  %8,  %9, %10, %11, %12, %13, %14, %15}, [%16];"
                 : "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3]), "=f"(tmp[4]), "=f"(tmp[5]), "=f"(tmp[6]), "=f"(tmp[7]),
                   "=f"(tmp[8]), "=f"(tmp[9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15])
                 : "r"(addr), "C"(SHAPE), "n"(NUM));
  if constexpr (NUM_REGS == 32)
    asm volatile("tcgen05.ld.sync.aligned%33.x%34.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];"
                 : "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3]), "=f"(tmp[4]), "=f"(tmp[5]), "=f"(tmp[6]), "=f"(tmp[7]),
                   "=f"(tmp[8]), "=f"(tmp[9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
                   "=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
                   "=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31])
                 : "r"(addr), "C"(SHAPE), "n"(NUM));
  if constexpr (NUM_REGS == 64)
    asm volatile("tcgen05.ld.sync.aligned%65.x%66.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];"
                 : "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3]), "=f"(tmp[4]), "=f"(tmp[5]), "=f"(tmp[6]), "=f"(tmp[7]),
                   "=f"(tmp[8]), "=f"(tmp[9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
                   "=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
                   "=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31]),
                   "=f"(tmp[32]), "=f"(tmp[33]), "=f"(tmp[34]), "=f"(tmp[35]), "=f"(tmp[36]), "=f"(tmp[37]), "=f"(tmp[38]), "=f"(tmp[39]),
                   "=f"(tmp[40]), "=f"(tmp[41]), "=f"(tmp[42]), "=f"(tmp[43]), "=f"(tmp[44]), "=f"(tmp[45]), "=f"(tmp[46]), "=f"(tmp[47]),
                   "=f"(tmp[48]), "=f"(tmp[49]), "=f"(tmp[50]), "=f"(tmp[51]), "=f"(tmp[52]), "=f"(tmp[53]), "=f"(tmp[54]), "=f"(tmp[55]),
                   "=f"(tmp[56]), "=f"(tmp[57]), "=f"(tmp[58]), "=f"(tmp[59]), "=f"(tmp[60]), "=f"(tmp[61]), "=f"(tmp[62]), "=f"(tmp[63])
                 : "r"(addr), "C"(SHAPE), "n"(NUM));
  if constexpr (NUM_REGS == 128)
    asm volatile("tcgen05.ld.sync.aligned%129.x%130.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, %65, %66, %67, %68, %69, %70, %71, "
                 " %72, %73, %74, %75, %76, %77, %78, %79, "
                 " %80, %81, %82, %83, %84, %85, %86, %87, "
                 " %88, %89, %90, %91, %92, %93, %94, %95, "
                 " %96, %97, %98,%99,%100,%101,%102,%103, "
                 "%104,%105,%106,%107,%108,%109,%110,%111, "
                 "%112,%113,%114,%115,%116,%117,%118,%119, "
                 "%120,%121,%122,%123,%124,%125,%126,%127}, [%128];"
                 : "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3]), "=f"(tmp[4]), "=f"(tmp[5]), "=f"(tmp[6]), "=f"(tmp[7]),
                   "=f"(tmp[8]), "=f"(tmp[9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
                   "=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
                   "=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31]),
                   "=f"(tmp[32]), "=f"(tmp[33]), "=f"(tmp[34]), "=f"(tmp[35]), "=f"(tmp[36]), "=f"(tmp[37]), "=f"(tmp[38]), "=f"(tmp[39]),
                   "=f"(tmp[40]), "=f"(tmp[41]), "=f"(tmp[42]), "=f"(tmp[43]), "=f"(tmp[44]), "=f"(tmp[45]), "=f"(tmp[46]), "=f"(tmp[47]),
                   "=f"(tmp[48]), "=f"(tmp[49]), "=f"(tmp[50]), "=f"(tmp[51]), "=f"(tmp[52]), "=f"(tmp[53]), "=f"(tmp[54]), "=f"(tmp[55]),
                   "=f"(tmp[56]), "=f"(tmp[57]), "=f"(tmp[58]), "=f"(tmp[59]), "=f"(tmp[60]), "=f"(tmp[61]), "=f"(tmp[62]), "=f"(tmp[63]),
                   "=f"(tmp[64]), "=f"(tmp[65]), "=f"(tmp[66]), "=f"(tmp[67]), "=f"(tmp[68]), "=f"(tmp[69]), "=f"(tmp[70]), "=f"(tmp[71]),
                   "=f"(tmp[72]), "=f"(tmp[73]), "=f"(tmp[74]), "=f"(tmp[75]), "=f"(tmp[76]), "=f"(tmp[77]), "=f"(tmp[78]), "=f"(tmp[79]),
                   "=f"(tmp[80]), "=f"(tmp[81]), "=f"(tmp[82]), "=f"(tmp[83]), "=f"(tmp[84]), "=f"(tmp[85]), "=f"(tmp[86]), "=f"(tmp[87]),
                   "=f"(tmp[88]), "=f"(tmp[89]), "=f"(tmp[90]), "=f"(tmp[91]), "=f"(tmp[92]), "=f"(tmp[93]), "=f"(tmp[94]), "=f"(tmp[95]),
                   "=f"(tmp[96]), "=f"(tmp[97]), "=f"(tmp[98]), "=f"(tmp[99]), "=f"(tmp[100]), "=f"(tmp[101]), "=f"(tmp[102]), "=f"(tmp[103]),
                   "=f"(tmp[104]), "=f"(tmp[105]), "=f"(tmp[106]), "=f"(tmp[107]), "=f"(tmp[108]), "=f"(tmp[109]), "=f"(tmp[110]), "=f"(tmp[111]),
                   "=f"(tmp[112]), "=f"(tmp[113]), "=f"(tmp[114]), "=f"(tmp[115]), "=f"(tmp[116]), "=f"(tmp[117]), "=f"(tmp[118]), "=f"(tmp[119]),
                   "=f"(tmp[120]), "=f"(tmp[121]), "=f"(tmp[122]), "=f"(tmp[123]), "=f"(tmp[124]), "=f"(tmp[125]), "=f"(tmp[126]), "=f"(tmp[127])
                 : "r"(addr), "C"(SHAPE), "n"(NUM));
}
template <int num>
__device__ inline void
tcgen05_ld_32x32b(float *tmp, uint32_t tmem_addr, int row, int col)
{
  tcgen05_ld<num, SHAPE::_32x32b, num>(tmp, tmem_addr, row, col);
}
template <int num>
__device__ inline void tcgen05_ld_16x128b(float *tmp, uint32_t tmem_addr, int row, int col)
{
  tcgen05_ld<num * 2, SHAPE::_16x128b, num>(tmp, tmem_addr, row, col);
}
template <int num>
__device__ inline void tcgen05_ld_16x256b(float *tmp, uint32_t tmem_addr, int row, int col)
{
  tcgen05_ld<num * 4, SHAPE::_16x256b, num>(tmp, tmem_addr, row, col);
}
template <typename T>
__device__ __inline__ T warp_uniform(T x) { return __shfl_sync(0xFFFF'FFFF, x, 0); }
__device__ inline int cluster_cta_rank()
{
  int rank;
  asm volatile("mov.u32 %0, %%cluster_ctarank;" : "=r"(rank));
  return rank;
}
__device__ inline int cluster_dim_x()
{
  int dim;
  asm volatile("mov.u32 %0, %%cluster_nctaid.x;" : "=r"(dim));
  return dim;
}
__device__ inline int cluster_dim_y()
{
  int dim;
  asm volatile("mov.u32 %0, %%cluster_nctaid.y;" : "=r"(dim));
  return dim;
}
__device__ inline int cluster_cta_x()
{
  int x;
  asm volatile("mov.u32 %0, %%cluster_ctaid.x;" : "=r"(x));
  return x;
}
__device__ inline int cluster_cta_y()
{
  int y;
  asm volatile("mov.u32 %0, %%cluster_ctaid.y;" : "=r"(y));
  return y;
}
__device__ inline void cluster_arrive_relaxed()
{
  asm volatile("barrier.cluster.arrive.relaxed.aligned;");
}
__device__ inline void cluster_wait_acquire()
{
  asm volatile("barrier.cluster.wait.acquire.aligned;");
}
__device__ inline void cluster_sync()
{
  cluster_arrive_relaxed();
  cluster_wait_acquire();
}
template <int CTA_GROUP = 1>
__device__ inline void tcgen05_alloc(int smem_holding_buf_addr, int num_cols)
{
  if constexpr (CTA_GROUP == 1)
  {
    asm volatile(
        "tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" ::"r"(smem_holding_buf_addr), "r"(num_cols)
        : "memory"
    );
  }
  else
  {
    asm volatile(
        "tcgen05.alloc.cta_group::2.sync.aligned.shared::cta.b32 [%0], %1;" ::"r"(smem_holding_buf_addr), "r"(num_cols)
        : "memory"
    );
  }
}
template <int CTA_GROUP = 1>
__device__ inline void tcgen05_dealloc(uint32_t tmem_addr, int num_cols)
{
  if constexpr (CTA_GROUP == 1)
  {
    asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" ::"r"(tmem_addr), "r"(num_cols));
  }
  else
  {
    asm volatile("tcgen05.dealloc.cta_group::2.sync.aligned.b32 %0, %1;" ::"r"(tmem_addr), "r"(num_cols));
  }
}
template <int CTA_GROUP = 1>
__device__ inline void tcgen05_relinquish_alloc_permit()
{
  if constexpr (CTA_GROUP == 1)
  {
    asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;");
  }
  else
  {
    asm volatile("tcgen05.relinquish_alloc_permit.cta_group::2.sync.aligned;");
  }
}
__device__ inline void tcgen05_wait_alloc()
{
  asm volatile("tcgen05.wait::ld.sync.aligned;");
}
constexpr int BLOCK_M = 128;
constexpr int BLOCK_N = 128;
constexpr int BLOCK_K = 256;
constexpr int MMA_K = 64;
constexpr int WARP_SIZE = 32;
constexpr int NUM_WARPS = 6;
constexpr int THREADS_PER_CTA = NUM_WARPS * WARP_SIZE;
constexpr int TMA_WARP = 4;
constexpr int MMA_WARP = 5;
constexpr int NUM_STAGES_MAX = 6;
constexpr int NUM_STAGES_DEFAULT = 5;
constexpr int CLUSTER_M = 1;
constexpr int CLUSTER_N = 1;
constexpr int CLUSTER_Z = 1;
constexpr int GROUP_M = 8;
constexpr int GROUP_N = 8;
struct GroupParams
{
  void *A_ptr;
  void *B_ptr;
  void *C_ptr;
  void *SFA_ptr;
  void *SFB_ptr;
  int M, N, K, L;
  uint32_t tile_offset;
  uint16_t tiles_m;
  uint16_t tiles_n;
  uint16_t padding;
};
constexpr int MAX_GROUPS = 16;
constexpr int TENSORMAPS_PER_GROUP = 5;
constexpr int TMAP_A_IDX = 0;
constexpr int TMAP_B_IDX = 1;
constexpr int TMAP_C_IDX = 2;
constexpr int TMAP_SFA_IDX = 3;
constexpr int TMAP_SFB_IDX = 4;
struct alignas(128) KernelParams {
    GroupParams groups[MAX_GROUPS];
    CUtensorMap tmaps[MAX_GROUPS * TENSORMAPS_PER_GROUP];
    uint32_t total_tiles;
    uint32_t num_groups;
};
constexpr int SF_VEC_SIZE = 16;
constexpr int NUM_MMA_K_ITERS = BLOCK_K / MMA_K;
constexpr int SF_SBO = 128;
constexpr int SF_SMEM_ADVANCE = 512;
constexpr int SF_TMEM_ADVANCE = 4;
constexpr int A_K_STRIDE = 2;
constexpr int B_K_STRIDE = 2;
constexpr int TMEM_ACC_COLS = 128;
constexpr int TMEM_SFA_COLS_LOGICAL = NUM_MMA_K_ITERS * 4;
constexpr int TMEM_SFB_COLS_LOGICAL = NUM_MMA_K_ITERS * 4;
constexpr int TMEM_SFA_COLS = 32;
constexpr int TMEM_SFB_COLS = 32;
constexpr int TMEM_TOTAL_COLS = TMEM_ACC_COLS + TMEM_SFA_COLS + TMEM_SFB_COLS;
constexpr int TMEM_ALLOC_COLS = TMEM_TOTAL_COLS;
constexpr int SMEM_A_SIZE = BLOCK_M * (BLOCK_K / 2);
constexpr int SMEM_B_SIZE = BLOCK_N * (BLOCK_K / 2);
constexpr int SMEM_SFA_SIZE = BLOCK_M * (BLOCK_K / 16);
constexpr int SMEM_SFB_SIZE = BLOCK_N * (BLOCK_K / 16);
constexpr int SMEM_STAGE_SIZE = SMEM_A_SIZE + SMEM_B_SIZE + SMEM_SFA_SIZE + SMEM_SFB_SIZE;
constexpr int TMA_A_BYTES = SMEM_A_SIZE;
constexpr int TMA_B_BYTES = SMEM_B_SIZE;
constexpr int TMA_SFA_BYTES = SMEM_SFA_SIZE;
constexpr int TMA_SFB_BYTES = SMEM_SFB_SIZE;
constexpr int TMA_BYTES_AB = TMA_A_BYTES + TMA_B_BYTES;
constexpr int TMA_BYTES_ALL = TMA_A_BYTES + TMA_B_BYTES + TMA_SFA_BYTES + TMA_SFB_BYTES;
struct SmemBuffers
{
  struct AlignedBuffA
  {
    alignas(128) char data[SMEM_A_SIZE];
  };
  struct AlignedBuffB
  {
    alignas(128) char data[SMEM_B_SIZE];
  };
  struct AlignedBuffSFA
  {
    alignas(128) char data[SMEM_SFA_SIZE];
  };
  struct AlignedBuffSFB
  {
    alignas(128) char data[SMEM_SFB_SIZE];
  };
  AlignedBuffA A_smem[NUM_STAGES_MAX];
  AlignedBuffB B_smem[NUM_STAGES_MAX];
  AlignedBuffSFA SFA_smem[NUM_STAGES_MAX];
  AlignedBuffSFB SFB_smem[NUM_STAGES_MAX];
  alignas(8) uint64_t full_mbar[NUM_STAGES_MAX];
  alignas(8) uint64_t empty_mbar[NUM_STAGES_MAX];
  alignas(8) uint64_t epilogue_mbar;
  alignas(8) uint64_t tmem_holding_buf;
};
__device__ inline uint64_t make_smem_desc_A(const void *smem_ptr)
{
  uint32_t addr = static_cast<uint32_t>(__cvta_generic_to_shared(smem_ptr));
  constexpr int SBO = 8 * 128;
  uint64_t desc = desc_encode(addr)
                  | (desc_encode(SBO) << 32ULL)
                  | (1ULL << 46ULL)
                  | (2ULL << 61ULL);
  return desc;
}
__device__ inline uint64_t make_smem_desc_B(const void *smem_ptr)
{
  uint32_t addr = static_cast<uint32_t>(__cvta_generic_to_shared(smem_ptr));
  constexpr int SBO = 8 * 128;
  uint64_t desc = desc_encode(addr)
                  | (desc_encode(SBO) << 32ULL)
                  | (1ULL << 46ULL)
                  | (2ULL << 61ULL);
  return desc;
}
__device__ inline uint64_t make_smem_desc(const void *smem_ptr, int row_stride_bytes)
{
  (void)row_stride_bytes;
  return make_smem_desc_B(smem_ptr);
}
__device__ inline uint64_t make_sf_smem_desc(const void *smem_ptr)
{
  uint32_t addr = static_cast<uint32_t>(__cvta_generic_to_shared(smem_ptr));
  constexpr uint64_t SBO = SF_SBO;
  uint64_t desc = desc_encode(addr) | (desc_encode(SBO) << 32ULL)
                  | (1ULL << 46ULL);
  return desc;
}
__device__ inline uint32_t make_mma_idesc(int mma_n = BLOCK_N)
{
  constexpr uint32_t MMA_M = BLOCK_M;
  const uint32_t MMA_N = mma_n;
  uint32_t idesc = (1U << 7U)
                             | (1U << 10U)
                             | (0U << 14U)
                             | ((MMA_N >> 3U) << 17U)
                             | (0U << 23U)
                             | ((MMA_M >> 7U) << 27U);
  return idesc;
}

"""

CUDA_SRC_KERNEL = r"""
template <int CLUSTER_N_PARAM, int NUM_STAGES_PARAM, int K_PARAM = 0, int BLOCK_N_PARAM = 128>
__global__ void __cluster_dims__(CLUSTER_M, CLUSTER_N_PARAM, CLUSTER_Z)
    __launch_bounds__(THREADS_PER_CTA)
        group_gemm_kernel_impl(
            const __grid_constant__ KernelParams kparams)
{
    constexpr int CLUSTER_N = CLUSTER_N_PARAM;
    constexpr int NUM_STAGES = NUM_STAGES_PARAM;
    constexpr int K_EXPECTED = K_PARAM;
    constexpr int BLOCK_N = BLOCK_N_PARAM;
    const int tid = threadIdx.x;
    const int warp_id = tid / WARP_SIZE;
    const int lane_id = tid % WARP_SIZE;
    const int cta_n = cluster_cta_y();
    const uint32_t work_idx = blockIdx.y / CLUSTER_N;
    if (work_idx >= kparams.total_tiles)
        return;
    extern __shared__ char smem_raw[];
    uintptr_t smem_addr = reinterpret_cast<uintptr_t>(smem_raw);
    uintptr_t aligned_addr = (smem_addr + 127) & ~uintptr_t(127);
    SmemBuffers *smem = reinterpret_cast<SmemBuffers *>(aligned_addr);
    auto get_mbar_addr = [](void *mbar) -> int
    {
        return static_cast<int>(__cvta_generic_to_shared(mbar));
    };
    if (tid == 0)
    {
        for (int s = 0; s < NUM_STAGES; s++)
        {
            mbarrier_init(get_mbar_addr(&smem->full_mbar[s]), 1);
            mbarrier_init(get_mbar_addr(&smem->empty_mbar[s]), 1);
        }
        mbarrier_init(get_mbar_addr(&smem->epilogue_mbar), 1);
        fence_mbarrier_init();
    }
    __syncthreads();
    __shared__ uint32_t tmem_base_addr;
    __shared__ uint32_t tmem_sfa_addr;
    __shared__ uint32_t tmem_sfb_addr;
    __shared__ uint32_t tmem_idesc;
    constexpr int16_t MCAST_MASK_ALL = (CLUSTER_N == 4) ? 0xF : (CLUSTER_N == 2) ? 0x3
                                                                                 : 0x1;
    uint32_t acc_tmem = 0, sfa_tmem = 0, sfb_tmem = 0;
    uint32_t idesc = 0;
    int group_id = 0;
    GroupParams params;
    int tile_m = 0, tile_n = 0;
    int num_k_iters = 0;
    int my_n_tile = 0;
    int total_n_tiles = 0;
    bool has_valid_work = false;
    {
        int lo = 0, hi = (int)kparams.num_groups - 1;
        while (lo < hi)
        {
            int mid = (lo + hi + 1) / 2;
            if (kparams.groups[mid].tile_offset <= work_idx)
            {
                lo = mid;
            }
            else
            {
                hi = mid - 1;
            }
        }
        group_id = lo;
    }
    params = kparams.groups[group_id];
    uint32_t local_idx = work_idx - params.tile_offset;
    tile_m = local_idx / params.tiles_n;
    tile_n = local_idx % params.tiles_n;
    if constexpr (K_EXPECTED > 0)
    {
        num_k_iters = K_EXPECTED / BLOCK_K;
    }
    else
    {
        num_k_iters = params.K / BLOCK_K;
    }
    my_n_tile = tile_n * CLUSTER_N + cta_n;
    total_n_tiles = params.N / BLOCK_N;
    has_valid_work = (my_n_tile < total_n_tiles);
    const void *tmap_A = nullptr;
    const void *tmap_B = nullptr;
    const void *tmap_SFA = nullptr;
    const void *tmap_SFB = nullptr;
    int coord_m = 0;
    int coord_n = 0;
    constexpr bool use_multicast = (CLUSTER_N > 1);
    if (has_valid_work)
    {
        tmap_A = &kparams.tmaps[group_id * TENSORMAPS_PER_GROUP + TMAP_A_IDX];
        tmap_B = &kparams.tmaps[group_id * TENSORMAPS_PER_GROUP + TMAP_B_IDX];
        tmap_SFA = &kparams.tmaps[group_id * TENSORMAPS_PER_GROUP + TMAP_SFA_IDX];
        tmap_SFB = &kparams.tmaps[group_id * TENSORMAPS_PER_GROUP + TMAP_SFB_IDX];
        if (tid < 4) {
            prefetch_tensormap(&kparams.tmaps[group_id * TENSORMAPS_PER_GROUP + tid]);
        }
        coord_m = tile_m * BLOCK_M;
        coord_n = my_n_tile * BLOCK_N;
    }
    __syncthreads();
    if (has_valid_work)
    {
        if (warp_id == MMA_WARP)
        {
            int holding_buf_addr = get_mbar_addr(&smem->tmem_holding_buf);
            tcgen05_alloc<1>(holding_buf_addr, TMEM_ACC_COLS);
            tcgen05_wait_alloc();
            if (lane_id == 0)
                acc_tmem = *reinterpret_cast<volatile uint32_t *>(&smem->tmem_holding_buf);
            tcgen05_alloc<1>(holding_buf_addr, TMEM_SFA_COLS);
            tcgen05_wait_alloc();
            if (lane_id == 0)
                sfa_tmem = *reinterpret_cast<volatile uint32_t *>(&smem->tmem_holding_buf);
            tcgen05_alloc<1>(holding_buf_addr, TMEM_SFB_COLS);
            tcgen05_wait_alloc();
            if (lane_id == 0)
            {
                sfb_tmem = *reinterpret_cast<volatile uint32_t *>(&smem->tmem_holding_buf);
                tmem_base_addr = acc_tmem;
                tmem_sfa_addr = sfa_tmem;
                tmem_sfb_addr = sfb_tmem;
                tmem_idesc = make_mma_idesc(BLOCK_N);
            }
        }
        __syncthreads();
        if (warp_id == MMA_WARP)
        {
            acc_tmem = __shfl_sync(0xFFFFFFFF, tmem_base_addr, 0);
            sfa_tmem = __shfl_sync(0xFFFFFFFF, tmem_sfa_addr, 0);
            sfb_tmem = __shfl_sync(0xFFFFFFFF, tmem_sfb_addr, 0);
            idesc = __shfl_sync(0xFFFFFFFF, tmem_idesc, 0);
        }
    }
    __syncthreads();
    cluster_sync();
    for (int k_iter = 0; k_iter < num_k_iters; k_iter++)
    {
        int stage = k_iter % NUM_STAGES;
        int coord_k = k_iter * BLOCK_K;
        int A_smem_addr = static_cast<int>(__cvta_generic_to_shared(smem->A_smem[stage].data));
        int B_smem_addr = static_cast<int>(__cvta_generic_to_shared(smem->B_smem[stage].data));
        int SFA_smem_addr = static_cast<int>(__cvta_generic_to_shared(smem->SFA_smem[stage].data));
        int SFB_smem_addr = static_cast<int>(__cvta_generic_to_shared(smem->SFB_smem[stage].data));
        int full_mbar_addr = get_mbar_addr(&smem->full_mbar[stage]);
        if (has_valid_work && warp_id == TMA_WARP)
        {
            if (elect_sync())
            {
                mbarrier_arrive_expect_tx(full_mbar_addr, TMA_BYTES_ALL);
            }
        }
        if (use_multicast)
        {
            cluster_sync();
        }
        if (has_valid_work && warp_id == TMA_WARP)
        {
            if (elect_sync())
            {
                if (k_iter >= NUM_STAGES)
                {
                    int empty_mbar_addr = get_mbar_addr(&smem->empty_mbar[stage]);
                    int phase = ((k_iter / NUM_STAGES) + 1) & 1;
                    mbarrier_wait(empty_mbar_addr, phase);
                }
                tma_2d_gmem2smem<1>(B_smem_addr, tmap_B, coord_k, coord_n, full_mbar_addr, EVICT_FIRST);
                int num_k_chunks = params.K / BLOCK_K;
                int off_sfb = (my_n_tile * num_k_chunks + k_iter) * SMEM_SFB_SIZE;
                tma_1d_gmem2smem<1>(SFB_smem_addr, tmap_SFB, off_sfb / 8, full_mbar_addr, EVICT_FIRST);
                if (use_multicast)
                {
                    if (cta_n == 0)
                    {
                        tma_2d_gmem2smem_mcast<1>(A_smem_addr, tmap_A, coord_k, coord_m,
                                                  full_mbar_addr, MCAST_MASK_ALL, EVICT_LAST);
                        int off_sfa = (tile_m * num_k_chunks + k_iter) * SMEM_SFA_SIZE;
                        tma_1d_gmem2smem_mcast<1>(SFA_smem_addr, tmap_SFA, off_sfa / 8,
                                                  full_mbar_addr, MCAST_MASK_ALL, EVICT_FIRST);
                    }
                }
                else
                {
                    tma_2d_gmem2smem<1>(A_smem_addr, tmap_A, coord_k, coord_m, full_mbar_addr, EVICT_LAST);
                    int off_sfa = (tile_m * num_k_chunks + k_iter) * SMEM_SFA_SIZE;
                    tma_1d_gmem2smem<1>(SFA_smem_addr, tmap_SFA, off_sfa / 8, full_mbar_addr, EVICT_FIRST);
                }
            }
        }
        if (has_valid_work && warp_id == MMA_WARP)
        {
            int full_mbar_addr = get_mbar_addr(&smem->full_mbar[stage]);
            int empty_mbar_addr = get_mbar_addr(&smem->empty_mbar[stage]);
            int phase = (k_iter / NUM_STAGES) & 1;
            mbarrier_wait(full_mbar_addr, phase);
            uint32_t sfa_tmem_base = tmem_sfa_addr;
            uint32_t sfb_tmem_base = tmem_sfb_addr;
            uint64_t a_desc = make_smem_desc_A(smem->A_smem[stage].data);
            uint64_t b_desc = make_smem_desc_B(smem->B_smem[stage].data);
            {
                constexpr int k = 0;
                uint64_t sfa_desc_0 = make_sf_smem_desc(
                    reinterpret_cast<const char *>(smem->SFA_smem[stage].data));
                uint64_t sfb_desc_0 = make_sf_smem_desc(
                    reinterpret_cast<const char *>(smem->SFB_smem[stage].data));
                uint64_t sfa_desc_1 = make_sf_smem_desc(
                    reinterpret_cast<const char *>(smem->SFA_smem[stage].data) + SF_SMEM_ADVANCE);
                uint64_t sfb_desc_1 = make_sf_smem_desc(
                    reinterpret_cast<const char *>(smem->SFB_smem[stage].data) + SF_SMEM_ADVANCE);
                if (elect_sync())
                {
                    tcgen05_cp_nvfp4<1>(sfa_tmem_base, sfa_desc_0);
                    tcgen05_cp_nvfp4<1>(sfb_tmem_base, sfb_desc_0);
                }
                tcgen05_fence_before_thread_sync();
                if (elect_sync())
                {
                    tcgen05_cp_nvfp4<1>(sfa_tmem_base + 1 * SF_TMEM_ADVANCE, sfa_desc_1);
                    tcgen05_cp_nvfp4<1>(sfb_tmem_base + 1 * SF_TMEM_ADVANCE, sfb_desc_1);
                }
                if (elect_sync())
                {
                    int enable_input_d = (k_iter > 0) ? 1 : 0;
                    tcgen05_mma_nvfp4<1>(acc_tmem, a_desc, b_desc, idesc,
                                         sfa_tmem_base, sfb_tmem_base, enable_input_d);
                }
                a_desc += A_K_STRIDE;
                b_desc += B_K_STRIDE;
            }
#pragma unroll
            for (int k = 1; k < NUM_MMA_K_ITERS; k++)
            {
                tcgen05_fence_before_thread_sync();
                if (k + 1 < NUM_MMA_K_ITERS)
                {
                    int sfa_smem_offset = (k + 1) * SF_SMEM_ADVANCE;
                    int sfb_smem_offset = (k + 1) * SF_SMEM_ADVANCE;
                    int next_scale_A_tmem = sfa_tmem_base + (k + 1) * SF_TMEM_ADVANCE;
                    int next_scale_B_tmem = sfb_tmem_base + (k + 1) * SF_TMEM_ADVANCE;
                    uint64_t sfa_desc_next = make_sf_smem_desc(
                        reinterpret_cast<const char *>(smem->SFA_smem[stage].data) + sfa_smem_offset);
                    uint64_t sfb_desc_next = make_sf_smem_desc(
                        reinterpret_cast<const char *>(smem->SFB_smem[stage].data) + sfb_smem_offset);
                    if (elect_sync())
                    {
                        tcgen05_cp_nvfp4<1>(next_scale_A_tmem, sfa_desc_next);
                        tcgen05_cp_nvfp4<1>(next_scale_B_tmem, sfb_desc_next);
                    }
                }
                int scale_A_tmem = sfa_tmem_base + k * SF_TMEM_ADVANCE;
                int scale_B_tmem = sfb_tmem_base + k * SF_TMEM_ADVANCE;
                if (elect_sync())
                {
                    tcgen05_mma_nvfp4<1>(acc_tmem, a_desc, b_desc, idesc,
                                         scale_A_tmem, scale_B_tmem, 1);
                }
                a_desc += A_K_STRIDE;
                b_desc += B_K_STRIDE;
            }
            if (elect_sync())
            {
                tcgen05_commit<1>(empty_mbar_addr);
            }
        }
    }
    __syncthreads();
    if (has_valid_work && warp_id == MMA_WARP)
    {
        tcgen05_fence_before_thread_sync();
        if (lane_id == 0)
        {
            mbarrier_arrive(get_mbar_addr(&smem->epilogue_mbar));
        }
    }
    if (has_valid_work && warp_id < 4)
    {
        mbarrier_wait(get_mbar_addr(&smem->epilogue_mbar), 0);
        tcgen05_fence_after_thread_sync();
        half *C_ptr = reinterpret_cast<half *>(params.C_ptr);
        int M = params.M;
        int N = params.N;
        int row_tile = warp_id;
        int base_row = row_tile * 32;
        for (int c_chunk = 0; c_chunk < BLOCK_N / 8; c_chunk++)
        {
            int base_col = c_chunk * 8;
            float tmp[8];
            tcgen05_ld_32x32b<8>(tmp, tmem_base_addr, base_row, base_col);
            tcgen05_wait_alloc();
            int local_row = base_row + lane_id;
            int global_row = coord_m + local_row;
            if (global_row < M)
            {
                half2 h2_tmp[4];
                #pragma unroll
                for (int i = 0; i < 4; i++)
                {
                    h2_tmp[i] = __float22half2_rn({tmp[i * 2], tmp[i * 2 + 1]});
                }
                reinterpret_cast<int4 *>(C_ptr + global_row * N + coord_n + base_col)[0] =
                    *reinterpret_cast<int4 *>(h2_tmp);
            }
        }
    }
    __syncthreads();
    if (has_valid_work && warp_id == MMA_WARP)
    {
        tcgen05_relinquish_alloc_permit<1>();
        tcgen05_dealloc<1>(sfb_tmem, TMEM_SFB_COLS);
        tcgen05_dealloc<1>(sfa_tmem, TMEM_SFA_COLS);
        tcgen05_dealloc<1>(acc_tmem, TMEM_ACC_COLS);
        tcgen05_wait_alloc();
    }
    __syncthreads();
    if constexpr (CLUSTER_N > 1)
    {
        cluster_sync();
    }
}
template __global__ void group_gemm_kernel_impl<1, 5, 0, 128>(const __grid_constant__ KernelParams kparams);
template __global__ void group_gemm_kernel_impl<1, 6, 0, 128>(const __grid_constant__ KernelParams kparams);
template __global__ void group_gemm_kernel_impl<2, 5, 0, 128>(const __grid_constant__ KernelParams kparams);
template __global__ void group_gemm_kernel_impl<2, 6, 0, 128>(const __grid_constant__ KernelParams kparams);
template __global__ void group_gemm_kernel_impl<4, 5, 0, 128>(const __grid_constant__ KernelParams kparams);
template __global__ void group_gemm_kernel_impl<4, 6, 0, 128>(const __grid_constant__ KernelParams kparams);
"""

CUDA_SRC_WRAPPER = r"""
#include <torch/library.h>
#include <ATen/core/Tensor.h>
#include <cuda.h>
#include <cuda_runtime.h>
template <int CLUSTER_N_PARAM, int NUM_STAGES_PARAM, int K_PARAM, int BLOCK_N_PARAM>
__global__ void group_gemm_kernel_impl(
    const __grid_constant__ KernelParams kparams);
constexpr int SMEM_A_ALIGNED = ((SMEM_A_SIZE + 127) / 128) * 128;
constexpr int SMEM_B_ALIGNED = ((SMEM_B_SIZE + 127) / 128) * 128;
constexpr int SMEM_SFA_ALIGNED = ((SMEM_SFA_SIZE + 127) / 128) * 128;
constexpr int SMEM_SFB_ALIGNED = ((SMEM_SFB_SIZE + 127) / 128) * 128;
constexpr int SMEM_DATA_SIZE = NUM_STAGES_MAX * (SMEM_A_ALIGNED + SMEM_B_ALIGNED + SMEM_SFA_ALIGNED + SMEM_SFB_ALIGNED);
constexpr int SMEM_BARRIER_SIZE = sizeof(uint64_t) * (NUM_STAGES_MAX * 2 + 2);
constexpr int SMEM_SIZE = SMEM_DATA_SIZE + SMEM_BARRIER_SIZE + 128;
void check_cu(CUresult err, const char *context)
{
    if (err == CUDA_SUCCESS)
        return;
    const char *error_msg_ptr;
    if (cuGetErrorString(err, &error_msg_ptr) != CUDA_SUCCESS)
        error_msg_ptr = "unable to get error string";
    TORCH_CHECK(false, context, ": ", error_msg_ptr);
}
void check_cuda(cudaError_t err, const char *context)
{
    if (err == cudaSuccess)
        return;
    TORCH_CHECK(false, context, ": ", cudaGetErrorString(err));
}
void init_A_tmap(
    CUtensorMap *tmap,
    const void *ptr,
    int M, int K, int L,
    int block_m, int block_k)
{
    (void)L;
    constexpr uint32_t rank = 2;
    uint64_t globalDim[rank] = {
        (uint64_t)(K),
        (uint64_t)(M)
    };
    uint64_t globalStrides[rank - 1] = {
        (uint64_t)(K / 2)
    };
    uint32_t boxDim[rank] = {
        (uint32_t)(block_k),
        (uint32_t)(block_m)
    };
    uint32_t elementStrides[rank] = {1, 1};
    auto err = cuTensorMapEncodeTiled(
        tmap,
        CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
        rank,
        const_cast<void *>(ptr),
        globalDim,
        globalStrides,
        boxDim,
        elementStrides,
        CU_TENSOR_MAP_INTERLEAVE_NONE,
        CU_TENSOR_MAP_SWIZZLE_128B,
        CU_TENSOR_MAP_L2_PROMOTION_NONE,
        CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
    check_cu(err, "cuTensorMapEncodeTiled for A (16U4, rank-2)");
}
void init_B_tmap(
    CUtensorMap *tmap,
    const void *ptr,
    int N, int K, int L,
    int block_n, int block_k)
{
    (void)L;
    constexpr uint32_t rank = 2;
    uint64_t globalDim[rank] = {
        (uint64_t)(K),
        (uint64_t)(N)
    };
    uint64_t globalStrides[rank - 1] = {
        (uint64_t)(K / 2)
    };
    uint32_t boxDim[rank] = {
        (uint32_t)(block_k),
        (uint32_t)(block_n)
    };
    uint32_t elementStrides[rank] = {1, 1};
    auto err = cuTensorMapEncodeTiled(
        tmap,
        CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
        rank,
        const_cast<void *>(ptr),
        globalDim,
        globalStrides,
        boxDim,
        elementStrides,
        CU_TENSOR_MAP_INTERLEAVE_NONE,
        CU_TENSOR_MAP_SWIZZLE_128B,
        CU_TENSOR_MAP_L2_PROMOTION_L2_64B,
        CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
    check_cu(err, "cuTensorMapEncodeTiled for B (16U4, rank-2)");
}
void init_C_tmap(
    CUtensorMap *tmap,
    const void *ptr,
    int M, int N, int L,
    int block_m, int block_n)
{
    constexpr uint32_t rank = 3;
    uint64_t globalDim[rank] = {
        (uint64_t)(N),
        (uint64_t)(M),
        (uint64_t)(L)
    };
    uint64_t globalStrides[rank - 1] = {
        (uint64_t)(N * sizeof(half)),
        (uint64_t)(M * N * sizeof(half))
    };
    uint32_t boxDim[rank] = {
        (uint32_t)(block_n < N ? block_n : N),
        (uint32_t)(block_m < M ? block_m : M),
        1
    };
    uint32_t elementStrides[rank] = {1, 1, 1};
    auto err = cuTensorMapEncodeTiled(
        tmap,
        CU_TENSOR_MAP_DATA_TYPE_FLOAT16,
        rank,
        const_cast<void *>(ptr),
        globalDim,
        globalStrides,
        boxDim,
        elementStrides,
        CU_TENSOR_MAP_INTERLEAVE_NONE,
        CU_TENSOR_MAP_SWIZZLE_NONE,
        CU_TENSOR_MAP_L2_PROMOTION_NONE,
        CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
    check_cu(err, "cuTensorMapEncodeTiled for C");
}
void init_SF_tmap(
    CUtensorMap *tmap,
    const void *ptr,
    int M, int K)
{
    constexpr uint32_t rank = 1;
    uint64_t global_size = (uint64_t)M * K / 16;
    uint64_t shared_size = SMEM_SFA_SIZE;
    uint64_t globalDim[rank] = {global_size / 8};
    uint64_t globalStrides[rank - 1] = {};
    uint32_t boxDim[rank] = {(uint32_t)(shared_size / 8)};
    uint32_t elementStrides[rank] = {1};
    auto err = cuTensorMapEncodeTiled(
        tmap,
        CU_TENSOR_MAP_DATA_TYPE_INT64,
        rank,
        const_cast<void *>(ptr),
        globalDim,
        globalStrides,
        boxDim,
        elementStrides,
        CU_TENSOR_MAP_INTERLEAVE_NONE,
        CU_TENSOR_MAP_SWIZZLE_NONE,
        CU_TENSOR_MAP_L2_PROMOTION_NONE,
        CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
    check_cu(err, "cuTensorMapEncodeTiled for SF (1D)");
}
int compute_cluster_n(const at::Tensor &problem_sizes)
{
    int num_groups = problem_sizes.size(0);
    auto sizes_acc = problem_sizes.accessor<int32_t, 2>();
    auto all_valid = [&](int cluster_n)
    {
        for (int g = 0; g < num_groups; g++)
        {
            int N = sizes_acc[g][1];
            if (N < cluster_n * BLOCK_N || N % (cluster_n * BLOCK_N) != 0)
                return false;
        }
        return true;
    };
    if (all_valid(4))
        return 4;
    if (all_valid(2))
        return 2;
    return 1;
}
int compute_num_stages(const at::Tensor &problem_sizes)
{
    int num_groups = problem_sizes.size(0);
    auto sizes_acc = problem_sizes.accessor<int32_t, 2>();
    float min_ai = FLT_MAX;
    for (int g = 0; g < num_groups; g++)
    {
        long long M = sizes_acc[g][0];
        long long N = sizes_acc[g][1];
        long long K = sizes_acc[g][2];
        double ops = 2.0 * M * N * K;
        double bytes = (M * K + N * K) * 0.5 + M * N * 2.0;
        float ai = (float)(ops / bytes);
        if (ai < min_ai)
            min_ai = ai;
    }
    return (min_ai > 150.0f) ? 6 : 5;
}
struct LaunchCache {
    uintptr_t ptrs[MAX_GROUPS * 5];
    int dims[MAX_GROUPS * 4];
    int num_groups;
    int cluster_n;
    KernelParams kparams;
    bool valid = false;
};
void group_gemm_launch(
    const at::Tensor &abc_ptrs,
    const at::Tensor &sf_ptrs,
    const at::Tensor &problem_sizes)
{
    int num_groups = problem_sizes.size(0);
    TORCH_CHECK(num_groups <= MAX_GROUPS, "Too many groups: ", num_groups);
    int cluster_n = compute_cluster_n(problem_sizes);
    int num_stages = compute_num_stages(problem_sizes);

    auto abc_acc = abc_ptrs.accessor<int64_t, 2>();
    auto sf_acc = sf_ptrs.accessor<int64_t, 2>();
    auto sizes_acc = problem_sizes.accessor<int32_t, 2>();
    static LaunchCache cache;
    bool hit = cache.valid && cache.num_groups == num_groups && cache.cluster_n == cluster_n;
    if (hit) {
        for (int g = 0; g < num_groups && hit; g++) {
            int gi = g * 5;
            hit = hit
                && cache.ptrs[gi]   == (uintptr_t)abc_acc[g][0]
                && cache.ptrs[gi+1] == (uintptr_t)abc_acc[g][1]
                && cache.ptrs[gi+2] == (uintptr_t)abc_acc[g][2]
                && cache.ptrs[gi+3] == (uintptr_t)sf_acc[g][0]
                && cache.ptrs[gi+4] == (uintptr_t)sf_acc[g][1]
                && cache.dims[g*4]   == sizes_acc[g][0]
                && cache.dims[g*4+1] == sizes_acc[g][1]
                && cache.dims[g*4+2] == sizes_acc[g][2]
                && cache.dims[g*4+3] == sizes_acc[g][3];
        }
    }
    KernelParams kparams;
    if (hit) {
        kparams = cache.kparams;
    } else {
        memset(&kparams, 0, sizeof(kparams));
        uint32_t total_tiles = 0;
        for (int g = 0; g < num_groups; g++)
        {
            int M = sizes_acc[g][0];
            int N = sizes_acc[g][1];
            int K = sizes_acc[g][2];
            int L = sizes_acc[g][3];
            int tiles_m = (M + BLOCK_M - 1) / BLOCK_M;
            int tiles_n = N / (BLOCK_N * cluster_n);
            int group_tiles = tiles_m * std::max(1, tiles_n);
            kparams.groups[g].tile_offset = total_tiles;
            kparams.groups[g].M = static_cast<uint16_t>(M);
            kparams.groups[g].N = static_cast<uint16_t>(N);
            kparams.groups[g].K = static_cast<uint16_t>(K);
            kparams.groups[g].tiles_m = static_cast<uint16_t>(tiles_m);
            kparams.groups[g].tiles_n = static_cast<uint16_t>(std::max(1, tiles_n));
            kparams.groups[g].L = static_cast<uint16_t>(L);
            kparams.groups[g].padding = 0;
            kparams.groups[g].A_ptr = reinterpret_cast<void *>(abc_acc[g][0]);
            kparams.groups[g].B_ptr = reinterpret_cast<void *>(abc_acc[g][1]);
            kparams.groups[g].C_ptr = reinterpret_cast<void *>(abc_acc[g][2]);
            kparams.groups[g].SFA_ptr = reinterpret_cast<void *>(sf_acc[g][0]);
            kparams.groups[g].SFB_ptr = reinterpret_cast<void *>(sf_acc[g][1]);
            total_tiles += group_tiles;
        }
        kparams.total_tiles = total_tiles;
        kparams.num_groups = static_cast<uint32_t>(num_groups);
        for (int g = 0; g < num_groups; g++)
        {
            int M = sizes_acc[g][0];
            int N = sizes_acc[g][1];
            int K = sizes_acc[g][2];
            int L = sizes_acc[g][3];
            void *A_ptr = reinterpret_cast<void *>(abc_acc[g][0]);
            void *B_ptr = reinterpret_cast<void *>(abc_acc[g][1]);
            void *C_ptr = reinterpret_cast<void *>(abc_acc[g][2]);
            void *SFA_ptr = reinterpret_cast<void *>(sf_acc[g][0]);
            void *SFB_ptr = reinterpret_cast<void *>(sf_acc[g][1]);
            int base_idx = g * TENSORMAPS_PER_GROUP;
            int num_m_tiles_a = (M + BLOCK_M - 1) / BLOCK_M;
            int num_m_tiles_b = (N + BLOCK_N - 1) / BLOCK_N;
            int M_padded = num_m_tiles_a * BLOCK_M;
            int N_padded = num_m_tiles_b * BLOCK_N;
            init_A_tmap(&kparams.tmaps[base_idx + TMAP_A_IDX], A_ptr, M_padded, K, L, BLOCK_M, BLOCK_K);
            init_B_tmap(&kparams.tmaps[base_idx + TMAP_B_IDX], B_ptr, N_padded, K, L, BLOCK_N, BLOCK_K);
            init_C_tmap(&kparams.tmaps[base_idx + TMAP_C_IDX], C_ptr, M, N, L, BLOCK_M, BLOCK_N);
            init_SF_tmap(&kparams.tmaps[base_idx + TMAP_SFA_IDX], SFA_ptr, M_padded, K);
            init_SF_tmap(&kparams.tmaps[base_idx + TMAP_SFB_IDX], SFB_ptr, N_padded, K);
        }
        cache.kparams = kparams;
        cache.num_groups = num_groups;
        cache.cluster_n = cluster_n;
        for (int g = 0; g < num_groups; g++) {
            int gi = g * 5;
            cache.ptrs[gi]   = (uintptr_t)abc_acc[g][0];
            cache.ptrs[gi+1] = (uintptr_t)abc_acc[g][1];
            cache.ptrs[gi+2] = (uintptr_t)abc_acc[g][2];
            cache.ptrs[gi+3] = (uintptr_t)sf_acc[g][0];
            cache.ptrs[gi+4] = (uintptr_t)sf_acc[g][1];
            cache.dims[g*4]   = sizes_acc[g][0];
            cache.dims[g*4+1] = sizes_acc[g][1];
            cache.dims[g*4+2] = sizes_acc[g][2];
            cache.dims[g*4+3] = sizes_acc[g][3];
        }
        cache.valid = true;
    }
    uint32_t grid_y = kparams.total_tiles * cluster_n;
    const void *kernel_func;
    if (cluster_n == 4)
    {
        if (num_stages >= 6)
            kernel_func = (const void *)group_gemm_kernel_impl<4, 6, 0, 128>;
        else
            kernel_func = (const void *)group_gemm_kernel_impl<4, 5, 0, 128>;
    }
    else if (cluster_n == 2)
    {
        if (num_stages >= 6)
            kernel_func = (const void *)group_gemm_kernel_impl<2, 6, 0, 128>;
        else
            kernel_func = (const void *)group_gemm_kernel_impl<2, 5, 0, 128>;
    }
    else
    {
        if (num_stages >= 6)
            kernel_func = (const void *)group_gemm_kernel_impl<1, 6, 0, 128>;
        else
            kernel_func = (const void *)group_gemm_kernel_impl<1, 5, 0, 128>;
    }
    static const void *configured_kernel = nullptr;
    if (configured_kernel != kernel_func)
    {
        check_cuda(cudaFuncSetAttribute(
                       kernel_func,
                       cudaFuncAttributeMaxDynamicSharedMemorySize,
                       SMEM_SIZE),
                   "cudaFuncSetAttribute for shared memory");
        configured_kernel = kernel_func;
    }
    dim3 grid(1, grid_y, 1);
    dim3 block(THREADS_PER_CTA, 1, 1);
    cudaLaunchConfig_t config = {};
    config.gridDim = grid;
    config.blockDim = block;
    config.dynamicSmemBytes = SMEM_SIZE;
    cudaLaunchAttribute attrs[1];
    attrs[0].id = cudaLaunchAttributeClusterDimension;
    attrs[0].val.clusterDim.x = CLUSTER_M;
    attrs[0].val.clusterDim.y = cluster_n;
    attrs[0].val.clusterDim.z = CLUSTER_Z;
    config.numAttrs = 1;
    config.attrs = attrs;
    using KernelFn = void(*)(const __grid_constant__ KernelParams);
    auto kfn = reinterpret_cast<KernelFn>(kernel_func);
    cudaError_t err = cudaLaunchKernelEx(&config, kfn, kparams);
    if (err != cudaSuccess)
    {
        TORCH_CHECK(false, "Kernel launch failed: ", cudaGetErrorString(err));
    }
}
TORCH_LIBRARY(group_gemm_module, m)
{
    m.def("group_gemm_launch(Tensor abc_ptrs, Tensor sf_ptrs, Tensor problem_sizes) -> ()");
    m.impl("group_gemm_launch", &group_gemm_launch);
}
"""

load_inline(
    name='group_gemm_module_compile_v7',
    cpp_sources='',
    cuda_sources=CUDA_SRC_UTILS + CUDA_SRC_KERNEL + CUDA_SRC_WRAPPER,
    verbose=True,
    is_python_module=False,
    no_implicit_headers=True,
    extra_cuda_cflags=[
        '-O3',
        '-gencode=arch=compute_100a,code=sm_100a',
        '--use_fast_math',
        '--expt-relaxed-constexpr',
        '--relocatable-device-code=false',
        '-lineinfo',
        '-diag-suppress=177',
        '-Xptxas=-v',
    ],
    extra_ldflags=['-lcuda']
)

_E2M1_LUT = torch.tensor([
    0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0,
    -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0
], dtype=torch.float16)

def dequantize_fp4(packed: torch.Tensor, rows: int, K: int, L: int, device) -> torch.Tensor:

    num_bytes = rows * (K // 2) * L

    packed_u8 = packed.contiguous().view(torch.uint8).reshape(num_bytes)

    low_nibble = (packed_u8 & 0x0F).long()
    high_nibble = ((packed_u8 >> 4) & 0x0F).long()

    lut = _E2M1_LUT.to(device)
    low_vals = lut[low_nibble]
    high_vals = lut[high_nibble]

    result = torch.stack([low_vals, high_vals], dim=1).reshape(-1)

    return result.reshape(rows, K, L)

def apply_scale_factors(tensor: torch.Tensor, sf: torch.Tensor, rows: int, K: int, L: int, device) -> torch.Tensor:

    sf_fp16 = sf.to(device=device, dtype=torch.float16).flatten()

    sf_k_blocks = K // 16
    total_sf = sf_fp16.numel()
    padded_rows = total_sf // (sf_k_blocks * L)

    sf_reshaped = sf_fp16.reshape(padded_rows, sf_k_blocks, L)

    sf_trimmed = sf_reshaped[:rows, :, :]

    sf_expanded = sf_trimmed.repeat_interleave(16, dim=1)

    return tensor * sf_expanded

def prepare_sf_for_tma(sf_raw, mn, k, block_size=128):

    BLOCK_M = 128

    sf_u8 = sf_raw.view(torch.uint8).clone()

    if len(sf_raw.shape) == 6:

        rest_m = sf_raw.shape[2]

        last_tile_m = mn - (rest_m - 1) * BLOCK_M
        if last_tile_m < BLOCK_M:
            for m4 in range(4):
                m4_base = m4 * 32
                if m4_base >= last_tile_m:
                    sf_u8[:, m4, rest_m - 1, :, :, :] = 0
                elif m4_base + 32 > last_tile_m:
                    valid = last_tile_m - m4_base
                    sf_u8[valid:, m4, rest_m - 1, :, :, :] = 0

        result = sf_u8.permute(5, 2, 4, 0, 1, 3).contiguous()
        return result.reshape(-1, 16)

    else:
        sf_k = k // 16
        L = sf_raw.shape[-1] if len(sf_raw.shape) > 2 else 1
        M = sf_u8.shape[0]
        M_padded = ((mn + block_size - 1) // block_size) * block_size
        if M < M_padded:
            pad_size = M_padded - M
            padding = torch.zeros((pad_size, sf_k, L), dtype=torch.uint8, device=sf_u8.device)
            sf_u8 = torch.cat([sf_u8, padding], dim=0)
        if L == 1:
            sf_u8 = sf_u8.squeeze(-1)
        return sf_u8.contiguous()

def pad_tensor_to_block(tensor, dim, block_size):

    current_size = tensor.shape[dim]
    target_size = ((current_size + block_size - 1) // block_size) * block_size
    if current_size == target_size:
        return tensor

    pad_shape = list(tensor.shape)
    pad_shape[dim] = target_size - current_size

    padding = torch.zeros(pad_shape, dtype=torch.uint8, device=tensor.device).view(tensor.dtype)
    return torch.cat([tensor, padding], dim=dim)

_launch_cache = dict()

def custom_kernel_cuda(data):

    abc_tensors, sfasfb_tensors, sfasfb_reordered, problem_sizes = data
    num_groups = len(problem_sizes)

    cache_key = tuple(
        (abc[0].data_ptr(), abc[1].data_ptr(), abc[2].data_ptr(),
         sf[0].data_ptr(), sf[1].data_ptr(), ps[0], ps[1], ps[2], ps[3])
        for abc, sf, ps in zip(abc_tensors, sfasfb_reordered, problem_sizes)
    )

    cached = _launch_cache.get(cache_key)
    if cached is not None:
        abc_ptrs_tensor, sf_ptrs_tensor, sizes_tensor, _refs = cached
        torch.ops.group_gemm_module.group_gemm_launch(abc_ptrs_tensor, sf_ptrs_tensor, sizes_tensor)
        return [t[2] for t in abc_tensors]

    abc_ptrs = []
    sf_ptrs = []
    sizes = []
    padded_tensors = []
    sf_prepared = []

    BLOCK_M = 128
    BLOCK_N = 128

    for i, ((a, b, c), (sfa_reordered, sfb_reordered), (M, N, K, L)) in enumerate(zip(abc_tensors, sfasfb_reordered, problem_sizes)):
        a_padded = pad_tensor_to_block(a, 0, BLOCK_M)
        b_padded = pad_tensor_to_block(b, 0, BLOCK_N)
        padded_tensors.append((a_padded, b_padded))
        abc_ptrs.append([a_padded.data_ptr(), b_padded.data_ptr(), c.data_ptr()])

        sfa_tma = prepare_sf_for_tma(sfa_reordered, M, K)
        sfb_tma = prepare_sf_for_tma(sfb_reordered, N, K)
        sf_prepared.append((sfa_tma, sfb_tma))
        sf_ptrs.append([sfa_tma.data_ptr(), sfb_tma.data_ptr()])
        sizes.append([M, N, K, L])

    abc_ptrs_tensor = torch.tensor(abc_ptrs, dtype=torch.int64)
    sf_ptrs_tensor = torch.tensor(sf_ptrs, dtype=torch.int64)
    sizes_tensor = torch.tensor(sizes, dtype=torch.int32)

    _launch_cache[cache_key] = (abc_ptrs_tensor, sf_ptrs_tensor, sizes_tensor,
                                (padded_tensors, sf_prepared))

    torch.ops.group_gemm_module.group_gemm_launch(abc_ptrs_tensor, sf_ptrs_tensor, sizes_tensor)
    return [t[2] for t in abc_tensors]

def custom_kernel_fallback(data):

    abc_tensors, sfasfb_tensors, _, problem_sizes = data

    results = []
    for i, ((a, b, c), (sfa, sfb), (M, N, K, L)) in enumerate(zip(abc_tensors, sfasfb_tensors, problem_sizes)):
        device = a.device

        a_fp16 = dequantize_fp4(a, M, K, L, device)
        b_fp16 = dequantize_fp4(b, N, K, L, device)

        a_scaled = apply_scale_factors(a_fp16, sfa, M, K, L, device)
        b_scaled = apply_scale_factors(b_fp16, sfb, N, K, L, device)

        a_batched = a_scaled.permute(2, 0, 1)
        b_batched = b_scaled.permute(2, 1, 0)

        c_batched = torch.bmm(a_batched.float(), b_batched.float())

        c_result = c_batched.permute(1, 2, 0).to(torch.float16)
        c.copy_(c_result)
        results.append(c)

    return results

USE_CUDA_KERNEL = os.environ.get('USE_CUDA_KERNEL', '1') == '1'

def custom_kernel(data):

    if USE_CUDA_KERNEL:
        return custom_kernel_cuda(data)
    else:
        return custom_kernel_fallback(data)
scrolls · 1425 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 399468.

⋯ diff truncated: revisions differ almost entirely

Best evidence level for this revision: reported

JSON