Skip to content
KernelIndex
Search⌘K

submission 505431

Joel🏴 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission-ptx.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-505431?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
25.2µs
#148 of 310
2026-02-21

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4440e4c11fbb40170747e614276ae403a019cd57c2c860286e3d47746b23d7ce
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)
fp4constexpr int BLOCK_K = 256; // Padded to 256 for 128B swizzle (256 FP4 / 2 = 128 bytes)
fused-epiloguealignas(8) uint64_t epilogue_mbar[2]; // MMA signals, epilogue waits
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; // Padded to 256 for 128B swizzle (256 FP4 / 2 = 128 bytes)
tile-m = 128constexpr int BLOCK_M = 128;
tile-n = 128constexpr int BLOCK_N = 128; // Per CTA; cluster covers 256
tmaasm volatile("cp.async.bulk.prefetch.L2.global.L2::cache_hint [%0], %1, %2;" ::"l"(src), "r"(size), "l"(cache_policy) : "memory");

Kernel source

submission-ptx.py2751 lines
#!POPCORN leaderboard nvfp4_group_gemm
#!POPCORN gpu B200

import os

import torch
from torch.utils.cpp_extension import load_inline

# =============================================================================
# CUDA Source: Utils (PTX helpers)
# =============================================================================

CUDA_SRC_UTILS = r"""
// utils.h - PTX utilities for nvfp4 group GEMM kernel (v1768978713)
// Note: No #pragma once since this file is concatenated into a single source

#include <cuda.h>
#include <cudaTypedefs.h>
#include <cuda_fp16.h>
#include <cstdint>
#include <cstdio>
#include <cstdlib>

#if 0
#define DEBUG_PRINT(...) printf(__VA_ARGS__)
#else
#define DEBUG_PRINT(...)
#endif

#if 0
#define DIAG_PRINT(...) printf(__VA_ARGS__)
#else
#define DIAG_PRINT(...)
#endif

// L2 Cache Policies
// https://github.com/NVIDIA/cutlass/blob/v4.3.2/include/cute/arch/copy_sm90_desc.hpp#L193-L197
// constexpr uint64_t EVICT_NORMAL = 0x1000000000000000;
constexpr uint64_t EVICT_FIRST = 0x12F0000000000000;
constexpr uint64_t EVICT_LAST = 0x14F0000000000000;

// Descriptor encoding for SMEM descriptors
__device__ inline constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; };

// =============================================================================
// Warp Election
// =============================================================================

// https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cute/arch/cluster_sm90.hpp#L180
__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;
}

// =============================================================================
// Mbarrier Operations
// =============================================================================

__device__ inline void mbarrier_init(int mbar_addr, int count)
{
  asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" ::"r"(mbar_addr), "r"(count));
}

// NOTE: using .shared::cluster
__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");
}

// https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cutlass/arch/barrier.h#L408
__device__ void mbarrier_wait(int mbar_addr, int phase)
{
  uint32_t ticks = 0x989680; // this is optional
  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));
}

// Simple arrive (no tx bytes)
__device__ inline void mbarrier_arrive(int mbar_addr)
{
  asm volatile("mbarrier.arrive.release.cta.shared::cta.b64 _, [%0];" ::"r"(mbar_addr) : "memory");
}

// Fence to ensure mbarrier init is visible across cluster
__device__ inline void fence_mbarrier_init()
{
  asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
}

// =============================================================================
// TMA Prefetch Operations
// =============================================================================

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

// =============================================================================
// TMA Load Operations (GMEM -> SMEM)
// =============================================================================

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

// =============================================================================
// tcgen05 Scale Factor Copy
// =============================================================================

template <int CTA_GROUP = 1>
__device__ inline void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc)
{
  // .32x128b corresponds to (32, 16) 8-bit scale -> 1 MMA for nvfp4.
  // .warpx4 duplicates data across 32-lane groups.
  asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" ::"r"(taddr), "l"(s_desc), "n"(CTA_GROUP));
}

// =============================================================================
// tcgen05 Commit (signal completion)
// =============================================================================

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

// =============================================================================
// tcgen05 Fence
// =============================================================================

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

// =============================================================================
// Collector Usage for A matrix reuse
// =============================================================================

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

// =============================================================================
// tcgen05 MMA Operations
// =============================================================================

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" // predicate register enable-input-d
      "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));
}

// =============================================================================
// tcgen05 Load (TMEM -> Registers)
// =============================================================================

// see https://docs.nvidia.com/cuda/inline-ptx-assembly/index.html
struct SHAPE
{
  static constexpr char _32x32b[] = ".32x32b";   // 32x1 tile for each warp
  static constexpr char _16x128b[] = ".16x128b"; // 16x4 tile
  static constexpr char _16x256b[] = ".16x256b"; // 16x8 tile
};

template <int NUM_REGS, const char *SHAPE, int NUM>
__device__ inline void tcgen05_ld(float *tmp, uint32_t tmem_addr, int row, int col)
{
  // Use addition as per reference to handle potential base address overlap
  // From docs/tcgen05-for-dummies.md
  // row << 16 puts row index in high 16 bits
  // tmem_addr is the base column index (alloc return value)
  // col is the column offset
  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)
{
  // each 32x32b tile uses 1 register per thread
  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)
{
  // each 16x128b tile uses 2 registers per thread
  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)
{
  // each 16x256b tile uses 4 registers per thread
  tcgen05_ld<num * 4, SHAPE::_16x256b, num>(tmp, tmem_addr, row, col);
}

// =============================================================================
// Utility Functions
// =============================================================================

template <typename T>
__device__ __inline__ T warp_uniform(T x) { return __shfl_sync(0xFFFF'FFFF, x, 0); }

// =============================================================================
// Cluster Utilities
// =============================================================================

// Get coordinate within cluster
__device__ inline int cluster_cta_rank()
{
  int rank;
  asm volatile("mov.u32 %0, %%cluster_ctarank;" : "=r"(rank));
  return rank;
}

// Get cluster dimension
__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;
}

// Get CTA coordinate within cluster
__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;
}

// Cluster barrier
__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();
}

// Arrive on mbarrier using cluster address space (for cross-CTA synchronization)
// Used to signal remote CTAs that local MMA has consumed a pipeline stage
__device__ inline void mbarrier_arrive_cluster(int mbar_addr)
{
  asm volatile("mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];" :: "r"(mbar_addr) : "memory");
}

// Map local shared memory address to remote CTA's equivalent in cluster address space
// Returns: cluster-scoped address usable with shared::cluster operations
__device__ inline int cluster_map_shared(int local_smem_addr, int remote_rank)
{
  int remote_addr;
  asm volatile("mapa.shared::cluster.u32 %0, %1, %2;"
               : "=r"(remote_addr)
               : "r"(local_smem_addr), "r"(remote_rank));
  return remote_addr;
}

// Named barrier for TMA-MMA synchronization (bar 1, 64 threads = 2 warps)
// TMA warp arrives after barrier init; MMA warp syncs before K-loop.
// Does NOT block epilogue warps (0-3).
__device__ inline void bar_arrive_tma_mma()
{
  asm volatile("bar.arrive 1, 64;" ::: "memory");
}

__device__ inline void bar_sync_tma_mma()
{
  asm volatile("bar.sync 1, 64;" ::: "memory");
}

// =============================================================================
// tcgen05 TMEM Allocation/Deallocation
// =============================================================================

// Allocate TMEM columns (writes TMEM address to shared memory holding buffer)
// num_cols must be a multiple of 32 (32 columns = 1 bank)
// cta_group::1 = independent per CTA, cta_group::2 = shared across 2 CTAs
// Returns the allocated TMEM address
template <int CTA_GROUP = 1>
__device__ inline void tcgen05_alloc(int smem_holding_buf_addr, int num_cols)
{
  // tcgen05.alloc writes the result to shared memory, not a register
  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" // <--- CRITICAL: Tells compiler memory was modified
    );
  }
  else
  {
    asm volatile(
        "tcgen05.alloc.cta_group::2.sync.aligned.shared::cta.b32 [%0], %1;" ::"r"(smem_holding_buf_addr), "r"(num_cols)
        : "memory" // <--- CRITICAL
    );
  }
  // CRITICAL FIX: Removed immediate ld.shared.
  // We must wait for tcgen05_wait_alloc() before reading.
}

// Deallocate TMEM
// tmem_addr: starting address of the allocation
// num_cols: number of columns to deallocate (must match allocation)
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));
  }
}

// Relinquish TMEM allocation permit (for epilogue)
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;");
  }
}

// Wait for TMEM allocation to complete
__device__ inline void tcgen05_wait_alloc()
{
  asm volatile("tcgen05.wait::ld.sync.aligned;");
}

constexpr int BLOCK_M = 128;
constexpr int BLOCK_N = 128; // Per CTA; cluster covers 256
constexpr int BLOCK_K = 256; // Padded to 256 for 128B swizzle (256 FP4 / 2 = 128 bytes)

constexpr int MMA_K = 64; // 32 bytes (2 units of 16 bytes)
constexpr int WARP_SIZE = 32;
constexpr int NUM_WARPS = 6;
constexpr int THREADS_PER_CTA = NUM_WARPS * WARP_SIZE; // 192

// Warp roles
// Warps 0-3: Epilogue (one per 32-row chunk of 128-row block)
constexpr int TMA_WARP = 4;
constexpr int MMA_WARP = 5;

// Pipeline stages
// NUM_STAGES_MAX is used for SMEM buffer allocation (arrays sized to this)
// With BLOCK_K=256, each stage is ~37KB. B200 has 228KB max -> 6 stages max (~221KB)
constexpr int NUM_STAGES_MAX = 6;
constexpr int NUM_STAGES_DEFAULT = 5;

// Cluster configuration
// Dynamic cluster selection: (1, 1, 1) for small N, (1, 2, 1) for large N
// CLUSTER_N=2 splits N across 2 CTAs, using TMA multicast for Matrix A and SFA
constexpr int CLUSTER_M = 1;
constexpr int CLUSTER_N = 1;
constexpr int CLUSTER_Z = 1;

// Grid grouping for L2 persistence
constexpr int GROUP_M = 8;
constexpr int GROUP_N = 8;

// =============================================================================
// Scale Factor Configuration
// =============================================================================

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

// Max groups we support
constexpr int MAX_GROUPS = 16;

// TensorMap constants (per group: A, B, C, SFA, SFB = 5 descriptors)
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;

// Fused mode TensorMap layout:
// tmaps[0] = fused A, tmaps[1] = fused SFA
// tmaps[FUSED_TMAP_B_BASE + g] = group g's B
// tmaps[FUSED_TMAP_SFB_BASE + g] = group g's SFB
constexpr int FUSED_TMAP_A = 0;
constexpr int FUSED_TMAP_SFA = 1;
constexpr int FUSED_TMAP_B_BASE = 2;
constexpr int FUSED_TMAP_SFB_BASE = 2 + MAX_GROUPS; // 18

// Max M-tiles for tile_group array
constexpr int MAX_M_TILES = 256;

// Unified kernel parameters passed as __grid_constant__
// Contains all group params + TensorMaps in a single struct
// Size: 16*72 + 80*128 + 8 + 256 + 8 ≈ 11,700 bytes (within 32KB limit)
struct alignas(128) KernelParams {
    GroupParams groups[MAX_GROUPS];
    CUtensorMap tmaps[MAX_GROUPS * TENSORMAPS_PER_GROUP];
    uint32_t total_tiles;
    uint32_t num_groups;

    // M-fusion fields (active when fused != 0)
    uint8_t tile_group[MAX_M_TILES]; // maps M-tile index -> original group_id
    uint16_t fused_tiles_m;          // total fused M-tiles
    uint16_t fused_tiles_n;          // N-tiles (per cluster column)
    uint32_t fused;                  // 1 = fused mode, 0 = legacy

    // Low-overhead prefetch support
    uint8_t next_tile_group[MAX_M_TILES]; // group_id for (tile_m + 1) % fused_tiles_m
    uint16_t next_tile_n[MAX_M_TILES];    // tile_n increment for tile_m + 1 (0 or 1)
};

// Scale Factor Layout for tcgen05 (matching gau.nernst reference):
// SF data is stored contiguously as [M_rows, K/16] FP8 values
// For BLOCK_M=128, BLOCK_K=256: 128 rows × 16 scale factors = 2048 bytes
constexpr int SF_VEC_SIZE = 16; // K-values per scale factor

// MMA K-loop constants
constexpr int NUM_MMA_K_ITERS = BLOCK_K / MMA_K; // 4 for BLOCK_K=256, MMA_K=64

// SF copy constants for blocked format (cuBLAS layout):
//
// KEY INSIGHT from PTX ISA Scale Factor diagrams (Fig 233, 242):
// - For scale_vec::4X/block16, MMA reads 4 TMEM columns at once
// - Each column holds 32 TMEM lanes (rows 0-31 only!)
// - M0-M31 → lanes 0-31, column X
// - M32-M63 → lanes 0-31, column X+1 (same lanes, different column!)
// - M64-M95 → lanes 0-31, column X+2
// - M96-M127 → lanes 0-31, column X+3
//
// tcgen05.cp.32x128b.warpx4 with SBO=128:
// - Each warp reads 128 bytes from base + warp_id * 128
// - Warp 0: bytes 0-127 (M0-M31), Warp 1: bytes 128-255 (M32-M63), etc.
// - All 128 TMEM lanes (0-127) get distinct SF data for M0-M127
// - Matches cuBLAS blocked format produced by to_blocked()/prepare_sf_for_tma()
//
// For BLOCK_K=256 (4 k_mma iterations):
// - k_mma=0: columns X to X+3 (SF for K=0-63)
// - k_mma=1: columns X+4 to X+7 (SF for K=64-127)
// - k_mma=2: columns X+8 to X+11 (SF for K=128-191)
// - k_mma=3: columns X+12 to X+15 (SF for K=192-255)
// Total: 16 TMEM columns per SF tensor
//
constexpr int SF_SBO = 128;          // 8*16 = 128 bytes stride between warp groups (matching gau.nernst)
constexpr int SF_SMEM_ADVANCE = 512; // 512 bytes per k_mma (32 rows × 16 bytes for warp 0)
constexpr int SF_TMEM_ADVANCE = 4;   // Advance 4 TMEM columns per k_mma iteration

// Descriptor advancement for A/B (in bits [0-13] units i.e. / 16 bytes)
// MMA_K=64 elements = 32 bytes.
// 32 bytes / 16 = 2 units.
constexpr int A_K_STRIDE = 2;
constexpr int B_K_STRIDE = 2;

// =============================================================================
// TMEM Configuration
// =============================================================================

// TMEM columns needed for accumulator + scale factors
// For scale_vec::4X/block16, each MMA reads 4 columns at once:
//   - Accumulator: BLOCK_N columns (128)
//   - SFA: 4 columns per k_mma × 4 k_mma iterations = 16 columns
//   - SFB: 4 columns per k_mma × 4 k_mma iterations = 16 columns
// Layout: [Acc: 0-127][SFA: 128-143][SFB: 144-159] (but allocate 32-col minimum)
constexpr int TMEM_ACC_COLS = 128; // Accumulator (128 columns for 128x128 tile)
// Each k_mma iteration needs 4 TMEM columns for SF (scale_vec::4X)
constexpr int TMEM_SFA_COLS_LOGICAL = NUM_MMA_K_ITERS * 4;                     // 4 iters × 4 cols = 16
constexpr int TMEM_SFB_COLS_LOGICAL = NUM_MMA_K_ITERS * 4;                     // 4 iters × 4 cols = 16
constexpr int TMEM_SFA_COLS = 32;                                              // Minimum allocation = 1 bank = 32 columns
constexpr int TMEM_SFB_COLS = 32;                                              // Minimum allocation = 1 bank = 32 columns
// Single-buffered accumulator: cta_group::1 limits TMEM to 256 columns per CTA.
// Double-buffered (2×128+32+32=320) exceeds this limit, causing tcgen05_alloc to stall.
// TMA-epilogue overlap is still achieved via double-buffered K-loop barriers (full/empty_mbar[2][*]).
constexpr int NUM_ACC_BUFS = 2;
constexpr int TMEM_TOTAL_COLS = NUM_ACC_BUFS * TMEM_ACC_COLS + TMEM_SFA_COLS + TMEM_SFB_COLS; // 2*128 + 32 + 32 = 320
constexpr int TMEM_ALLOC_COLS = (TMEM_TOTAL_COLS <= 256) ? 256 : 512; // 512 columns (power of 2)

// =============================================================================
// Shared Memory Layout
// =============================================================================

// SMEM sizes per stage (bytes)
// A: BLOCK_M x BLOCK_K/2 (FP4 packed) = 128 x 128 = 16KB
// B: BLOCK_N x BLOCK_K/2 (FP4 packed) = 128 x 128 = 16KB
// SFA: BLOCK_M x BLOCK_K/16 (FP8) = 128 x 16 = 2KB (but padded to 16B/row)
// SFB: BLOCK_N x BLOCK_K/16 (FP8) = 128 x 16 = 2KB (but padded to 16B/row)
// Total per stage: ~40KB

constexpr int SMEM_A_SIZE = BLOCK_M * (BLOCK_K / 2); // 16384 bytes
constexpr int SMEM_B_SIZE = BLOCK_N * (BLOCK_K / 2); // 16384 bytes

// Scale factor SMEM size (matching gau.nernst reference):
// Contiguous layout: [128 rows × BLOCK_K/16 scale factors] = 128 * 16 = 2048 bytes
// For BLOCK_K=256: 128 * (256/16) = 128 * 16 = 2048 bytes
constexpr int SMEM_SFA_SIZE = BLOCK_M * (BLOCK_K / 16); // 2048 bytes
constexpr int SMEM_SFB_SIZE = BLOCK_N * (BLOCK_K / 16); // 2048 bytes
constexpr int SMEM_STAGE_SIZE = SMEM_A_SIZE + SMEM_B_SIZE + SMEM_SFA_SIZE + SMEM_SFB_SIZE;

// TMA transaction sizes (for mbarrier expect_tx)
constexpr int TMA_A_BYTES = SMEM_A_SIZE;                                                 // 16384 bytes
constexpr int TMA_B_BYTES = SMEM_B_SIZE;                                                 // 16384 bytes
constexpr int TMA_SFA_BYTES = SMEM_SFA_SIZE;                                             // 2048 bytes
constexpr int TMA_SFB_BYTES = SMEM_SFB_SIZE;                                             // 2048 bytes
constexpr int TMA_BYTES_AB = TMA_A_BYTES + TMA_B_BYTES;                                  // 32768 bytes
constexpr int TMA_BYTES_ALL = TMA_A_BYTES + TMA_B_BYTES + TMA_SFA_BYTES + TMA_SFB_BYTES; // 36864 bytes

template <int STAGES>
struct SmemBuffersT
{
  // [CRITICAL] Data buffers MUST come FIRST to ensure 128-byte alignment
  // Putting them first means they inherit the struct's base alignment
  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[STAGES];
  AlignedBuffB B_smem[STAGES];
  AlignedBuffSFA SFA_smem[STAGES];
  AlignedBuffSFB SFB_smem[STAGES];

  // Mbarriers and metadata come AFTER data buffers (8-byte alignment is sufficient)
  // Double-buffered barriers: [bp] where bp = tile_counter & 1
  // Even/odd tiles use separate barrier sets — no re-init conflicts during overlap
  alignas(8) uint64_t full_mbar[2][STAGES];  // TMA signals, MMA waits
  alignas(8) uint64_t empty_mbar[2][STAGES]; // MMA signals, TMA waits
  alignas(8) uint64_t epilogue_mbar[2];               // MMA signals, epilogue waits
  alignas(8) uint64_t epilogue_done_mbar[2];          // Epilogue signals, MMA waits (TMEM free)
  alignas(8) uint64_t tmem_holding_buf;               // Used by tcgen05_alloc

  // Flags for epilogue barrier readiness (set by MMA after reinit, checked by epilogue)
  volatile int epi_barriers_ready[2];
};

// Backward-compatible alias for maximum stage count
using SmemBuffers = SmemBuffersT<NUM_STAGES_MAX>;

// =============================================================================
// SMEM Descriptor Building
// =============================================================================

// Build 64-bit SMEM descriptor for tcgen05 MMA operand A
//
// Descriptor format (64-bit):
//   Bits [0-13]:  Base Address >> 4
//   Bits [16-29]: LBO = 0 (hardware-implied for K-major with swizzle)
//   Bits [32-45]: SBO >> 4 = 1024 >> 4 = 64
//   Bit 46:       1 = K-major (K contiguous)
//   Bits [61-63]: Swizzle mode 2 = 128B swizzle
//
// Matches gau.nernst reference (line 515):
//   constexpr uint64_t AB_desc = (desc_encode(8 * 128) << 32) | (1 << 46) | (2 << 61);
//
// PTX ISA 9.1 canonical layout for K-major, 128B swizzle:
//   ((8, m), (T, 2k)) : ((8T, SBO), (1, T))
// Inner M dimension is 8 rows. SBO = 8 rows × bytes_per_row.
// For BLOCK_K=256 FP4 = 128 bytes/row: SBO = 8 * 128 = 1024.
//
// LBO field = 0: PTX ISA states "LBO encoding = 1 (assumed)" for K-major
// with swizzle. Hardware uses implicit LBO; descriptor field must be 0.
__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; // 1024 = 8 rows × 128 bytes/row

  uint64_t desc = desc_encode(addr)
                  // Bits 16-29: LBO = 0 (hardware-implied for K-major + swizzle)
                  | (desc_encode(SBO) << 32ULL) // Bits 32-45: SBO = 64
                  | (1ULL << 46ULL)             // Bit 46: K-major
                  | (2ULL << 61ULL);            // Bits 61-63: 128B swizzle
  return desc;
}

// Build 64-bit SMEM descriptor for tcgen05 MMA operand B
// Same format as A - K-major with 128B swizzle (see make_smem_desc_A)
__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; // 1024 = 8 rows × 128 bytes/row

  uint64_t desc = desc_encode(addr)
                  // Bits 16-29: LBO = 0 (hardware-implied for K-major + swizzle)
                  | (desc_encode(SBO) << 32ULL) // Bits 32-45: SBO = 64
                  | (1ULL << 46ULL)             // Bit 46: K-major
                  | (2ULL << 61ULL);            // Bits 61-63: 128B swizzle
  return desc;
}

// Legacy function for backward compatibility
__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); // Default to N-major
}

// Build SMEM descriptor for scale factors (tcgen05.cp.32x128b.warpx4)
//
// With blocked format and SBO=128:
//   - Each warp reads from base + warp_id * 128 bytes
//   - Warp 0: M0-M31, Warp 1: M32-M63, Warp 2: M64-M95, Warp 3: M96-M127
//   - All 128 TMEM lanes get distinct SF data for M0-M127
//   - Matches cuBLAS blocked format from to_blocked()/prepare_sf_for_tma()
//
// Descriptor format (64-bit):
//   Bits [0-13]:  Base Address >> 4
//   Bits [32-45]: SBO (Stride Byte Offset) >> 4 = 128 >> 4 = 8
//   Bit 46:       Mode bit = 1 (no swizzle)
__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;                                // 128 bytes stride between warp groups
  uint64_t desc = desc_encode(addr) | (desc_encode(SBO) << 32ULL) // SBO at bits 32-45
                  | (1ULL << 46ULL);                              // Mode = 1 (no swizzle)
  return desc;
}

// Build instruction descriptor for tcgen05 MMA
// See: https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-instruction-descriptor
// Reference: PTX ISA 9.1 Table 44 - Instruction descriptor format for .kind::mxf4nvf4
// Matches working implementations (hekailove, gau.nernst)
__device__ inline uint32_t make_mma_idesc(int mma_n = BLOCK_N)
{
  // For cta_group::1 (independent CTAs), dimensions are BLOCK_M x BLOCK_N
  constexpr uint32_t MMA_M = BLOCK_M; // 128
  const uint32_t MMA_N = mma_n;

  // NVFP4 instruction descriptor encoding for .kind::mxf4nvf4:
  // Table 44 from PTX ISA 9.1:
  // - Bits 7-9:   atype (E2M1 = 1)
  // - Bits 10-11: btype (E2M1 = 1)
  // - Bit 12:     Reserved (0)
  // - Bit 13:     Negate A Matrix (0 = no negate)
  // - Bit 14:     Negate B Matrix (0 = no negate)
  //   NOTE: B.T is achieved via data layout (B stored as [N,K]), NOT via negate bit
  //   Reference (gau.nernst) does NOT set bit 14.
  // - Bits 17-22: N >> 3 (output tile N dimension / 8)
  // - Bit 23:     stype (UE4M3 = 0 for mxf4nvf4)
  // - Bits 27-28: M >> 7 (output tile M dimension / 128)
  uint32_t idesc = (1U << 7U)                // atype = E2M1
                             | (1U << 10U)             // btype = E2M1
                             | (0U << 14U)             // No negate B (matches gau.nernst)
                             | ((MMA_N >> 3U) << 17U)  // N / 8 = 128/8 = 16 or 64/8 = 8
                             | (0U << 23U)             // stype = UE4M3
                             | ((MMA_M >> 7U) << 27U); // M / 128 = 128/128 = 1
  return idesc;
}

// =============================================================================
// Descriptor Validation Helpers
// =============================================================================

// Validate SMEM descriptor fields (for debugging)
// Returns: true if all fields match expected values
__device__ inline void validate_smem_desc(uint64_t desc, const char *name, int expected_sbo, int expected_mode, int expected_swizzle)
{
  uint64_t sbo = ((desc >> 32) & 0x3FFFULL) << 4; // Bits 32-45 (shifted back)
  uint64_t mode_bit = (desc >> 46) & 0x1ULL;      // Bit 46
  uint64_t swizzle = (desc >> 61) & 0x7ULL;       // Bits 61-63

  DIAG_PRINT("[%s] SBO=%llu (expected %d) %s\n", name,
             (unsigned long long)sbo, expected_sbo,
             sbo == (uint64_t)expected_sbo ? "OK" : "MISMATCH");
  DIAG_PRINT("[%s] Mode=%llu (expected %d) %s\n", name,
             (unsigned long long)mode_bit, expected_mode,
             mode_bit == (uint64_t)expected_mode ? "OK" : "MISMATCH");
  DIAG_PRINT("[%s] Swizzle=%llu (expected %d) %s\n", name,
             (unsigned long long)swizzle, expected_swizzle,
             swizzle == (uint64_t)expected_swizzle ? "OK" : "MISMATCH");
}

// Print SF descriptor debug info:
// - K-offset stride (SBO field)
// - Scale factor count (derived from MMA_K and block size)
__device__ inline void print_sf_descriptor_info(uint64_t sf_desc, int mma_k, int block_size)
{
  uint64_t sf_sbo = ((sf_desc >> 32) & 0x3FFFULL) << 4;
  int sf_per_mma = mma_k / block_size;

  DIAG_PRINT("[SF_INFO] K-offset stride (SBO): %llu bytes\n", (unsigned long long)sf_sbo);
  DIAG_PRINT("[SF_INFO] Scale factors per MMA: %d (MMA_K=%d / block_size=%d)\n",
             sf_per_mma, mma_k, block_size);
}

"""

# =============================================================================
# CUDA Source: Kernel
# =============================================================================

CUDA_SRC_KERNEL = r"""
// kernel.cu - Group GEMM kernel with persistent grid
// Phase 7: Single-buffered TMEM acc (256-col cta_group::1 limit), double-buffered K-loop barriers
// TMA-epilogue overlap: TMA loads tile N+1 while epilogue drains tile N

// =============================================================================
// Kernel Configuration
// =============================================================================

// utils.h included above

// Debug: only CTA 0, lane 0
#define DBG(fmt, ...) do { } while(0)

// =============================================================================
// Templated Kernel for Dynamic Cluster Selection
// =============================================================================
// CLUSTER_N_PARAM: 1 for small N (no multicast), 2 for large N (multicast A)
// NUM_STAGES_PARAM: Pipeline depth (5 for memory-bound, 6 for math-bound)
// K_PARAM: Template specialization for K (0 means use runtime K)
// BLOCK_N_PARAM: Tile width (usually 128 or 64)
// Separate instantiations allow compile-time optimization of cluster-specific code

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)
{
    // Use template parameter for cluster-dependent code
    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;

    // =========================================================================
    // Thread/Block/Cluster Identification
    // =========================================================================

    const int tid = threadIdx.x;
    const int warp_id = tid / WARP_SIZE;
    const int lane_id = tid % WARP_SIZE;

    // Cluster position
    // With cluster (1, 2, 1), cta_y is 0 or 1 - determines which N-portion this CTA handles
    const int cta_n = cluster_cta_y(); // 0 or 1 for cluster (1, 2, 1)

    // =========================================================================
    // Persistent grid: cluster loops over tiles
    // =========================================================================
    const uint32_t total_clusters = gridDim.y / CLUSTER_N;
    const uint32_t cluster_id = blockIdx.y / CLUSTER_N;

    // =========================================================================
    // Shared Memory Setup
    // =========================================================================

    extern __shared__ char smem_raw[];

    // [CRITICAL FIX] Manually align smem_raw to 128 bytes
    // Dynamic shared memory is only 8-byte aligned by default
    uintptr_t smem_addr = reinterpret_cast<uintptr_t>(smem_raw);
    uintptr_t aligned_addr = (smem_addr + 127) & ~uintptr_t(127);
    SmemBuffersT<NUM_STAGES> *smem = reinterpret_cast<SmemBuffersT<NUM_STAGES> *>(aligned_addr);

    // Get SMEM addresses for barriers
    auto get_mbar_addr = [](void *mbar) -> int
    {
        return static_cast<int>(__cvta_generic_to_shared(mbar));
    };

    // Multicast mask: all CTAs in cluster
    constexpr int16_t MCAST_MASK_ALL = (CLUSTER_N == 4) ? 0xF : (CLUSTER_N == 2) ? 0x3
                                                                                  : 0x1;

    // [UNIFIED FLOW FIX] use_multicast must be consistent across all CTAs in cluster
    constexpr bool use_multicast = (CLUSTER_N > 1);

    // =====================================================================
    // TMEM Allocation — ONCE before the tile loop
    // =====================================================================
    // TMEM Allocation (Single 512-column block)
    // =====================================================================

    uint32_t acc_tmem[NUM_ACC_BUFS] = {0, 128};
    uint32_t sfa_tmem = 256;
    uint32_t sfb_tmem = 288;
    uint32_t idesc = 0;

    if (warp_id == MMA_WARP)
    {
        // Allocate single 512-column TMEM block.
        // The CTA fully owns its cta_group::1 partition, so TMEM implicitly starts at offset 0.
        // tcgen05.alloc writes the base address to shared memory — we pass a valid smem addr but ignore the result.
        int alloc_dst = static_cast<int>(__cvta_generic_to_shared(smem));
        tcgen05_alloc<1>(alloc_dst, TMEM_ALLOC_COLS);
        tcgen05_wait_alloc();
    }
    
    if (warp_id == MMA_WARP)
    {
        idesc = make_mma_idesc(BLOCK_N);
    }

    // Sync CTA to ensure TMEM allocation is complete before any warp accesses it
    __syncthreads();

    // =========================================================================
    // PROLOGUE: Initialize first tile's barriers
    // =========================================================================

    if (tid == 0)
    {
        // Init K-loop barriers for buffer parity 0 (first tile)
        // empty_mbar init count = CLUSTER_N for multicast: TMA waits for ALL CTAs' MMA
        // to consume a stage before reusing it for multicast (prevents cross-CTA phase races)
        for (int s = 0; s < NUM_STAGES; s++)
        {
            mbarrier_init(get_mbar_addr(&smem->full_mbar[0][s]), 1);
            mbarrier_init(get_mbar_addr(&smem->empty_mbar[0][s]), use_multicast ? CLUSTER_N : 1);
        }
        // Init epilogue barriers for both bp indices (tiles 0..1 skip MMA reinit)
        for (int a = 0; a < 2; a++)
        {
            mbarrier_init(get_mbar_addr(&smem->epilogue_mbar[a]), 1); // 1 MMA warp
            mbarrier_init(get_mbar_addr(&smem->epilogue_done_mbar[a]), 4); // 4 epilogue warps
            smem->epi_barriers_ready[a] = 1;
        }
        fence_mbarrier_init();
        asm volatile("fence.proxy.async.shared::cta;" ::: "memory");
    }

    // Full CTA sync — one-time prologue sync (all 192 threads)
    __syncthreads();

    // Batch expect_tx for first tile's first epoch (multicast path)
    // Arms full_mbar for all stages so TMA can multicast without per-k_iter cluster_sync.
    // Done BEFORE cluster_sync so one sync covers both barrier init + expect_tx.
    if constexpr (use_multicast)
    {
        if (warp_id == TMA_WARP && elect_sync())
        {
            for (int s = 0; s < NUM_STAGES; s++)
            {
                mbarrier_arrive_expect_tx(get_mbar_addr(&smem->full_mbar[0][s]), TMA_BYTES_ALL);
            }
        }
    }

    // Cluster sync to ensure all CTAs have barriers ready + expect_tx armed
    // CRITICAL: barrier.cluster requires ALL threads (.aligned), not just one warp
    if constexpr (use_multicast)
    {
        DBG("PRO: pre csync\n");
        cluster_sync();
        DBG("PRO: post csync\n");
    }

    // =========================================================================
    // Persistent Tile Loop
    // =========================================================================

    int tile_counter = 0;

    // Block distribution: each cluster gets consecutive tiles for B L2 reuse.
    // With M-fast ordering, consecutive tiles share the same tile_n (same B data).
    // Cyclic (stride) distribution scatters tiles across N, thrashing L2.
    const uint32_t tiles_per_cluster = (kparams.total_tiles + total_clusters - 1) / total_clusters;
    const uint32_t work_start = cluster_id * tiles_per_cluster;
    const uint32_t work_end = min(work_start + tiles_per_cluster, kparams.total_tiles);

    for (uint32_t work_idx = work_start; work_idx < work_end; work_idx++)
    {
        int bp = tile_counter & 1;       // barrier parity for K-loop SMEM barriers (full/empty_mbar)
        // Acc/epilogue always use index 0 (single-buffered TMEM, cta_group::1 = 256 cols max)
        // if (tid == 0 && blockIdx.y == 0) printf("=== tile=%d bp=%d ===\n", tile_counter, bp);

        // -----------------------------------------------------------------
        // Work Decoding - Find group and tile coordinates
        // -----------------------------------------------------------------

        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 epi_coord_m_offset = 0; // offset to subtract from coord_m for C addressing (fused mode)

        if (kparams.fused)
        {
            // Fused mode: single rectangular tile grid, O(1) group lookup
            int fused_tiles_m = kparams.fused_tiles_m;
            tile_n = work_idx / fused_tiles_m;
            tile_m = work_idx % fused_tiles_m;
            group_id = kparams.tile_group[tile_m];
        }
        else
        {
            // Legacy mode: binary search to find which group this tile belongs to
            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;
        }

        // Load group parameters
        params = kparams.groups[group_id];

        if (!kparams.fused)
        {
            // Legacy: decode tile indices within group
            uint32_t local_idx = work_idx - params.tile_offset;
            // M-major ordering: M varies fast so consecutive tiles share same B in L2
            tile_n = local_idx / params.tiles_m;
            tile_m = local_idx % params.tiles_m;
        }
        else
        {
            // Fused: tile_offset stores cumulative M offset for epilogue C addressing
            epi_coord_m_offset = params.tile_offset * BLOCK_M;
        }

        // K-iteration determination (compile-time if specialized)
        if constexpr (K_EXPECTED > 0)
        {
            num_k_iters = K_EXPECTED / BLOCK_K;
        }
        else
        {
            num_k_iters = params.K / BLOCK_K;
        }

        // N-tile Bounds Check (Cluster Safety)
        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);

        // -----------------------------------------------------------------
        // Setup pointers and coordinates for this tile
        // -----------------------------------------------------------------

        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;

        if (has_valid_work)
        {
            if (kparams.fused)
            {
                // Fused mode: single A/SFA tmap, per-group B/SFB tmaps
                tmap_A = &kparams.tmaps[FUSED_TMAP_A];
                tmap_B = &kparams.tmaps[FUSED_TMAP_B_BASE + group_id];
                tmap_SFA = &kparams.tmaps[FUSED_TMAP_SFA];
                tmap_SFB = &kparams.tmaps[FUSED_TMAP_SFB_BASE + group_id];

                // Prefetch the 4 tmaps used by this tile
                if (tid == 0) prefetch_tensormap(&kparams.tmaps[FUSED_TMAP_A]);
                if (tid == 1) prefetch_tensormap(&kparams.tmaps[FUSED_TMAP_B_BASE + group_id]);
                if (tid == 2) prefetch_tensormap(&kparams.tmaps[FUSED_TMAP_SFA]);
                if (tid == 3) prefetch_tensormap(&kparams.tmaps[FUSED_TMAP_SFB_BASE + group_id]);
            }
            else
            {
                // Legacy mode: per-group tmaps
                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]);
                }
            }

            // Calculate tile coordinates
            coord_m = tile_m * BLOCK_M;
            coord_n = my_n_tile * BLOCK_N;
        }

        // -----------------------------------------------------------------
        // MMA warp: wait for acc buffer to be free, reinit epilogue barriers
        // Single-buffered acc: MMA must wait for previous tile's epilogue every time.
        // -----------------------------------------------------------------

        if (warp_id == MMA_WARP)
        {
            if (tile_counter >= 1)
            {
                DBG("MMA: wait epi_done\n");
                // Serialized: always wait for previous tile's epilogue (bp^1).
                // Concurrent double-buffered would guard >= NUM_ACC_BUFS and wait_bp = bp.
                mbarrier_wait(get_mbar_addr(&smem->epilogue_done_mbar[bp ^ 1]), 0);
                DBG("MMA: epi_done passed, reinit\n");

                // Reinit epilogue barriers for this tile's bp (MMA owns them now).
                // Phase resets to 0 — epilogue always waits with phase 0 after reinit.
                if (lane_id == 0)
                {
                    mbarrier_init(get_mbar_addr(&smem->epilogue_mbar[bp]), 1);
                    mbarrier_init(get_mbar_addr(&smem->epilogue_done_mbar[bp]), 4);
                    fence_mbarrier_init();
                    asm volatile("fence.proxy.async.shared::cta;" ::: "memory");
                    smem->epi_barriers_ready[bp] = 1;
                }

                // CRITICAL: Canonical PTX return handoff fence (epilogue→MMA)
                tcgen05_fence_after_thread_sync();

                DBG("MMA: reinit done\n");
            }

            // Wait for TMA to finish K-loop barrier init for this tile's bp
            if (tile_counter > 0)
            {
                DBG("MMA: wait bar_sync\n");
                bar_sync_tma_mma();
                DBG("MMA: bar_sync ok\n");
            }
        }

        // =====================================================================
        // Main Pipelined K-Loop
        // =====================================================================

        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;

            // -----------------------------------------------------------------
            // STEP 1: TMA warp issues loads for A, B, SFA, SFB (all async)
            // -----------------------------------------------------------------

            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[bp][stage]);

            // STEP 1a: TMA warp sets up barrier expectations
            // Multicast path: expect_tx is batched in prologue/post-K-loop + rolling in MMA.
            // Non-multicast path: per-k_iter expect_tx (no cross-CTA coordination needed).
            if constexpr (!use_multicast)
            {
                if (has_valid_work && warp_id == TMA_WARP)
                {
                    if (elect_sync())
                    {
                        if (k_iter == 0) DBG("TMA: expect_tx[%d][0]\n", bp);
                        mbarrier_arrive_expect_tx(full_mbar_addr, TMA_BYTES_ALL);
                    }
                }
            }

            // STEP 1b: Per-k_iter cluster_sync REMOVED.
            // Cross-CTA synchronization now handled by:
            // - Batched expect_tx (prologue/post-K-loop) + rolling expect_tx (MMA path)
            // - Cross-CTA empty_mbar arrives (MMA) with init_count=CLUSTER_N
            // This restores pipeline depth from 1 to NUM_STAGES for multicast.

            // STEP 1c: TMA warp issues actual TMA loads
            if (has_valid_work && warp_id == TMA_WARP)
            {
                if (elect_sync())
                {
                    // Wait for stage buffer to be consumed BEFORE issuing new TMA
                    if (k_iter >= NUM_STAGES)
                    {
                        int empty_mbar_addr = get_mbar_addr(&smem->empty_mbar[bp][stage]);
                        int phase = ((k_iter / NUM_STAGES) + 1) & 1;
                        mbarrier_wait(empty_mbar_addr, phase);
                    }

                    // Load B and SFB (Unicast, issued by all CTAs)
                    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);

                    // Load A and SFA (Multicast gating: only rank 0 issues)
                    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
                    {
                        // Unicast fallback
                        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 (k_iter == 0) DBG("TMA: loads issued\n");
                }
            }

            // -----------------------------------------------------------------
            // STEP 2: MMA warp waits for TMA, then SF copy and MMA
            // -----------------------------------------------------------------

            if (has_valid_work && warp_id == MMA_WARP)
            {
                uint32_t cur_acc = acc_tmem[bp % NUM_ACC_BUFS]; // Use double-buffered index
                int full_mbar_addr_mma = get_mbar_addr(&smem->full_mbar[bp][stage]);
                int empty_mbar_addr = get_mbar_addr(&smem->empty_mbar[bp][stage]);

                int phase = (k_iter / NUM_STAGES) & 1;

                // MMA warp waits for TMA data
                if (k_iter == 0) DBG("MMA: wait full[%d][0] ph=%d\n", bp, phase);
                mbarrier_wait(full_mbar_addr_mma, phase);
                if (k_iter == 0) DBG("MMA: full ok\n");

                // =========================================================
                // MMA K-LOOP (Unrolled first iteration)
                // =========================================================

                uint32_t sfa_tmem_base = sfa_tmem;
                uint32_t sfb_tmem_base = sfb_tmem;

                // Base A/B descriptors
                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);

                // --- ITERATION 0 ---
                {
                    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));

                    // Next iteration descriptors (for interleaving)
                    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);

                    // Initial SF copy for k=0
                    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();

                    // Issue SF copy for k=1 WHILE k=0 MMA is running
                    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>(cur_acc, a_desc, b_desc, idesc,
                                             sfa_tmem_base, sfb_tmem_base, enable_input_d);
                    }

                    a_desc += A_K_STRIDE;
                    b_desc += B_K_STRIDE;
                }

// --- ITERATIONS 1..3 ---
#pragma unroll
                for (int k = 1; k < NUM_MMA_K_ITERS; k++)
                {
                    tcgen05_fence_before_thread_sync();

                    // Issue SF copy for k+1 WHILE k MMA is running
                    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>(cur_acc, 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);
                }

                // Cross-CTA multicast synchronization (replaces per-k_iter cluster_sync)
                // Order is critical: expect_tx MUST precede remote arrive so that
                // remote CTA's full_mbar is armed before the arrive triggers TMA.
                if constexpr (use_multicast)
                {
                    // Rolling expect_tx: arm full_mbar for this stage's next use
                    if (k_iter + NUM_STAGES < num_k_iters)
                    {
                        if (elect_sync())
                        {
                            mbarrier_arrive_expect_tx(full_mbar_addr_mma, TMA_BYTES_ALL);
                        }
                    }

                    // Arrive on remote CTAs' empty_mbar to signal SMEM is free
                    // for multicast reuse. Each CTA's empty_mbar has init_count=CLUSTER_N,
                    // so TMA waits for ALL CTAs to consume before reusing a stage.
                    if (elect_sync())
                    {
                        int local_empty_addr = get_mbar_addr(&smem->empty_mbar[bp][stage]);
                        #pragma unroll
                        for (int r = 0; r < CLUSTER_N; r++)
                        {
                            if (r != cta_n)
                            {
                                int remote_addr = cluster_map_shared(local_empty_addr, r);
                                mbarrier_arrive_cluster(remote_addr);
                            }
                        }
                    }
                }
            }
        }

        // =====================================================================
        // POST K-LOOP: Warp-specialized, NO __syncthreads
        // =====================================================================

        // Check if there is a next tile
        bool has_next_tile = (work_idx + 1 < work_end);
        int next_bp = bp ^ 1;

        // -----------------------------------------------------------------
        // MMA WARP: Signal epilogue that accumulator is ready
        // -----------------------------------------------------------------
        if (warp_id == MMA_WARP)
        {
            if (has_valid_work)
            {
                tcgen05_fence_before_thread_sync();

                if (lane_id == 0)
                {
                    DBG("MMA: done, signal epi\n");
                    mbarrier_arrive(get_mbar_addr(&smem->epilogue_mbar[bp]));
                }
            }
        }

        // -----------------------------------------------------------------
        // TMA WARP: Init next tile's K-loop barriers
        // -----------------------------------------------------------------
        if (warp_id == TMA_WARP && has_next_tile)
        {
            DBG("TMA: reinit[%d]\n", next_bp);
            if (lane_id == 0)
            {
                for (int s = 0; s < NUM_STAGES; s++)
                {
                    mbarrier_init(get_mbar_addr(&smem->full_mbar[next_bp][s]), 1);
                    mbarrier_init(get_mbar_addr(&smem->empty_mbar[next_bp][s]), use_multicast ? CLUSTER_N : 1);
                }
                fence_mbarrier_init();
                asm volatile("fence.proxy.async.shared::cta;" ::: "memory");
            }

            // Batch expect_tx for next tile's first epoch (multicast path)
            // Done here (before cluster_sync) so one sync covers reinit + expect_tx.
            if constexpr (use_multicast)
            {
                if (elect_sync())
                {
                    for (int s = 0; s < NUM_STAGES; s++)
                    {
                        mbarrier_arrive_expect_tx(get_mbar_addr(&smem->full_mbar[next_bp][s]), TMA_BYTES_ALL);
                    }
                }
            }

            // L2 Prefetch for next tile's A+B+SFA+SFB (fire-and-forget, no sync needed)
            // Fires while epilogue is running — by the time next K-loop starts, data is warm in L2.
            if (elect_sync() && kparams.fused && work_idx + 1 < work_end)
            {
                int next_group_id = kparams.next_tile_group[tile_m]; // precomputed by wrapper
                int next_n_increment = kparams.next_tile_n[tile_m];  // 0=same N, 1=next N
                int next_my_n_tile = (tile_n + next_n_increment) * CLUSTER_N + cta_n;
                int next_tile_m = (work_idx + 1) % kparams.fused_tiles_m;
                int next_coord_n = next_my_n_tile * BLOCK_N;
                int next_coord_m = next_tile_m * BLOCK_M;
                int next_k_iters = kparams.groups[next_group_id].K / BLOCK_K;

                const void* next_tmap_B   = &kparams.tmaps[FUSED_TMAP_B_BASE   + next_group_id];
                const void* next_tmap_SFB = &kparams.tmaps[FUSED_TMAP_SFB_BASE + next_group_id];
                const void* next_tmap_A   = &kparams.tmaps[FUSED_TMAP_A];
                const void* next_tmap_SFA = &kparams.tmaps[FUSED_TMAP_SFA];

                // Prefetch B+SFB: each CTA prefetches its own N-slice (all ranks)
                // Prefetch up to NUM_STAGES B stages (covers full pipeline depth)
                #pragma unroll
                for (int k_pf = 0; k_pf < NUM_STAGES && k_pf < next_k_iters; k_pf++)
                {
                    tma_2d_prefetch(next_tmap_B, k_pf * BLOCK_K, next_coord_n, EVICT_FIRST);
                    int off_sfb = (next_my_n_tile * next_k_iters + k_pf) * SMEM_SFB_SIZE;
                    tma_1d_prefetch(next_tmap_SFB, off_sfb / 8, EVICT_FIRST);
                }

                // Prefetch A+SFA: multicast-gated to rank-0 only (avoids duplicate HBM reads)
                // Non-multicast: all ranks prefetch their own (but A is same so 1 rank is enough)
                // Prefetch NUM_STAGES stages of A: L2 is large enough (96KB << 4MB per cluster)
                // and A uses EVICT_LAST so it won't displace B tiles.
                if (!use_multicast || cta_n == 0)
                {
                    #pragma unroll
                    for (int k_pf = 0; k_pf < NUM_STAGES && k_pf < next_k_iters; k_pf++)
                    {
                        tma_2d_prefetch(next_tmap_A, k_pf * BLOCK_K, next_coord_m, EVICT_LAST);
                        int off_sfa = (next_tile_m * next_k_iters + k_pf) * SMEM_SFA_SIZE;
                        tma_1d_prefetch(next_tmap_SFA, off_sfa / 8, EVICT_LAST);
                    }
                }
            }
        }

        // ALL threads: cluster_sync ensures all CTAs see reinitialized barriers
        if (has_next_tile && use_multicast)
        {
            if (warp_id == TMA_WARP) DBG("TMA: csync\n");
            cluster_sync();
        }

        // TMA-MMA bar_sync: ensures MMA doesn't start K-loop before barriers ready
        if (warp_id == TMA_WARP && has_next_tile)
        {
            DBG("TMA: wait bar_sync\n");
            bar_sync_tma_mma();
            DBG("TMA: bar_sync ok\n");
        }

        // -----------------------------------------------------------------
        // EPILOGUE WARPS (0-3): Drain TMEM to global memory
        // -----------------------------------------------------------------

        if (warp_id < 4)
        {
            if (has_valid_work)
            {
                // Spin-check that MMA has reinit'd the barriers for this bp.
                // Necessary because epilogue warps can reach here before MMA's lane 0
                // has completed mbarrier_init + epi_barriers_ready write.
                while (smem->epi_barriers_ready[bp] == 0) {}

                if (warp_id == 0) DBG("EPI: wait epi\n");
                // Wait for MMA to signal accumulator is ready.
                // Phase 0: barrier is always reinit'd (parity reset) before this wait.
                mbarrier_wait(get_mbar_addr(&smem->epilogue_mbar[bp]), 0);
                if (warp_id == 0) DBG("EPI: draining\n");

                // CRITICAL: Fence required between MMA/TMA and tcgen05_ld
                tcgen05_fence_after_thread_sync();

                // Distribute work: 1 tile per warp
                // Warps 0, 1, 2, 3 handle 32 rows each -> 128 rows total
                int row_tile = warp_id;
                int base_row = row_tile * 32;

                // Phase 1: Drain and store in CHUNKS to reduce register pressure
                // BLOCK_N=128, handle in 4 chunks of 32 columns each.
                // This reduces tile_data from 64 registers/thread to 32.
                constexpr int CHUNK_SIZE = 32;
                constexpr int NUM_CHUNKS_PER_ITER = CHUNK_SIZE / 8;
                constexpr int NUM_ITERATIONS = BLOCK_N / CHUNK_SIZE;
                float tile_data[NUM_CHUNKS_PER_ITER][8];

                half *C_ptr_dst = reinterpret_cast<half *>(params.C_ptr);
                int local_row = base_row + lane_id;
                int c_row = coord_m - epi_coord_m_offset + local_row;

                #pragma unroll
                for (int chunk_idx = 0; chunk_idx < NUM_ITERATIONS; chunk_idx++)
                {
                    int chunk_col_base = chunk_idx * CHUNK_SIZE;

                    // Load chunk from TMEM to registers
                    #pragma unroll
                    for (int c_chunk = 0; c_chunk < NUM_CHUNKS_PER_ITER; c_chunk++)
                    {
                        tcgen05_ld_32x32b<8>(tile_data[c_chunk], acc_tmem[bp % NUM_ACC_BUFS], base_row, chunk_col_base + c_chunk * 8);
                    }
                    tcgen05_wait_alloc(); // wait for this chunk's load

                    // Store chunk to GMEM
                    if (c_row < params.M)
                    {
                        #pragma unroll
                        for (int c_chunk = 0; c_chunk < NUM_CHUNKS_PER_ITER; c_chunk++)
                        {
                            int base_col = chunk_col_base + c_chunk * 8;
                            // int4 is 16-byte aligned by type; half2[4] is only 4-byte aligned
                            // and can spill to local memory causing misaligned address fault.
                            int4 out;
                            #pragma unroll
                            for (int i = 0; i < 4; i++)
                            {
                                reinterpret_cast<half2 *>(&out)[i] = __float22half2_rn({tile_data[c_chunk][i * 2], tile_data[c_chunk][i * 2 + 1]});
                            }
                            reinterpret_cast<int4 *>(C_ptr_dst + c_row * params.N + coord_n + base_col)[0] = out;
                        }
                    }
                }

                // Canonical PTX return handoff fence: guarantee ld visibility before MMA reuse
                tcgen05_wait_alloc(); 
                tcgen05_fence_before_thread_sync();

                // Signal that this warp's TMEM drain is complete
                if (lane_id == 0)
                {
                    if (warp_id == 0) DBG("EPI: done\n");
                    mbarrier_arrive(get_mbar_addr(&smem->epilogue_done_mbar[bp]));
                }
                // Clear the ready flag so this bp can be safely reinit'd next time
                smem->epi_barriers_ready[bp] = 0;
            }
            else
            {
                // No valid work but still must arrive to avoid epilogue_done_mbar deadlock
                if (lane_id == 0)
                {
                    if (warp_id == 0) DBG("EPI: done(nw)\n");
                    mbarrier_arrive(get_mbar_addr(&smem->epilogue_done_mbar[bp]));
                }
            }
        }

        tile_counter++;
    }

    // =========================================================================
    // LAST TILE CLEANUP: Wait for final epilogue to complete
    // =========================================================================

    // MMA warp must wait for last epilogue_done before deallocating TMEM.
    // Phase 0: last tile's epi_done was reinit'd by MMA guard (parity reset to 0),
    // then epilogue arrives 4x → parity 1. Wait with phase 0 → passes.
    // Special case tile 0 (no reinit): prologue init'd it to parity 0, epilogue → 1.
    if (warp_id == MMA_WARP && tile_counter > 0)
    {
        int wait_bp = (tile_counter - 1) & 1;
        mbarrier_wait(get_mbar_addr(&smem->epilogue_done_mbar[wait_bp]), 0);
    }

    // Full sync before TMEM deallocation
    __syncthreads();

    // =========================================================================
    // TMEM Deallocation — ONCE after the tile loop
    // =========================================================================

    if (warp_id == MMA_WARP)
    {
        tcgen05_relinquish_alloc_permit<1>();
        tcgen05_dealloc<1>(0, TMEM_ALLOC_COLS);
        tcgen05_wait_alloc();
    }

    __syncthreads();
    if constexpr (CLUSTER_N > 1)
    {
        cluster_sync();
    }
}

// =============================================================================
// Explicit Template Instantiations
// =============================================================================

// Format: <CLUSTER_N, NUM_STAGES, K_EXPECTED, BLOCK_N>
// K=0 fallback only — K-specialized templates cause icache pressure regression.

template __global__ void group_gemm_kernel_impl<1, 4, 0, 128>(const __grid_constant__ KernelParams kparams);
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, 4, 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, 4, 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 Source: Wrapper
# =============================================================================

CUDA_SRC_WRAPPER = r"""
// wrapper.cu - PyTorch interface for group GEMM kernel
// Phase 6B: Epilogue-TMA overlap — persistent grid with capped clusters

#include <torch/library.h>
#include <ATen/core/Tensor.h>
#include <cuda.h>
#include <cuda_runtime.h>

// utils.h included above

// =============================================================================
// Forward Declarations for Templated Kernel
// =============================================================================

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

// =============================================================================
// Shared Memory Configuration
// =============================================================================
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);
// Double-buffered barriers: full_mbar[2][N] + empty_mbar[2][N] + epilogue_mbar[2]
// + epilogue_done_mbar[2] + tmem_holding_buf + epi_barriers_ready[2]
constexpr int SMEM_BARRIER_SIZE = sizeof(uint64_t) * (NUM_STAGES_MAX * 4 + 4 + 1) + sizeof(int) * 2;
constexpr int SMEM_SIZE = SMEM_DATA_SIZE + SMEM_BARRIER_SIZE + 128; // Max (NUM_STAGES_MAX)

// Dynamic SMEM size computation for variable pipeline depth
inline int compute_smem_size(int num_stages) {
    int data = num_stages * (SMEM_A_ALIGNED + SMEM_B_ALIGNED + SMEM_SFA_ALIGNED + SMEM_SFB_ALIGNED);
    int barriers = sizeof(uint64_t) * (num_stages * 4 + 4 + 1) + sizeof(int) * 2;
    return data + barriers + 128;  // +128 for alignment padding
}

// =============================================================================
// CUDA Driver API Error Checking
// =============================================================================

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

// =============================================================================
// TensorMap Creation Functions
// =============================================================================

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_L2_64B, // A is small, reused across N-tiles
        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, // B reused across M-tiles with block distribution
        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_L2_64B,
        CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
    check_cu(err, "cuTensorMapEncodeTiled for SF (1D)");
}

// =============================================================================
// Dynamic Cluster Selection
// =============================================================================

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

    // Compute min K across groups — cross-CTA empty_mbar overhead is proportionally
    // larger for few k_iters. Cap cluster size for low-K shapes where multicast
    // savings don't outweigh the per-k_iter remote arrive cost.
    int min_K = INT_MAX;
    for (int g = 0; g < num_groups; g++)
    {
        int K = sizes_acc[g][2];
        if (K < min_K) min_K = K;
    }

    // CLUSTER_N=4 with short K loops (K=2048, 8 k_iters) causes ~80% regression:
    // empty_mbar init_count=4 forces TMA to wait for all 4 CTAs, and synchronization
    // overhead on short pipelines dominates over A bandwidth savings. Tested: BM2 41→75µs.
    // CLUSTER_N=2 for BM4 (K=1536, 6 k_iters): neutral (10.6→10.7µs), keep it.
    int max_cluster = (min_K >= 4096) ? 4 : (min_K >= 1024) ? 2 : 1;

    if (max_cluster >= 4 && all_valid(4))
        return 4;
    if (max_cluster >= 2 && 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>();

    // Select pipeline depth based on arithmetic intensity.
    // Higher AI (compute-bound) → more stages to hide latency.
    // Lower AI (memory-bound) → fewer stages (diminishing returns).
    // TODO: Enable stages 3-4 for 2-CTA occupancy after kernel validation.
    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;
    }

    DEBUG_PRINT("[PIPELINE] min_ai=%.2f\n", min_ai);
    if (min_ai > 150.0f) return 6;
    if (min_ai > 100.0f) return 5;
    return 4;
}

// =============================================================================
// Launch Cache
// =============================================================================

struct LaunchCache {
    uintptr_t ptrs[MAX_GROUPS * 5];  // A,B,C,SFA,SFB per group
    int dims[MAX_GROUPS * 4];        // M,N,K,L per group
    int num_groups;
    int cluster_n;
    KernelParams kparams;
    bool valid = false;
};

// =============================================================================
// Main Kernel Launch Function
// =============================================================================

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

    // Determine optimal cluster size based on problem dimensions
    int cluster_n = compute_cluster_n(problem_sizes);

    // Determine optimal pipeline depth based on arithmetic intensity
    int num_stages = compute_num_stages(problem_sizes);
    // Use dynamic SMEM size for actual num_stages.
    int smem_size = compute_smem_size(num_stages);
    DEBUG_PRINT("[PIPELINE] selected num_stages=%d, smem=%d bytes\n", num_stages, smem_size);

    // Access data on CPU
    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>();

    // =========================================================================
    // Check launch cache
    // =========================================================================
    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) {
        // Reuse cached params (avoids cuTensorMapEncodeTiled calls)
        kparams = cache.kparams;
    } else {
        // Build KernelParams from scratch
        memset(&kparams, 0, sizeof(kparams));

        // Check if all groups share same N, K, L (eligible for M-fusion)
        bool can_fuse = (num_groups > 1);
        if (can_fuse) {
            int N0 = sizes_acc[0][1], K0 = sizes_acc[0][2], L0 = sizes_acc[0][3];
            for (int g = 1; g < num_groups; g++) {
                if (sizes_acc[g][1] != N0 || sizes_acc[g][2] != K0 || sizes_acc[g][3] != L0) {
                    can_fuse = false;
                    break;
                }
            }
        }

        if (can_fuse) {
            // =========================================================
            // Fused mode: single A/SFA TensorMap, per-group B/SFB
            // =========================================================
            int N = sizes_acc[0][1];
            int K = sizes_acc[0][2];
            int L = sizes_acc[0][3];

            // Python passes fused data as num_groups entries where:
            //   - All groups share the same A_ptr (fused A) and SFA_ptr (fused SFA)
            //   - Each group has its own B_ptr, SFB_ptr, C_ptr
            //   - sizes[g] = [M_g (actual), N, K, L]  (per-group actual M)
            // M_fused = sum of ceil(M_g/BLOCK_M)*BLOCK_M across all groups

            // Compute M_fused from per-group Ms (each padded to BLOCK_M)
            int M_fused = 0;
            for (int g = 0; g < num_groups; g++) {
                M_fused += ((sizes_acc[g][0] + BLOCK_M - 1) / BLOCK_M) * BLOCK_M;
            }
            void *A_fused_ptr = reinterpret_cast<void *>(abc_acc[0][0]);
            void *SFA_fused_ptr = reinterpret_cast<void *>(sf_acc[0][0]);

            int tiles_m_total = (M_fused + BLOCK_M - 1) / BLOCK_M;
            int tiles_n = N / (BLOCK_N * cluster_n);

            TORCH_CHECK(tiles_m_total <= MAX_M_TILES,
                "Too many M-tiles for fused mode: ", tiles_m_total);

            kparams.fused = 1;
            kparams.fused_tiles_m = static_cast<uint16_t>(tiles_m_total);
            kparams.fused_tiles_n = static_cast<uint16_t>(std::max(1, tiles_n));
            kparams.total_tiles = tiles_m_total * std::max(1, tiles_n);
            kparams.num_groups = static_cast<uint32_t>(num_groups);

            // Build tile_group[] mapping and per-group params
            int tile_m_cursor = 0;
            for (int g = 0; g < num_groups; g++) {
                int M_g = sizes_acc[g][0];
                int group_tiles_m = (M_g + BLOCK_M - 1) / BLOCK_M;

                // In fused mode, tile_offset stores the cumulative M-tile offset
                // (used in epilogue to convert fused coord_m to per-group C row)
                kparams.groups[g].tile_offset = tile_m_cursor;
                kparams.groups[g].M = static_cast<uint16_t>(M_g);
                kparams.groups[g].N = static_cast<uint16_t>(N);
                kparams.groups[g].K = static_cast<uint16_t>(K);
                kparams.groups[g].L = static_cast<uint16_t>(L);
                kparams.groups[g].tiles_m = static_cast<uint16_t>(group_tiles_m);
                kparams.groups[g].tiles_n = static_cast<uint16_t>(std::max(1, tiles_n));
                kparams.groups[g].padding = 0;

                // C_ptr is per-group for epilogue writeback
                kparams.groups[g].C_ptr = reinterpret_cast<void *>(abc_acc[g][2]);
                // A/B/SF ptrs stored for reference but tmaps are separate
                kparams.groups[g].A_ptr = A_fused_ptr;
                kparams.groups[g].B_ptr = reinterpret_cast<void *>(abc_acc[g][1]);
                kparams.groups[g].SFA_ptr = SFA_fused_ptr;
                kparams.groups[g].SFB_ptr = reinterpret_cast<void *>(sf_acc[g][1]);

                for (int t = 0; t < group_tiles_m; t++) {
                    kparams.tile_group[tile_m_cursor + t] = static_cast<uint8_t>(g);
                }
                tile_m_cursor += group_tiles_m;
            }

            // Populate prefetch helpers: next tile's group and n-increment
            for (int t = 0; t < tiles_m_total; t++) {
                int next_t = (t + 1);
                if (next_t < tiles_m_total) {
                    kparams.next_tile_group[t] = kparams.tile_group[next_t];
                    kparams.next_tile_n[t] = 0; // stays in same N-tile column
                } else {
                    kparams.next_tile_group[t] = kparams.tile_group[0];
                    kparams.next_tile_n[t] = 1; // wraps to next N-tile column
                }
            }

            // Fused A TensorMap (covers entire fused M_total)
            int M_fused_padded = tiles_m_total * BLOCK_M;
            init_A_tmap(&kparams.tmaps[FUSED_TMAP_A], A_fused_ptr, M_fused, K, L, BLOCK_M, BLOCK_K);
            init_SF_tmap(&kparams.tmaps[FUSED_TMAP_SFA], SFA_fused_ptr, M_fused_padded, K);

            // Per-group B/SFB TensorMaps
            for (int g = 0; g < num_groups; g++) {
                void *B_ptr = reinterpret_cast<void *>(abc_acc[g][1]);
                void *SFB_ptr = reinterpret_cast<void *>(sf_acc[g][1]);
                int N_g = N;
                int N_padded = ((N_g + BLOCK_N - 1) / BLOCK_N) * BLOCK_N;

                init_B_tmap(&kparams.tmaps[FUSED_TMAP_B_BASE + g], B_ptr, N_g, K, L, BLOCK_N, BLOCK_K);
                init_SF_tmap(&kparams.tmaps[FUSED_TMAP_SFB_BASE + g], SFB_ptr, N_padded, K);
            }

            DEBUG_PRINT("[FUSED] M_fused=%d, tiles_m=%d, tiles_n=%d, total_tiles=%u\n",
                       M_fused, tiles_m_total, tiles_n, kparams.total_tiles);
        } else {
            // =========================================================
            // Legacy mode: per-group TensorMaps
            // =========================================================
            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);

            // Fill TensorMaps
            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;

                // A/B use OOB_FILL_ZERO: pass actual M,N — TMA zero-fills partial tiles
                init_A_tmap(&kparams.tmaps[base_idx + TMAP_A_IDX], A_ptr, M, K, L, BLOCK_M, BLOCK_K);
                init_B_tmap(&kparams.tmaps[base_idx + TMAP_B_IDX], B_ptr, N, 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);
            }
        }

        // Store to cache
        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;
    }

    // Persistent grid: cap clusters to max resident on device.
    // Each cluster processes multiple tiles when grid is smaller than total_tiles.
    // Query SM count once (cached).
    static int sm_count = 0;
    if (sm_count == 0)
    {
        int dev;
        cudaGetDevice(&dev);
        cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, dev);
    }

    // Max clusters = SM count / cluster_n (each cluster occupies cluster_n SMs)
    // B200 has 192 SMs. With cluster_n=2, max = 96 clusters.
    uint32_t max_clusters = sm_count / cluster_n;
    uint32_t num_clusters = std::min(kparams.total_tiles, max_clusters);
    uint32_t grid_y = num_clusters * cluster_n;

    DEBUG_PRINT("[LAUNCH] num_groups=%d, total_tiles=%u, grid_y=%u\n",
           num_groups, kparams.total_tiles, grid_y);
    DEBUG_PRINT("[LAUNCH] smem_size=%d bytes, THREADS_PER_CTA=%d\n", smem_size, THREADS_PER_CTA);
    DEBUG_PRINT("[LAUNCH] Cluster dims: (%d, %d, %d)\n", CLUSTER_M, cluster_n, CLUSTER_Z);

    // Select kernel function based on cluster_n and num_stages
    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 if (num_stages >= 5)
            kernel_func = (const void *)group_gemm_kernel_impl<4, 5, 0, 128>;
        else
            kernel_func = (const void *)group_gemm_kernel_impl<4, 4, 0, 128>;
    }
    else if (cluster_n == 2)
    {
        if (num_stages >= 6)
            kernel_func = (const void *)group_gemm_kernel_impl<2, 6, 0, 128>;
        else if (num_stages >= 5)
            kernel_func = (const void *)group_gemm_kernel_impl<2, 5, 0, 128>;
        else
            kernel_func = (const void *)group_gemm_kernel_impl<2, 4, 0, 128>;
    }
    else
    {
        if (num_stages >= 6)
            kernel_func = (const void *)group_gemm_kernel_impl<1, 6, 0, 128>;
        else if (num_stages >= 5)
            kernel_func = (const void *)group_gemm_kernel_impl<1, 5, 0, 128>;
        else
            kernel_func = (const void *)group_gemm_kernel_impl<1, 4, 0, 128>;
    }

    // Set maximum dynamic shared memory size (cached per kernel variant)
    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;
    }

    // Grid dimensions: (1, grid_y, 1) for cluster (1, cluster_n, 1)
    dim3 grid(1, grid_y, 1);
    dim3 block(THREADS_PER_CTA, 1, 1);

    // Launch with cluster
    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;

    // Use cudaLaunchKernelEx (variadic) - launches on default CUDA context
    // matching reference implementations (gau.nernst, shiyegao, hekailove)
    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 Registration
// =============================================================================

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

"""

# =============================================================================
# Compile the inline CUDA module
# =============================================================================

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

# =============================================================================
# FP4 E2M1 Dequantization
# =============================================================================

# E2M1 lookup table: 4-bit index -> float value
# Format: 1 sign bit, 2 exponent bits, 1 mantissa bit, bias=1
_E2M1_LUT = torch.tensor([
    0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0,      # positive values (0-7)
    -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0  # negative values (8-15)
], dtype=torch.float16)

def dequantize_fp4(packed: torch.Tensor, rows: int, K: int, L: int, device) -> torch.Tensor:
    """
    Dequantize packed FP4 tensor (float4_e2m1fn_x2) to FP16.

    packed: tensor with shape [rows, K//2, L] containing packed FP4 pairs
    Returns: [rows, K, L] in FP16
    """
    # Total number of byte pairs (each byte = 2 FP4 values)
    num_bytes = rows * (K // 2) * L

    # Flatten and view as uint8 to handle any exotic dtype layout
    packed_u8 = packed.contiguous().view(torch.uint8).reshape(num_bytes)

    # Extract low and high nibbles (2 FP4 values per byte)
    low_nibble = (packed_u8 & 0x0F).long()
    high_nibble = ((packed_u8 >> 4) & 0x0F).long()

    # Lookup dequantized values
    lut = _E2M1_LUT.to(device)
    low_vals = lut[low_nibble]
    high_vals = lut[high_nibble]

    # Interleave low and high values
    result = torch.stack([low_vals, high_vals], dim=1).reshape(-1)

    # Reshape to [rows, K, L]
    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:
    """
    Apply block scale factors to tensor.

    tensor: [rows, K, L]
    sf: scale factor tensor (may have padded/tiled layout)

    Returns: [rows, K, L] with scale factors applied
    """
    # Convert scale factor to fp16 and move to GPU (sfasfb_tensors are on CPU)
    sf_fp16 = sf.to(device=device, dtype=torch.float16).flatten()

    # Calculate dimensions - sf may be padded to power of 2
    sf_k_blocks = K // 16
    total_sf = sf_fp16.numel()
    padded_rows = total_sf // (sf_k_blocks * L)

    # Reshape to [padded_rows, K//16, L]
    sf_reshaped = sf_fp16.reshape(padded_rows, sf_k_blocks, L)

    # Take only the rows we need (remove padding)
    sf_trimmed = sf_reshaped[:rows, :, :]

    # Expand each scale factor to cover 16 K elements
    sf_expanded = sf_trimmed.repeat_interleave(16, dim=1)  # [rows, K, L]

    return tensor * sf_expanded


# Cache for prepared launch data (avoids re-preparing on repeated calls)
_launch_cache = dict()

def _merge_groups(abc_tensors, sfasfb_reordered, problem_sizes):
    """
    M-fusion: when all groups share the same N, K, L, fuse them into a single
    GEMM by concatenating A/SFA along M (padded to 128 per group).
    The kernel uses per-tile B selection via tile_group[] mapping.

    Falls back to per-group pass-through when groups have different N, K, or L.

    Returns (merged_abc_ptrs, merged_sf_ptrs, merged_sizes, refs, c_copyback).
    """
    num_groups = len(problem_sizes)

    # Check if all groups share same N, K, L (eligible for M-fusion)
    N0, K0, L0 = problem_sizes[0][1], problem_sizes[0][2], problem_sizes[0][3]
    can_fuse = num_groups > 1 and all(
        ps[1] == N0 and ps[2] == K0 and ps[3] == L0
        for ps in problem_sizes[1:]
    )

    if can_fuse:
        # M-fusion: concatenate A/SFA along M, keep per-group B/SFB/C
        # cuBLAS SF reorder format already pads M to 128 (rest_m = ceil(M/128)),
        # so each group's SFA is already correctly sized for TMA.
        # A needs explicit padding to BLOCK_M=128 boundary.
        BLOCK_M = 128
        a_parts = []
        sfa_parts = []

        for idx in range(num_groups):
            a_i = abc_tensors[idx][0]
            M_i = problem_sizes[idx][0]
            pad_m = ((M_i + BLOCK_M - 1) // BLOCK_M) * BLOCK_M

            # Pad A to BLOCK_M boundary. FP4 packed dtype doesn't support fill_,
            # so allocate as uint8 (same byte layout), zero-fill, then view as FP4.
            if pad_m > M_i:
                a_padded = torch.zeros(pad_m, a_i.shape[1], *a_i.shape[2:],
                                       dtype=torch.uint8, device=a_i.device).view(a_i.dtype)
                a_padded[:M_i] = a_i
            else:
                a_padded = a_i

            # SFA: cuBLAS reorder already has ceil(M/128)*128 worth of SF rows
            sfa_tma_i = sfasfb_reordered[idx][0].view(torch.uint8).permute(5, 2, 4, 0, 1, 3).reshape(-1, 16)

            a_parts.append(a_padded)
            sfa_parts.append(sfa_tma_i)

        # Concatenate A and SFA along M dimension
        a_fused = torch.cat(a_parts, dim=0)
        sfa_fused = torch.cat(sfa_parts, dim=0)

        # Per-group SFB preparation
        sfb_tmas = []
        for idx in range(num_groups):
            sfb_tma = sfasfb_reordered[idx][1].view(torch.uint8).permute(5, 2, 4, 0, 1, 3).reshape(-1, 16)
            sfb_tmas.append(sfb_tma)

        # Build per-group abc_ptrs/sf_ptrs/sizes arrays
        # All groups share fused A_ptr and SFA_ptr; each has own B, SFB, C
        abc_ptrs = []
        sf_ptrs = []
        sizes = []
        c_copyback = []

        for idx in range(num_groups):
            M_i = problem_sizes[idx][0]
            abc_ptrs.append([a_fused.data_ptr(), abc_tensors[idx][1].data_ptr(),
                            abc_tensors[idx][2].data_ptr()])
            sf_ptrs.append([sfa_fused.data_ptr(), sfb_tmas[idx].data_ptr()])
            sizes.append([M_i, N0, K0, L0])

        # Keep references alive
        refs = [a_fused, sfa_fused] + sfb_tmas + [a_parts, sfa_parts]

        abc_ptrs_t = torch.tensor(abc_ptrs, dtype=torch.int64)
        sf_ptrs_t = torch.tensor(sf_ptrs, dtype=torch.int64)
        sizes_t = torch.tensor(sizes, dtype=torch.int32)
        return abc_ptrs_t, sf_ptrs_t, sizes_t, refs, c_copyback

    # Non-fusable: pass through each group individually
    abc_ptrs = []
    sf_ptrs = []
    sizes = []
    refs = []
    c_copyback = []

    for idx in range(num_groups):
        a, b, c = abc_tensors[idx]
        sfa_r = sfasfb_reordered[idx][0]
        sfb_r = sfasfb_reordered[idx][1]
        M = problem_sizes[idx][0]
        N, K, L = problem_sizes[idx][1], problem_sizes[idx][2], problem_sizes[idx][3]

        abc_ptrs.append([a.data_ptr(), b.data_ptr(), c.data_ptr()])
        sfa_tma = sfa_r.view(torch.uint8).permute(5, 2, 4, 0, 1, 3).reshape(-1, 16)
        sfb_tma = sfb_r.view(torch.uint8).permute(5, 2, 4, 0, 1, 3).reshape(-1, 16)
        refs.append((sfa_r, sfb_r, sfa_tma, sfb_tma))
        sf_ptrs.append([sfa_tma.data_ptr(), sfb_tma.data_ptr()])
        sizes.append([M, N, K, L])

    abc_ptrs_t = torch.tensor(abc_ptrs, dtype=torch.int64)
    sf_ptrs_t = torch.tensor(sf_ptrs, dtype=torch.int64)
    sizes_t = torch.tensor(sizes, dtype=torch.int32)
    return abc_ptrs_t, sf_ptrs_t, sizes_t, refs, c_copyback


def _copyback_c(c_copyback):
    """Copy merged C slices back to original C tensors."""
    for c_merged, slices in c_copyback:
        for c_orig, start, end in slices:
            c_orig.copy_(c_merged[start:end])


def custom_kernel_cuda(data):
    """
    Group GEMM kernel with opportunistic M-fusion for shared-B groups.
    Caches prepared tensors to avoid redundant work on repeated calls.
    """
    abc_tensors, sfasfb_tensors, sfasfb_reordered, problem_sizes = data
    num_groups = len(problem_sizes)

    # Build cache key from tensor data pointers and 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_t, sf_ptrs_t, sizes_t, _refs, c_copyback = cached
        torch.ops.group_gemm_module.group_gemm_launch(abc_ptrs_t, sf_ptrs_t, sizes_t)
        _copyback_c(c_copyback)
        return [t[2] for t in abc_tensors]

    # Cache miss: merge groups sharing B, then prepare
    abc_ptrs_t, sf_ptrs_t, sizes_t, refs, c_copyback = _merge_groups(
        abc_tensors, sfasfb_reordered, problem_sizes)

    _launch_cache[cache_key] = (abc_ptrs_t, sf_ptrs_t, sizes_t, refs, c_copyback)

    torch.ops.group_gemm_module.group_gemm_launch(abc_ptrs_t, sf_ptrs_t, sizes_t)
    _copyback_c(c_copyback)
    return [t[2] for t in abc_tensors]


def custom_kernel_fallback(data):
    """
    Group GEMM kernel using PyTorch fallback implementation.
    """
    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

        # Dequantize FP4 to FP16 using known dimensions
        a_fp16 = dequantize_fp4(a, M, K, L, device)  # [M, K, L]
        b_fp16 = dequantize_fp4(b, N, K, L, device)  # [N, K, L]

        # Apply scale factors
        a_scaled = apply_scale_factors(a_fp16, sfa, M, K, L, device)  # [M, K, L]
        b_scaled = apply_scale_factors(b_fp16, sfb, N, K, L, device)  # [N, K, L]

        # GEMM: C[m,n,l] = sum_k A[m,k,l] * B[n,k,l]
        # For L batches: C = A @ B.transpose(-2, -1)
        # Reshape for batched matmul: [L, M, K] @ [L, K, N] -> [L, M, N]
        a_batched = a_scaled.permute(2, 0, 1)  # [L, M, K]
        b_batched = b_scaled.permute(2, 1, 0)  # [L, K, N] (transpose K and N)

        c_batched = torch.bmm(a_batched.float(), b_batched.float())  # [L, M, N]

        # Store result back in c tensor
        c_result = c_batched.permute(1, 2, 0).to(torch.float16)  # [M, N, L]
        c.copy_(c_result)
        results.append(c)

    return results


# Select implementation based on environment variable
# USE_CUDA_KERNEL=0 to use fallback, otherwise use CUDA kernel (default)
USE_CUDA_KERNEL = os.environ.get('USE_CUDA_KERNEL', '1') == '1'

def custom_kernel(data):
    """
    Group GEMM kernel entry point for POPCORN benchmark.

    Computes C = A @ B.T for each group with block-scaled FP4 inputs.

    Input format:
        data = (abc_tensors, sfasfb_tensors, sfasfb_reordered, problem_sizes)
    Where:
        abc_tensors: list of tuples (a, b, c) - ON GPU
        sfasfb_tensors: list of tuples (sfa, sfb) - reference format, may be CPU
        sfasfb_reordered: list of tuples - cuBLAS format, ON GPU
        problem_sizes: list of tuples (M, N, K, L)
    """
    if USE_CUDA_KERNEL:
        return custom_kernel_cuda(data)
    else:
        return custom_kernel_fallback(data)
scrolls · 2751 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 489762.

⋯ 632 unchanged lines
constexpr int TMAP_SFA_IDX = 3;
constexpr int TMAP_SFB_IDX = 4;
+ // Fused mode TensorMap layout:
+ // tmaps[0] = fused A, tmaps[1] = fused SFA
+ // tmaps[FUSED_TMAP_B_BASE + g] = group g's B
+ // tmaps[FUSED_TMAP_SFB_BASE + g] = group g's SFB
+ constexpr int FUSED_TMAP_A = 0;
+ constexpr int FUSED_TMAP_SFA = 1;
+ constexpr int FUSED_TMAP_B_BASE = 2;
+ constexpr int FUSED_TMAP_SFB_BASE = 2 + MAX_GROUPS; // 18
+
+ // Max M-tiles for tile_group array
+ constexpr int MAX_M_TILES = 256;
+
// Unified kernel parameters passed as __grid_constant__
// Contains all group params + TensorMaps in a single struct
- // Size: 16*72 + 80*128 + 8 ≈ 11,400 bytes (within 32KB limit)
+ // Size: 16*72 + 80*128 + 8 + 256 + 8 ≈ 11,700 bytes (within 32KB limit)
struct alignas(128) KernelParams {
GroupParams groups[MAX_GROUPS];
CUtensorMap tmaps[MAX_GROUPS * TENSORMAPS_PER_GROUP];
uint32_t total_tiles;
uint32_t num_groups;
+
+ // M-fusion fields (active when fused != 0)
+ uint8_t tile_group[MAX_M_TILES]; // maps M-tile index -> original group_id
+ uint16_t fused_tiles_m; // total fused M-tiles
+ uint16_t fused_tiles_n; // N-tiles (per cluster column)
+ uint32_t fused; // 1 = fused mode, 0 = legacy
+
+ // Low-overhead prefetch support
+ uint8_t next_tile_group[MAX_M_TILES]; // group_id for (tile_m + 1) % fused_tiles_m
+ uint16_t next_tile_n[MAX_M_TILES]; // tile_n increment for tile_m + 1 (0 or 1)
};
// Scale Factor Layout for tcgen05 (matching gau.nernst reference):
⋯ 56 unchanged lines
// Single-buffered accumulator: cta_group::1 limits TMEM to 256 columns per CTA.
// Double-buffered (2×128+32+32=320) exceeds this limit, causing tcgen05_alloc to stall.
// TMA-epilogue overlap is still achieved via double-buffered K-loop barriers (full/empty_mbar[2][*]).
- constexpr int NUM_ACC_BUFS = 1;
+ constexpr int NUM_ACC_BUFS = 2;
+ constexpr int TMEM_TOTAL_COLS = NUM_ACC_BUFS * TMEM_ACC_COLS + TMEM_SFA_COLS + TMEM_SFB_COLS; // 2*128 + 32 + 32 = 320
+ constexpr int TMEM_ALLOC_COLS = (TMEM_TOTAL_COLS <= 256) ? 256 : 512; // 512 columns (power of 2)
- constexpr int TMEM_TOTAL_COLS = NUM_ACC_BUFS * TMEM_ACC_COLS + TMEM_SFA_COLS + TMEM_SFB_COLS; // 1*128 + 32 + 32 = 192
- constexpr int TMEM_ALLOC_COLS = TMEM_TOTAL_COLS; // 192 columns (fits in 256-col cta_group::1 limit)
-
// =============================================================================
// Shared Memory Layout
// =============================================================================
⋯ 23 unchanged lines
constexpr int TMA_BYTES_AB = TMA_A_BYTES + TMA_B_BYTES; // 32768 bytes
constexpr int TMA_BYTES_ALL = TMA_A_BYTES + TMA_B_BYTES + TMA_SFA_BYTES + TMA_SFB_BYTES; // 36864 bytes
- struct SmemBuffers
+ template <int STAGES>
+ struct SmemBuffersT
{
// [CRITICAL] Data buffers MUST come FIRST to ensure 128-byte alignment
// Putting them first means they inherit the struct's base alignment
⋯ 14 unchanged lines
alignas(128) char data[SMEM_SFB_SIZE];
};
- // PHASE 3: Allocate for maximum stages (8), actual usage determined by template param
- AlignedBuffA A_smem[NUM_STAGES_MAX];
- AlignedBuffB B_smem[NUM_STAGES_MAX];
- AlignedBuffSFA SFA_smem[NUM_STAGES_MAX];
- AlignedBuffSFB SFB_smem[NUM_STAGES_MAX];
+ AlignedBuffA A_smem[STAGES];
+ AlignedBuffB B_smem[STAGES];
+ AlignedBuffSFA SFA_smem[STAGES];
+ AlignedBuffSFB SFB_smem[STAGES];
// Mbarriers and metadata come AFTER data buffers (8-byte alignment is sufficient)
// Double-buffered barriers: [bp] where bp = tile_counter & 1
// Even/odd tiles use separate barrier sets — no re-init conflicts during overlap
- alignas(8) uint64_t full_mbar[2][NUM_STAGES_MAX]; // TMA signals, MMA waits
- alignas(8) uint64_t empty_mbar[2][NUM_STAGES_MAX]; // MMA signals, TMA waits
+ alignas(8) uint64_t full_mbar[2][STAGES]; // TMA signals, MMA waits
+ alignas(8) uint64_t empty_mbar[2][STAGES]; // MMA signals, TMA waits
alignas(8) uint64_t epilogue_mbar[2]; // MMA signals, epilogue waits
alignas(8) uint64_t epilogue_done_mbar[2]; // Epilogue signals, MMA waits (TMEM free)
alignas(8) uint64_t tmem_holding_buf; // Used by tcgen05_alloc
- // Flags for epilogue barrier readiness (set by TMA after init, checked by epilogue)
+ // Flags for epilogue barrier readiness (set by MMA after reinit, checked by epilogue)
volatile int epi_barriers_ready[2];
};
+ // Backward-compatible alias for maximum stage count
+ using SmemBuffers = SmemBuffersT<NUM_STAGES_MAX>;
+
// =============================================================================
// SMEM Descriptor Building
// =============================================================================
⋯ 209 unchanged lines
// Dynamic shared memory is only 8-byte aligned by default
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);
+ SmemBuffersT<NUM_STAGES> *smem = reinterpret_cast<SmemBuffersT<NUM_STAGES> *>(aligned_addr);
// Get SMEM addresses for barriers
auto get_mbar_addr = [](void *mbar) -> int
⋯ 11 unchanged lines
// =====================================================================
// TMEM Allocation — ONCE before the tile loop
// =====================================================================
+ // TMEM Allocation (Single 512-column block)
+ // =====================================================================
- __shared__ uint32_t tmem_base_addr[NUM_ACC_BUFS]; // Double-buffered accumulator bases
- __shared__ uint32_t tmem_sfa_addr;
- __shared__ uint32_t tmem_sfb_addr;
- __shared__ uint32_t tmem_idesc;
-
- uint32_t acc_tmem[NUM_ACC_BUFS] = {};
- uint32_t sfa_tmem = 0, sfb_tmem = 0;
+ uint32_t acc_tmem[NUM_ACC_BUFS] = {0, 128};
+ uint32_t sfa_tmem = 256;
+ uint32_t sfb_tmem = 288;
uint32_t idesc = 0;
if (warp_id == MMA_WARP)
{
- int holding_buf_addr = get_mbar_addr(&smem->tmem_holding_buf);
-
- // 1. Allocate double-buffered Accumulators (128 cols each)
- for (int a = 0; a < NUM_ACC_BUFS; a++)
- {
- tcgen05_alloc<1>(holding_buf_addr, TMEM_ACC_COLS);
- tcgen05_wait_alloc();
- if (lane_id == 0)
- acc_tmem[a] = *reinterpret_cast<volatile uint32_t *>(&smem->tmem_holding_buf);
- }
-
- // 2. Allocate SFA (32 cols minimum - 1 bank)
- tcgen05_alloc<1>(holding_buf_addr, TMEM_SFA_COLS);
+ // Allocate single 512-column TMEM block.
+ // The CTA fully owns its cta_group::1 partition, so TMEM implicitly starts at offset 0.
+ // tcgen05.alloc writes the base address to shared memory — we pass a valid smem addr but ignore the result.
+ int alloc_dst = static_cast<int>(__cvta_generic_to_shared(smem));
+ tcgen05_alloc<1>(alloc_dst, TMEM_ALLOC_COLS);
tcgen05_wait_alloc();
- if (lane_id == 0)
- sfa_tmem = *reinterpret_cast<volatile uint32_t *>(&smem->tmem_holding_buf);
-
- // 3. Allocate SFB (32 cols minimum - 1 bank)
- 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);
-
- // Store to shared memory for ALL warps to access
- for (int a = 0; a < NUM_ACC_BUFS; a++)
- tmem_base_addr[a] = acc_tmem[a];
- tmem_sfa_addr = sfa_tmem;
- tmem_sfb_addr = sfb_tmem;
- tmem_idesc = make_mma_idesc(BLOCK_N);
- }
}
-
- // Sync CTA to ensure shared TMEM addresses are visible
- __syncthreads();
-
- // Version marker (disabled for perf)
- // if (tid == 0 && blockIdx.y == 0)
- // printf("KERNEL v713: ACC_BUFS=%d TMEM=%d cols\n", NUM_ACC_BUFS, TMEM_TOTAL_COLS);
-
- // All warps read TMEM addresses from shared memory
+
if (warp_id == MMA_WARP)
{
- for (int a = 0; a < NUM_ACC_BUFS; a++)
- acc_tmem[a] = __shfl_sync(0xFFFFFFFF, tmem_base_addr[a], 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);
+ idesc = make_mma_idesc(BLOCK_N);
}
+ // Sync CTA to ensure TMEM allocation is complete before any warp accesses it
+ __syncthreads();
+
// =========================================================================
// PROLOGUE: Initialize first tile's barriers
// =========================================================================
⋯ 11 unchanged lines
// Init epilogue barriers for both bp indices (tiles 0..1 skip MMA reinit)
for (int a = 0; a < 2; a++)
{
- mbarrier_init(get_mbar_addr(&smem->epilogue_mbar[a]), 1);
+ mbarrier_init(get_mbar_addr(&smem->epilogue_mbar[a]), 1); // 1 MMA warp
mbarrier_init(get_mbar_addr(&smem->epilogue_done_mbar[a]), 4); // 4 epilogue warps
smem->epi_barriers_ready[a] = 1;
}
⋯ 33 unchanged lines
int tile_counter = 0;
- for (uint32_t work_idx = cluster_id; work_idx < kparams.total_tiles; work_idx += total_clusters)
+ // Block distribution: each cluster gets consecutive tiles for B L2 reuse.
+ // With M-fast ordering, consecutive tiles share the same tile_n (same B data).
+ // Cyclic (stride) distribution scatters tiles across N, thrashing L2.
+ const uint32_t tiles_per_cluster = (kparams.total_tiles + total_clusters - 1) / total_clusters;
+ const uint32_t work_start = cluster_id * tiles_per_cluster;
+ const uint32_t work_end = min(work_start + tiles_per_cluster, kparams.total_tiles);
+
+ for (uint32_t work_idx = work_start; work_idx < work_end; work_idx++)
{
int bp = tile_counter & 1; // barrier parity for K-loop SMEM barriers (full/empty_mbar)
// Acc/epilogue always use index 0 (single-buffered TMEM, cta_group::1 = 256 cols max)
⋯ 10 unchanged lines
int my_n_tile = 0;
int total_n_tiles = 0;
bool has_valid_work = false;
+ int epi_coord_m_offset = 0; // offset to subtract from coord_m for C addressing (fused mode)
- // Binary search to find which group this tile belongs to
+ if (kparams.fused)
{
+ // Fused mode: single rectangular tile grid, O(1) group lookup
+ int fused_tiles_m = kparams.fused_tiles_m;
+ tile_n = work_idx / fused_tiles_m;
+ tile_m = work_idx % fused_tiles_m;
+ group_id = kparams.tile_group[tile_m];
+ }
+ else
+ {
+ // Legacy mode: binary search to find which group this tile belongs to
int lo = 0, hi = (int)kparams.num_groups - 1;
while (lo < hi)
{
⋯ 13 unchanged lines
// Load group parameters
params = kparams.groups[group_id];
- // Decode tile indices within group
- uint32_t local_idx = work_idx - params.tile_offset;
- // M-major ordering: M varies fast so consecutive tiles share same B in L2
- tile_n = local_idx / params.tiles_m;
- tile_m = local_idx % params.tiles_m;
+ if (!kparams.fused)
+ {
+ // Legacy: decode tile indices within group
+ uint32_t local_idx = work_idx - params.tile_offset;
+ // M-major ordering: M varies fast so consecutive tiles share same B in L2
+ tile_n = local_idx / params.tiles_m;
+ tile_m = local_idx % params.tiles_m;
+ }
+ else
+ {
+ // Fused: tile_offset stores cumulative M offset for epilogue C addressing
+ epi_coord_m_offset = params.tile_offset * BLOCK_M;
+ }
// K-iteration determination (compile-time if specialized)
if constexpr (K_EXPECTED > 0)
⋯ 23 unchanged lines
if (has_valid_work)
{
- // Get TensorMap pointers for this group from __grid_constant__ array
- 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 (kparams.fused)
+ {
+ // Fused mode: single A/SFA tmap, per-group B/SFB tmaps
+ tmap_A = &kparams.tmaps[FUSED_TMAP_A];
+ tmap_B = &kparams.tmaps[FUSED_TMAP_B_BASE + group_id];
+ tmap_SFA = &kparams.tmaps[FUSED_TMAP_SFA];
+ tmap_SFB = &kparams.tmaps[FUSED_TMAP_SFB_BASE + group_id];
- // Prefetch TensorMaps to warm up TMA path
- if (tid < 4) {
- prefetch_tensormap(&kparams.tmaps[group_id * TENSORMAPS_PER_GROUP + tid]);
+ // Prefetch the 4 tmaps used by this tile
+ if (tid == 0) prefetch_tensormap(&kparams.tmaps[FUSED_TMAP_A]);
+ if (tid == 1) prefetch_tensormap(&kparams.tmaps[FUSED_TMAP_B_BASE + group_id]);
+ if (tid == 2) prefetch_tensormap(&kparams.tmaps[FUSED_TMAP_SFA]);
+ if (tid == 3) prefetch_tensormap(&kparams.tmaps[FUSED_TMAP_SFB_BASE + group_id]);
}
+ else
+ {
+ // Legacy mode: per-group tmaps
+ 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]);
+ }
+ }
+
// Calculate tile coordinates
coord_m = tile_m * BLOCK_M;
coord_n = my_n_tile * BLOCK_N;
⋯ 6 unchanged lines
if (warp_id == MMA_WARP)
{
- if (tile_counter >= NUM_ACC_BUFS)
+ if (tile_counter >= 1)
{
DBG("MMA: wait epi_done\n");
- // Wait for previous tile's epilogue to finish draining TMEM acc[0]
+ // Serialized: always wait for previous tile's epilogue (bp^1).
+ // Concurrent double-buffered would guard >= NUM_ACC_BUFS and wait_bp = bp.
mbarrier_wait(get_mbar_addr(&smem->epilogue_done_mbar[bp ^ 1]), 0);
DBG("MMA: epi_done passed, reinit\n");
- // Reinit epilogue barriers for this tile (MMA owns them now)
+ // Reinit epilogue barriers for this tile's bp (MMA owns them now).
+ // Phase resets to 0 — epilogue always waits with phase 0 after reinit.
if (lane_id == 0)
{
mbarrier_init(get_mbar_addr(&smem->epilogue_mbar[bp]), 1);
⋯ 2 unchanged lines
asm volatile("fence.proxy.async.shared::cta;" ::: "memory");
smem->epi_barriers_ready[bp] = 1;
}
+
+ // CRITICAL: Canonical PTX return handoff fence (epilogue→MMA)
+ tcgen05_fence_after_thread_sync();
+
+ DBG("MMA: reinit done\n");
}
// Wait for TMA to finish K-loop barrier init for this tile's bp
⋯ 96 unchanged lines
if (has_valid_work && warp_id == MMA_WARP)
{
- uint32_t cur_acc = acc_tmem[0]; // Single-buffered accumulator
+ uint32_t cur_acc = acc_tmem[bp % NUM_ACC_BUFS]; // Use double-buffered index
int full_mbar_addr_mma = get_mbar_addr(&smem->full_mbar[bp][stage]);
int empty_mbar_addr = get_mbar_addr(&smem->empty_mbar[bp][stage]);
⋯ 8 unchanged lines
// MMA K-LOOP (Unrolled first iteration)
// =========================================================
- uint32_t sfa_tmem_base = tmem_sfa_addr;
- uint32_t sfb_tmem_base = tmem_sfb_addr;
+ uint32_t sfa_tmem_base = sfa_tmem;
+ uint32_t sfb_tmem_base = sfb_tmem;
// Base A/B descriptors
uint64_t a_desc = make_smem_desc_A(smem->A_smem[stage].data);
⋯ 122 unchanged lines
// =====================================================================
// Check if there is a next tile
- uint32_t next_work_idx = work_idx + total_clusters;
- bool has_next_tile = (next_work_idx < kparams.total_tiles);
+ bool has_next_tile = (work_idx + 1 < work_end);
int next_bp = bp ^ 1;
// -----------------------------------------------------------------
⋯ 42 unchanged lines
}
}
}
+
+ // L2 Prefetch for next tile's A+B+SFA+SFB (fire-and-forget, no sync needed)
+ // Fires while epilogue is running — by the time next K-loop starts, data is warm in L2.
+ if (elect_sync() && kparams.fused && work_idx + 1 < work_end)
+ {
+ int next_group_id = kparams.next_tile_group[tile_m]; // precomputed by wrapper
+ int next_n_increment = kparams.next_tile_n[tile_m]; // 0=same N, 1=next N
+ int next_my_n_tile = (tile_n + next_n_increment) * CLUSTER_N + cta_n;
+ int next_tile_m = (work_idx + 1) % kparams.fused_tiles_m;
+ int next_coord_n = next_my_n_tile * BLOCK_N;
+ int next_coord_m = next_tile_m * BLOCK_M;
+ int next_k_iters = kparams.groups[next_group_id].K / BLOCK_K;
+
+ const void* next_tmap_B = &kparams.tmaps[FUSED_TMAP_B_BASE + next_group_id];
+ const void* next_tmap_SFB = &kparams.tmaps[FUSED_TMAP_SFB_BASE + next_group_id];
+ const void* next_tmap_A = &kparams.tmaps[FUSED_TMAP_A];
+ const void* next_tmap_SFA = &kparams.tmaps[FUSED_TMAP_SFA];
+
+ // Prefetch B+SFB: each CTA prefetches its own N-slice (all ranks)
+ // Prefetch up to NUM_STAGES B stages (covers full pipeline depth)
+ #pragma unroll
+ for (int k_pf = 0; k_pf < NUM_STAGES && k_pf < next_k_iters; k_pf++)
+ {
+ tma_2d_prefetch(next_tmap_B, k_pf * BLOCK_K, next_coord_n, EVICT_FIRST);
+ int off_sfb = (next_my_n_tile * next_k_iters + k_pf) * SMEM_SFB_SIZE;
+ tma_1d_prefetch(next_tmap_SFB, off_sfb / 8, EVICT_FIRST);
+ }
+
+ // Prefetch A+SFA: multicast-gated to rank-0 only (avoids duplicate HBM reads)
+ // Non-multicast: all ranks prefetch their own (but A is same so 1 rank is enough)
+ // Prefetch NUM_STAGES stages of A: L2 is large enough (96KB << 4MB per cluster)
+ // and A uses EVICT_LAST so it won't displace B tiles.
+ if (!use_multicast || cta_n == 0)
+ {
+ #pragma unroll
+ for (int k_pf = 0; k_pf < NUM_STAGES && k_pf < next_k_iters; k_pf++)
+ {
+ tma_2d_prefetch(next_tmap_A, k_pf * BLOCK_K, next_coord_m, EVICT_LAST);
+ int off_sfa = (next_tile_m * next_k_iters + k_pf) * SMEM_SFA_SIZE;
+ tma_1d_prefetch(next_tmap_SFA, off_sfa / 8, EVICT_LAST);
+ }
+ }
+ }
}
// ALL threads: cluster_sync ensures all CTAs see reinitialized barriers
⋯ 14 unchanged lines
// -----------------------------------------------------------------
// EPILOGUE WARPS (0-3): Drain TMEM to global memory
// -----------------------------------------------------------------
+
if (warp_id < 4)
{
if (has_valid_work)
{
- // Spin-check that barriers are ready
+ // Spin-check that MMA has reinit'd the barriers for this bp.
+ // Necessary because epilogue warps can reach here before MMA's lane 0
+ // has completed mbarrier_init + epi_barriers_ready write.
while (smem->epi_barriers_ready[bp] == 0) {}
if (warp_id == 0) DBG("EPI: wait epi\n");
- // Wait for MMA to signal accumulator is ready
+ // Wait for MMA to signal accumulator is ready.
+ // Phase 0: barrier is always reinit'd (parity reset) before this wait.
mbarrier_wait(get_mbar_addr(&smem->epilogue_mbar[bp]), 0);
if (warp_id == 0) DBG("EPI: draining\n");
// CRITICAL: Fence required between MMA/TMA and tcgen05_ld
tcgen05_fence_after_thread_sync();
- half *C_ptr = reinterpret_cast<half *>(params.C_ptr);
- int M = params.M;
- int N = params.N;
-
// Distribute work: 1 tile per warp
// Warps 0, 1, 2, 3 handle 32 rows each -> 128 rows total
int row_tile = warp_id;
int base_row = row_tile * 32;
- // Phase 1: Drain ALL TMEM to registers (fast, ~300 cycles)
- // Signal epilogue_done immediately after — MMA can start next tile
- // while Phase 2 stores overlap in background.
- constexpr int NUM_CHUNKS = BLOCK_N / 8;
- float tile_data[NUM_CHUNKS][8];
+ // Phase 1: Drain and store in CHUNKS to reduce register pressure
+ // BLOCK_N=128, handle in 4 chunks of 32 columns each.
+ // This reduces tile_data from 64 registers/thread to 32.
+ constexpr int CHUNK_SIZE = 32;
+ constexpr int NUM_CHUNKS_PER_ITER = CHUNK_SIZE / 8;
+ constexpr int NUM_ITERATIONS = BLOCK_N / CHUNK_SIZE;
+ float tile_data[NUM_CHUNKS_PER_ITER][8];
+ half *C_ptr_dst = reinterpret_cast<half *>(params.C_ptr);
+ int local_row = base_row + lane_id;
+ int c_row = coord_m - epi_coord_m_offset + local_row;
+
#pragma unroll
- for (int c_chunk = 0; c_chunk < NUM_CHUNKS; c_chunk++)
+ for (int chunk_idx = 0; chunk_idx < NUM_ITERATIONS; chunk_idx++)
{
- tcgen05_ld_32x32b<8>(tile_data[c_chunk], tmem_base_addr[0], base_row, c_chunk * 8);
- tcgen05_wait_alloc();
+ int chunk_col_base = chunk_idx * CHUNK_SIZE;
+
+ // Load chunk from TMEM to registers
+ #pragma unroll
+ for (int c_chunk = 0; c_chunk < NUM_CHUNKS_PER_ITER; c_chunk++)
+ {
+ tcgen05_ld_32x32b<8>(tile_data[c_chunk], acc_tmem[bp % NUM_ACC_BUFS], base_row, chunk_col_base + c_chunk * 8);
+ }
+ tcgen05_wait_alloc(); // wait for this chunk's load
+
+ // Store chunk to GMEM
+ if (c_row < params.M)
+ {
+ #pragma unroll
+ for (int c_chunk = 0; c_chunk < NUM_CHUNKS_PER_ITER; c_chunk++)
+ {
+ int base_col = chunk_col_base + c_chunk * 8;
+ // int4 is 16-byte aligned by type; half2[4] is only 4-byte aligned
+ // and can spill to local memory causing misaligned address fault.
+ int4 out;
+ #pragma unroll
+ for (int i = 0; i < 4; i++)
+ {
+ reinterpret_cast<half2 *>(&out)[i] = __float22half2_rn({tile_data[c_chunk][i * 2], tile_data[c_chunk][i * 2 + 1]});
+ }
+ reinterpret_cast<int4 *>(C_ptr_dst + c_row * params.N + coord_n + base_col)[0] = out;
+ }
+ }
}
+ // Canonical PTX return handoff fence: guarantee ld visibility before MMA reuse
+ tcgen05_wait_alloc();
+ tcgen05_fence_before_thread_sync();
+
// Signal that this warp's TMEM drain is complete
- // All 4 epilogue warps must arrive (init count=4) before MMA can reuse TMEM
if (lane_id == 0)
{
if (warp_id == 0) DBG("EPI: done\n");
mbarrier_arrive(get_mbar_addr(&smem->epilogue_done_mbar[bp]));
}
+ // Clear the ready flag so this bp can be safely reinit'd next time
smem->epi_barriers_ready[bp] = 0;
-
- // Phase 2: Convert and store (overlapped with next tile's MMA)
- int local_row = base_row + lane_id;
- int global_row = coord_m + local_row;
-
- if (global_row < M)
- {
- #pragma unroll
- for (int c_chunk = 0; c_chunk < NUM_CHUNKS; c_chunk++)
- {
- int base_col = c_chunk * 8;
- half2 h2_tmp[4];
- #pragma unroll
- for (int i = 0; i < 4; i++)
- {
- h2_tmp[i] = __float22half2_rn({tile_data[c_chunk][i * 2], tile_data[c_chunk][i * 2 + 1]});
- }
- reinterpret_cast<int4 *>(C_ptr + global_row * N + coord_n + base_col)[0] =
- *reinterpret_cast<int4 *>(h2_tmp);
- }
- }
}
else
{
⋯ 13 unchanged lines
// LAST TILE CLEANUP: Wait for final epilogue to complete
// =========================================================================
- // MMA warp must wait for last epilogue_done before deallocating TMEM
+ // MMA warp must wait for last epilogue_done before deallocating TMEM.
+ // Phase 0: last tile's epi_done was reinit'd by MMA guard (parity reset to 0),
+ // then epilogue arrives 4x → parity 1. Wait with phase 0 → passes.
+ // Special case tile 0 (no reinit): prologue init'd it to parity 0, epilogue → 1.
if (warp_id == MMA_WARP && tile_counter > 0)
{
- mbarrier_wait(get_mbar_addr(&smem->epilogue_done_mbar[(tile_counter - 1) & 1]), 0);
+ int wait_bp = (tile_counter - 1) & 1;
+ mbarrier_wait(get_mbar_addr(&smem->epilogue_done_mbar[wait_bp]), 0);
}
// Full sync before TMEM deallocation
⋯ 6 unchanged lines
if (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);
- for (int a = NUM_ACC_BUFS - 1; a >= 0; a--)
- tcgen05_dealloc<1>(acc_tmem[a], TMEM_ACC_COLS);
-
+ tcgen05_dealloc<1>(0, TMEM_ALLOC_COLS);
tcgen05_wait_alloc();
}
⋯ 11 unchanged lines
// Format: <CLUSTER_N, NUM_STAGES, K_EXPECTED, BLOCK_N>
// K=0 fallback only — K-specialized templates cause icache pressure regression.
+ template __global__ void group_gemm_kernel_impl<1, 4, 0, 128>(const __grid_constant__ KernelParams kparams);
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, 4, 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, 4, 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);
⋯ 33 unchanged lines
// Double-buffered barriers: full_mbar[2][N] + empty_mbar[2][N] + epilogue_mbar[2]
// + epilogue_done_mbar[2] + tmem_holding_buf + epi_barriers_ready[2]
constexpr int SMEM_BARRIER_SIZE = sizeof(uint64_t) * (NUM_STAGES_MAX * 4 + 4 + 1) + sizeof(int) * 2;
- constexpr int SMEM_SIZE = SMEM_DATA_SIZE + SMEM_BARRIER_SIZE + 128;
+ constexpr int SMEM_SIZE = SMEM_DATA_SIZE + SMEM_BARRIER_SIZE + 128; // Max (NUM_STAGES_MAX)
+ // Dynamic SMEM size computation for variable pipeline depth
+ inline int compute_smem_size(int num_stages) {
+ int data = num_stages * (SMEM_A_ALIGNED + SMEM_B_ALIGNED + SMEM_SFA_ALIGNED + SMEM_SFB_ALIGNED);
+ int barriers = sizeof(uint64_t) * (num_stages * 4 + 4 + 1) + sizeof(int) * 2;
+ return data + barriers + 128; // +128 for alignment padding
+ }
+
// =============================================================================
// CUDA Driver API Error Checking
// =============================================================================
⋯ 55 unchanged lines
elementStrides,
CU_TENSOR_MAP_INTERLEAVE_NONE,
CU_TENSOR_MAP_SWIZZLE_128B,
- CU_TENSOR_MAP_L2_PROMOTION_NONE,
+ CU_TENSOR_MAP_L2_PROMOTION_L2_64B, // A is small, reused across N-tiles
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
check_cu(err, "cuTensorMapEncodeTiled for A (16U4, rank-2)");
}
⋯ 34 unchanged lines
elementStrides,
CU_TENSOR_MAP_INTERLEAVE_NONE,
CU_TENSOR_MAP_SWIZZLE_128B,
- CU_TENSOR_MAP_L2_PROMOTION_L2_64B, // Promote B to L2 for reuse across M-tiles
+ CU_TENSOR_MAP_L2_PROMOTION_L2_64B, // B reused across M-tiles with block distribution
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
check_cu(err, "cuTensorMapEncodeTiled for B (16U4, rank-2)");
}
⋯ 65 unchanged lines
elementStrides,
CU_TENSOR_MAP_INTERLEAVE_NONE,
CU_TENSOR_MAP_SWIZZLE_NONE,
- CU_TENSOR_MAP_L2_PROMOTION_NONE,
+ CU_TENSOR_MAP_L2_PROMOTION_L2_64B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
check_cu(err, "cuTensorMapEncodeTiled for SF (1D)");
}
⋯ 28 unchanged lines
if (K < min_K) min_K = K;
}
- // CLUSTER_N=4 needs 3 remote arrives/k_iter — only worthwhile for high-K
- // CLUSTER_N=2 needs 1 remote arrive/k_iter — moderate overhead
- int max_cluster = (min_K >= 4096) ? 4 : (min_K >= 2048) ? 2 : 1;
+ // CLUSTER_N=4 with short K loops (K=2048, 8 k_iters) causes ~80% regression:
+ // empty_mbar init_count=4 forces TMA to wait for all 4 CTAs, and synchronization
+ // overhead on short pipelines dominates over A bandwidth savings. Tested: BM2 41→75µs.
+ // CLUSTER_N=2 for BM4 (K=1536, 6 k_iters): neutral (10.6→10.7µs), keep it.
+ int max_cluster = (min_K >= 4096) ? 4 : (min_K >= 1024) ? 2 : 1;
if (max_cluster >= 4 && all_valid(4))
return 4;
⋯ 7 unchanged lines
int num_groups = problem_sizes.size(0);
auto sizes_acc = problem_sizes.accessor<int32_t, 2>();
+ // Select pipeline depth based on arithmetic intensity.
+ // Higher AI (compute-bound) → more stages to hide latency.
+ // Lower AI (memory-bound) → fewer stages (diminishing returns).
+ // TODO: Enable stages 3-4 for 2-CTA occupancy after kernel validation.
float min_ai = FLT_MAX;
for (int g = 0; g < num_groups; g++)
{
⋯ 10 unchanged lines
}
DEBUG_PRINT("[PIPELINE] min_ai=%.2f\n", min_ai);
- return (min_ai > 150.0f) ? 6 : 5;
+ if (min_ai > 150.0f) return 6;
+ if (min_ai > 100.0f) return 5;
+ return 4;
}
// =============================================================================
⋯ 26 unchanged lines
// Determine optimal pipeline depth based on arithmetic intensity
int num_stages = compute_num_stages(problem_sizes);
- DEBUG_PRINT("[PIPELINE] selected num_stages=%d\n", num_stages);
+ // Use dynamic SMEM size for actual num_stages.
+ int smem_size = compute_smem_size(num_stages);
+ DEBUG_PRINT("[PIPELINE] selected num_stages=%d, smem=%d bytes\n", num_stages, smem_size);
// Access data on CPU
auto abc_acc = abc_ptrs.accessor<int64_t, 2>();
⋯ 31 unchanged lines
// Build KernelParams from scratch
memset(&kparams, 0, sizeof(kparams));
- uint32_t total_tiles = 0;
+ // Check if all groups share same N, K, L (eligible for M-fusion)
+ bool can_fuse = (num_groups > 1);
+ if (can_fuse) {
+ int N0 = sizes_acc[0][1], K0 = sizes_acc[0][2], L0 = sizes_acc[0][3];
+ for (int g = 1; g < num_groups; g++) {
+ if (sizes_acc[g][1] != N0 || sizes_acc[g][2] != K0 || sizes_acc[g][3] != L0) {
+ can_fuse = false;
+ break;
+ }
+ }
+ }
- 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];
+ if (can_fuse) {
+ // =========================================================
+ // Fused mode: single A/SFA TensorMap, per-group B/SFB
+ // =========================================================
+ int N = sizes_acc[0][1];
+ int K = sizes_acc[0][2];
+ int L = sizes_acc[0][3];
- int tiles_m = (M + BLOCK_M - 1) / BLOCK_M;
+ // Python passes fused data as num_groups entries where:
+ // - All groups share the same A_ptr (fused A) and SFA_ptr (fused SFA)
+ // - Each group has its own B_ptr, SFB_ptr, C_ptr
+ // - sizes[g] = [M_g (actual), N, K, L] (per-group actual M)
+ // M_fused = sum of ceil(M_g/BLOCK_M)*BLOCK_M across all groups
+
+ // Compute M_fused from per-group Ms (each padded to BLOCK_M)
+ int M_fused = 0;
+ for (int g = 0; g < num_groups; g++) {
+ M_fused += ((sizes_acc[g][0] + BLOCK_M - 1) / BLOCK_M) * BLOCK_M;
+ }
+ void *A_fused_ptr = reinterpret_cast<void *>(abc_acc[0][0]);
+ void *SFA_fused_ptr = reinterpret_cast<void *>(sf_acc[0][0]);
+
+ int tiles_m_total = (M_fused + 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;
+ TORCH_CHECK(tiles_m_total <= MAX_M_TILES,
+ "Too many M-tiles for fused mode: ", tiles_m_total);
- 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]);
+ kparams.fused = 1;
+ kparams.fused_tiles_m = static_cast<uint16_t>(tiles_m_total);
+ kparams.fused_tiles_n = static_cast<uint16_t>(std::max(1, tiles_n));
+ kparams.total_tiles = tiles_m_total * std::max(1, tiles_n);
+ kparams.num_groups = static_cast<uint32_t>(num_groups);
- total_tiles += group_tiles;
- }
- kparams.total_tiles = total_tiles;
- kparams.num_groups = static_cast<uint32_t>(num_groups);
+ // Build tile_group[] mapping and per-group params
+ int tile_m_cursor = 0;
+ for (int g = 0; g < num_groups; g++) {
+ int M_g = sizes_acc[g][0];
+ int group_tiles_m = (M_g + BLOCK_M - 1) / BLOCK_M;
- // Fill TensorMaps
- 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];
+ // In fused mode, tile_offset stores the cumulative M-tile offset
+ // (used in epilogue to convert fused coord_m to per-group C row)
+ kparams.groups[g].tile_offset = tile_m_cursor;
+ kparams.groups[g].M = static_cast<uint16_t>(M_g);
+ kparams.groups[g].N = static_cast<uint16_t>(N);
+ kparams.groups[g].K = static_cast<uint16_t>(K);
+ kparams.groups[g].L = static_cast<uint16_t>(L);
+ kparams.groups[g].tiles_m = static_cast<uint16_t>(group_tiles_m);
+ kparams.groups[g].tiles_n = static_cast<uint16_t>(std::max(1, tiles_n));
+ kparams.groups[g].padding = 0;
- 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]);
+ // C_ptr is per-group for epilogue writeback
+ kparams.groups[g].C_ptr = reinterpret_cast<void *>(abc_acc[g][2]);
+ // A/B/SF ptrs stored for reference but tmaps are separate
+ kparams.groups[g].A_ptr = A_fused_ptr;
+ kparams.groups[g].B_ptr = reinterpret_cast<void *>(abc_acc[g][1]);
+ kparams.groups[g].SFA_ptr = SFA_fused_ptr;
+ kparams.groups[g].SFB_ptr = reinterpret_cast<void *>(sf_acc[g][1]);
- int base_idx = g * TENSORMAPS_PER_GROUP;
+ for (int t = 0; t < group_tiles_m; t++) {
+ kparams.tile_group[tile_m_cursor + t] = static_cast<uint8_t>(g);
+ }
+ tile_m_cursor += group_tiles_m;
+ }
- int num_m_tiles_a = (M + BLOCK_M - 1) / BLOCK_M;
- int num_m_tiles_b = (N + BLOCK_N - 1) / BLOCK_N;
+ // Populate prefetch helpers: next tile's group and n-increment
+ for (int t = 0; t < tiles_m_total; t++) {
+ int next_t = (t + 1);
+ if (next_t < tiles_m_total) {
+ kparams.next_tile_group[t] = kparams.tile_group[next_t];
+ kparams.next_tile_n[t] = 0; // stays in same N-tile column
+ } else {
+ kparams.next_tile_group[t] = kparams.tile_group[0];
+ kparams.next_tile_n[t] = 1; // wraps to next N-tile column
+ }
+ }
- int M_padded = num_m_tiles_a * BLOCK_M;
- int N_padded = num_m_tiles_b * BLOCK_N;
+ // Fused A TensorMap (covers entire fused M_total)
+ int M_fused_padded = tiles_m_total * BLOCK_M;
+ init_A_tmap(&kparams.tmaps[FUSED_TMAP_A], A_fused_ptr, M_fused, K, L, BLOCK_M, BLOCK_K);
+ init_SF_tmap(&kparams.tmaps[FUSED_TMAP_SFA], SFA_fused_ptr, M_fused_padded, K);
- 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);
+ // Per-group B/SFB TensorMaps
+ for (int g = 0; g < num_groups; g++) {
+ void *B_ptr = reinterpret_cast<void *>(abc_acc[g][1]);
+ void *SFB_ptr = reinterpret_cast<void *>(sf_acc[g][1]);
+ int N_g = N;
+ int N_padded = ((N_g + BLOCK_N - 1) / BLOCK_N) * BLOCK_N;
+
+ init_B_tmap(&kparams.tmaps[FUSED_TMAP_B_BASE + g], B_ptr, N_g, K, L, BLOCK_N, BLOCK_K);
+ init_SF_tmap(&kparams.tmaps[FUSED_TMAP_SFB_BASE + g], SFB_ptr, N_padded, K);
+ }
+
+ DEBUG_PRINT("[FUSED] M_fused=%d, tiles_m=%d, tiles_n=%d, total_tiles=%u\n",
+ M_fused, tiles_m_total, tiles_n, kparams.total_tiles);
+ } else {
+ // =========================================================
+ // Legacy mode: per-group TensorMaps
+ // =========================================================
+ 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);
+
+ // Fill TensorMaps
+ 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;
+
+ // A/B use OOB_FILL_ZERO: pass actual M,N — TMA zero-fills partial tiles
+ init_A_tmap(&kparams.tmaps[base_idx + TMAP_A_IDX], A_ptr, M, K, L, BLOCK_M, BLOCK_K);
+ init_B_tmap(&kparams.tmaps[base_idx + TMAP_B_IDX], B_ptr, N, 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);
+ }
}
// Store to cache
⋯ 34 unchanged lines
DEBUG_PRINT("[LAUNCH] num_groups=%d, total_tiles=%u, grid_y=%u\n",
num_groups, kparams.total_tiles, grid_y);
- DEBUG_PRINT("[LAUNCH] SMEM_SIZE=%d bytes, THREADS_PER_CTA=%d\n", SMEM_SIZE, THREADS_PER_CTA);
+ DEBUG_PRINT("[LAUNCH] smem_size=%d bytes, THREADS_PER_CTA=%d\n", smem_size, THREADS_PER_CTA);
DEBUG_PRINT("[LAUNCH] Cluster dims: (%d, %d, %d)\n", CLUSTER_M, cluster_n, CLUSTER_Z);
// Select kernel function based on cluster_n and num_stages
⋯ 2 unchanged lines
{
if (num_stages >= 6)
kernel_func = (const void *)group_gemm_kernel_impl<4, 6, 0, 128>;
- else
+ else if (num_stages >= 5)
kernel_func = (const void *)group_gemm_kernel_impl<4, 5, 0, 128>;
+ else
+ kernel_func = (const void *)group_gemm_kernel_impl<4, 4, 0, 128>;
}
else if (cluster_n == 2)
{
if (num_stages >= 6)
kernel_func = (const void *)group_gemm_kernel_impl<2, 6, 0, 128>;
- else
+ else if (num_stages >= 5)
kernel_func = (const void *)group_gemm_kernel_impl<2, 5, 0, 128>;
+ else
+ kernel_func = (const void *)group_gemm_kernel_impl<2, 4, 0, 128>;
}
else
{
if (num_stages >= 6)
kernel_func = (const void *)group_gemm_kernel_impl<1, 6, 0, 128>;
- else
+ else if (num_stages >= 5)
kernel_func = (const void *)group_gemm_kernel_impl<1, 5, 0, 128>;
+ else
+ kernel_func = (const void *)group_gemm_kernel_impl<1, 4, 0, 128>;
}
// Set maximum dynamic shared memory size (cached per kernel variant)
⋯ 3 unchanged lines
check_cuda(cudaFuncSetAttribute(
kernel_func,
cudaFuncAttributeMaxDynamicSharedMemorySize,
- SMEM_SIZE),
+ smem_size),
"cudaFuncSetAttribute for shared memory");
configured_kernel = kernel_func;
}
⋯ 6 unchanged lines
cudaLaunchConfig_t config = {};
config.gridDim = grid;
config.blockDim = block;
- config.dynamicSmemBytes = SMEM_SIZE;
+ config.dynamicSmemBytes = smem_size;
cudaLaunchAttribute attrs[1];
attrs[0].id = cudaLaunchAttributeClusterDimension;
⋯ 117 unchanged lines
return tensor * sf_expanded
- def prepare_sf_for_tma(sf_raw, mn, k, block_size=128):
+
+ # Cache for prepared launch data (avoids re-preparing on repeated calls)
+ _launch_cache = dict()
+
+ def _merge_groups(abc_tensors, sfasfb_reordered, problem_sizes):
"""
- Prepare blocked scale factor tensor for tcgen05.cp.32x128b.warpx4 with SBO=128.
+ M-fusion: when all groups share the same N, K, L, fuse them into a single
+ GEMM by concatenating A/SFA along M (padded to 128 per group).
+ The kernel uses per-tile B selection via tile_group[] mapping.
- Input: sfasfb_reordered [32, 4, rest_m, 4, rest_k, L] (cuBLAS blocked format)
- Output: [L * rest_m * rest_k * 32, 16] for TMA loading
+ Falls back to per-group pass-through when groups have different N, K, or L.
- Vectorized: replaces triple-nested Python loop with single permute+reshape.
- The loop extracted sf_u8[:,:,mm,:,kk,L_idx] → [32,4,4], reshaped to [32,16].
- This is equivalent to permuting dims (5,2,4,0,1,3) → [L,rest_m,rest_k,32,4,4]
- then reshaping to [-1, 16].
+ Returns (merged_abc_ptrs, merged_sf_ptrs, merged_sizes, refs, c_copyback).
"""
- BLOCK_M = 128
+ num_groups = len(problem_sizes)
- sf_u8 = sf_raw.view(torch.uint8).clone()
+ # Check if all groups share same N, K, L (eligible for M-fusion)
+ N0, K0, L0 = problem_sizes[0][1], problem_sizes[0][2], problem_sizes[0][3]
+ can_fuse = num_groups > 1 and all(
+ ps[1] == N0 and ps[2] == K0 and ps[3] == L0
+ for ps in problem_sizes[1:]
+ )
- if len(sf_raw.shape) == 6:
- # Shape: [mm32=32, mm4=4, rest_m, kk4=4, rest_k, L]
- rest_m = sf_raw.shape[2]
+ if can_fuse:
+ # M-fusion: concatenate A/SFA along M, keep per-group B/SFB/C
+ # cuBLAS SF reorder format already pads M to 128 (rest_m = ceil(M/128)),
+ # so each group's SFA is already correctly sized for TMA.
+ # A needs explicit padding to BLOCK_M=128 boundary.
+ BLOCK_M = 128
+ a_parts = []
+ sfa_parts = []
- # Zero out OOB M positions for partial last tile (vectorized)
- 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
+ for idx in range(num_groups):
+ a_i = abc_tensors[idx][0]
+ M_i = problem_sizes[idx][0]
+ pad_m = ((M_i + BLOCK_M - 1) // BLOCK_M) * BLOCK_M
- # Permute [mm32, mm4, rest_m, kk4, rest_k, L] → [L, rest_m, rest_k, mm32, mm4, kk4]
- # Then reshape: merge (mm4, kk4) → 16 bytes per row, flatten outer dims
- result = sf_u8.permute(5, 2, 4, 0, 1, 3).contiguous()
- return result.reshape(-1, 16)
+ # Pad A to BLOCK_M boundary. FP4 packed dtype doesn't support fill_,
+ # so allocate as uint8 (same byte layout), zero-fill, then view as FP4.
+ if pad_m > M_i:
+ a_padded = torch.zeros(pad_m, a_i.shape[1], *a_i.shape[2:],
+ dtype=torch.uint8, device=a_i.device).view(a_i.dtype)
+ a_padded[:M_i] = a_i
+ else:
+ a_padded = a_i
- 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()
+ # SFA: cuBLAS reorder already has ceil(M/128)*128 worth of SF rows
+ sfa_tma_i = sfasfb_reordered[idx][0].view(torch.uint8).permute(5, 2, 4, 0, 1, 3).reshape(-1, 16)
+ a_parts.append(a_padded)
+ sfa_parts.append(sfa_tma_i)
- def pad_tensor_to_block(tensor, dim, block_size):
- """
- Pad tensor along specified dimension to be a multiple of 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
+ # Concatenate A and SFA along M dimension
+ a_fused = torch.cat(a_parts, dim=0)
+ sfa_fused = torch.cat(sfa_parts, dim=0)
- # Create padding shape
- pad_shape = list(tensor.shape)
- pad_shape[dim] = target_size - current_size
-
- # Workaround: "fill_cuda" not implemented for Float4_e2m1fn_x2
- # Create as uint8 (underlying storage) and view as target dtype
- padding = torch.zeros(pad_shape, dtype=torch.uint8, device=tensor.device).view(tensor.dtype)
- return torch.cat([tensor, padding], dim=dim)
+ # Per-group SFB preparation
+ sfb_tmas = []
+ for idx in range(num_groups):
+ sfb_tma = sfasfb_reordered[idx][1].view(torch.uint8).permute(5, 2, 4, 0, 1, 3).reshape(-1, 16)
+ sfb_tmas.append(sfb_tma)
+ # Build per-group abc_ptrs/sf_ptrs/sizes arrays
+ # All groups share fused A_ptr and SFA_ptr; each has own B, SFB, C
+ abc_ptrs = []
+ sf_ptrs = []
+ sizes = []
+ c_copyback = []
- # Cache for prepared launch data (avoids re-preparing on repeated calls)
- _launch_cache = dict()
+ for idx in range(num_groups):
+ M_i = problem_sizes[idx][0]
+ abc_ptrs.append([a_fused.data_ptr(), abc_tensors[idx][1].data_ptr(),
+ abc_tensors[idx][2].data_ptr()])
+ sf_ptrs.append([sfa_fused.data_ptr(), sfb_tmas[idx].data_ptr()])
+ sizes.append([M_i, N0, K0, L0])
+ # Keep references alive
+ refs = [a_fused, sfa_fused] + sfb_tmas + [a_parts, sfa_parts]
+
+ abc_ptrs_t = torch.tensor(abc_ptrs, dtype=torch.int64)
+ sf_ptrs_t = torch.tensor(sf_ptrs, dtype=torch.int64)
+ sizes_t = torch.tensor(sizes, dtype=torch.int32)
+ return abc_ptrs_t, sf_ptrs_t, sizes_t, refs, c_copyback
+
+ # Non-fusable: pass through each group individually
+ abc_ptrs = []
+ sf_ptrs = []
+ sizes = []
+ refs = []
+ c_copyback = []
+
+ for idx in range(num_groups):
+ a, b, c = abc_tensors[idx]
+ sfa_r = sfasfb_reordered[idx][0]
+ sfb_r = sfasfb_reordered[idx][1]
+ M = problem_sizes[idx][0]
+ N, K, L = problem_sizes[idx][1], problem_sizes[idx][2], problem_sizes[idx][3]
+
+ abc_ptrs.append([a.data_ptr(), b.data_ptr(), c.data_ptr()])
+ sfa_tma = sfa_r.view(torch.uint8).permute(5, 2, 4, 0, 1, 3).reshape(-1, 16)
+ sfb_tma = sfb_r.view(torch.uint8).permute(5, 2, 4, 0, 1, 3).reshape(-1, 16)
+ refs.append((sfa_r, sfb_r, sfa_tma, sfb_tma))
+ sf_ptrs.append([sfa_tma.data_ptr(), sfb_tma.data_ptr()])
+ sizes.append([M, N, K, L])
+
+ abc_ptrs_t = torch.tensor(abc_ptrs, dtype=torch.int64)
+ sf_ptrs_t = torch.tensor(sf_ptrs, dtype=torch.int64)
+ sizes_t = torch.tensor(sizes, dtype=torch.int32)
+ return abc_ptrs_t, sf_ptrs_t, sizes_t, refs, c_copyback
+
+
+ def _copyback_c(c_copyback):
+ """Copy merged C slices back to original C tensors."""
+ for c_merged, slices in c_copyback:
+ for c_orig, start, end in slices:
+ c_orig.copy_(c_merged[start:end])
+
+
def custom_kernel_cuda(data):
"""
- Group GEMM kernel using optimized CUDA implementation.
+ Group GEMM kernel with opportunistic M-fusion for shared-B groups.
Caches prepared tensors to avoid redundant work on repeated calls.
"""
abc_tensors, sfasfb_tensors, sfasfb_reordered, problem_sizes = data
⋯ 8 unchanged lines
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)
+ abc_ptrs_t, sf_ptrs_t, sizes_t, _refs, c_copyback = cached
+ torch.ops.group_gemm_module.group_gemm_launch(abc_ptrs_t, sf_ptrs_t, sizes_t)
+ _copyback_c(c_copyback)
return [t[2] for t in abc_tensors]
- # Cache miss: prepare everything
- abc_ptrs = []
- sf_ptrs = []
- sizes = []
- padded_tensors = []
- sf_prepared = []
+ # Cache miss: merge groups sharing B, then prepare
+ abc_ptrs_t, sf_ptrs_t, sizes_t, refs, c_copyback = _merge_groups(
+ abc_tensors, sfasfb_reordered, problem_sizes)
- BLOCK_M = 128
- BLOCK_N = 128
+ _launch_cache[cache_key] = (abc_ptrs_t, sf_ptrs_t, sizes_t, refs, c_copyback)
- 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)
-
- # Cache for future calls (keep refs alive to prevent GC)
- _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)
+ torch.ops.group_gemm_module.group_gemm_launch(abc_ptrs_t, sf_ptrs_t, sizes_t)
+ _copyback_c(c_copyback)
return [t[2] for t in abc_tensors]
⋯ diff truncated
scrolls · 1201 diff lines total

Best evidence level for this revision: reported

JSON