Skip to content
KernelIndex
Search⌘K

submission 487344

macto · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sub_v4.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-487344?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
16.9µs
#45 of 310
2026-02-09

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:90018c450157f068379d12833ea1220d70f8f4be69d1d293fc85942a2bdb8ebb
license declaredunknown
license concludedunknown
authorsmacto
imported2026-08-15

Techniques

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

mbarrier__device__ __forceinline__ void mbarrier_init(int mbar_addr, int count) {
persistent-kernelnamespace persistent {
shared-memoryextern __shared__ __align__(1024) char smem_ptr[];
stages = 6constexpr int NUM_STAGES = 6;
tcgen05asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc), "n"(CTA_GROUP));
tile-k = 256constexpr int BLOCK_K = 256;
tile-m = 128constexpr int BLOCK_M = 128;
tile-n = 128constexpr int BLOCK_N = 128;
tmaasm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;"
vector-width = half2half2 val = __floats2half2_rn(f0, f1);

Kernel source

sub_v4.py1993 lines
from __future__ import annotations

from typing import Dict, List, Tuple

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t


CPP_SRC = r"""
#include <torch/extension.h>
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {}
"""

CUDA_SRC = r"""
#include <torch/types.h>
#include <cuda.h>
#include <cuda_runtime.h>

#include <torch/types.h>
#include <cuda.h>
#include <cuda_runtime.h>

#include <cudaTypedefs.h>
#include <cuda_fp16.h>

#include <stddef.h>
#include <stdint.h>
#include <torch/library.h>

// --------------------------
// Common helpers
// --------------------------

__device__ inline constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; };

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

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

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

__device__ __forceinline__ void mbarrier_wait(int mbar_addr, int phase) {
  uint32_t ticks = 0x989680;
  asm volatile(
    "{\n\t"
    ".reg .pred P1;\n\t"
    "LAB_WAIT:\n\t"
    "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t"
    "@!P1 bra.uni LAB_WAIT;\n\t"
    "}"
    :: "r"(mbar_addr), "r"(phase), "r"(ticks)
  );
}

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

__device__ __forceinline__ 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::1.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)
              : "memory");
}

template <int CTA_GROUP = 1>
__device__ __forceinline__ void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
  asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc), "n"(CTA_GROUP));
}

template <int CTA_GROUP = 1>
__device__ __forceinline__ void tcgen05_mma_nvfp4(
  int d_tmem,
  uint64_t a_desc,
  uint64_t b_desc,
  uint32_t i_desc,
  int scale_A_tmem,
  int scale_B_tmem,
  int enable_input_d
) {
  asm volatile(
    "{\n\t"
    ".reg .pred p;\n\t"
    "setp.ne.b32 p, %6, 0;\n\t"
    "tcgen05.mma.cta_group::%7.kind::mxf4nvf4.block_scale.block16 [%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)
  );
}

struct SHAPE { static constexpr char _16x256b[] = ".16x256b"; };

template <int NUM_REGS, const char *SHAPE_, int NUM>
__device__ __forceinline__ void tcgen05_ld(float *tmp, int row, int col) {
  const int addr = (row << 16) | col;
  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));
  }
}

__device__ __forceinline__ void tcgen05_ld_16x256bx8(float *tmp, int row, int col) {
  tcgen05_ld<32, SHAPE::_16x256b, 8>(tmp, row, col);
}

template <int num>
__device__ __forceinline__ void tcgen05_ld_16x256b(float *tmp, int row, int col) {
  tcgen05_ld<num * 4, SHAPE::_16x256b, num>(tmp, row, col);
}

__device__ __forceinline__ void store_cs_half2(half *ptr, float f0, float f1) {
  half2 val = __floats2half2_rn(f0, f1);
  asm volatile("st.cs.b32 [%0], %1;" :: "l"(ptr), "r"(*(uint32_t*)&val) : "memory");
}

static void check_cu(CUresult err) {
  if (err == CUDA_SUCCESS) return;
  const char *error_msg_ptr = nullptr;
  cuGetErrorString(err, &error_msg_ptr);
  TORCH_CHECK(false, "cuTensorMap error: ", (error_msg_ptr ? error_msg_ptr : "unknown"));
}

static void check_cuda(cudaError_t err) {
  if (err == cudaSuccess) return;
  TORCH_CHECK(false, cudaGetErrorString(err));
}

// Shared meta types
constexpr int TMAP_CACHE_CAP = 64;
constexpr int PTR_HT_CAP = 128;

static inline uint32_t hash_u64(uint64_t x) {
  // MurmurHash3 finalizer
  x ^= x >> 33;
  x *= 0xff51afd7ed558ccdULL;
  x ^= x >> 33;
  x *= 0xc4ceb9fe1a85ec53ULL;
  x ^= x >> 33;
  return (uint32_t)x;
}

template <int HT_CAP>
static inline int ht_find(const uint64_t *keys, const uint8_t *vals, uint64_t key) {
  const uint32_t mask = (uint32_t)HT_CAP - 1U;
  uint32_t idx = hash_u64(key) & mask;
  #pragma unroll 1
  for (int p = 0; p < HT_CAP; p++) {
    const uint64_t k = keys[idx];
    if (k == key) return (int)vals[idx];
    if (k == 0) return -1;
    idx = (idx + 1U) & mask;
  }
  return -1;
}

template <int HT_CAP>
static inline void ht_insert(uint64_t *keys, uint8_t *vals, uint64_t key, uint8_t val) {
  const uint32_t mask = (uint32_t)HT_CAP - 1U;
  uint32_t idx = hash_u64(key) & mask;
  #pragma unroll 1
  for (int p = 0; p < HT_CAP; p++) {
    const uint64_t k = keys[idx];
    if (k == 0 || k == key) {
      keys[idx] = key;
      vals[idx] = val;
      return;
    }
    idx = (idx + 1U) & mask;
  }
}

struct __align__(16) Meta {
  uint64_t C[8];
  uint64_t SFA[8];
  uint64_t SFB[8];
  int M[8];
  int N[8];
  int K[8];
  uint8_t A_slot[8];
  uint8_t B_slot[8];
  int offsets[9];
  int num_groups;
  uint64_t tiles_ptr;
  int tiles_count;
};

struct __align__(64) DeviceBlob {
  CUtensorMap A[8 * TMAP_CACHE_CAP];
  CUtensorMap B[8 * TMAP_CACHE_CAP];
  Meta meta;
};

static void init_AB_tmap(
  CUtensorMap *tmap,
  const void *ptr,
  uint64_t global_height, uint64_t global_width,
  uint32_t shared_height, uint32_t shared_width
) {
  constexpr uint32_t rank = 3;
  uint64_t globalDim[rank]       = {256ULL, global_height, global_width / 256ULL};
  uint64_t globalStrides[rank-1] = {global_width / 2ULL, 128ULL};
  uint32_t boxDim[rank]          = {256U, shared_height, shared_width / 256U};
  uint32_t elementStrides[rank]  = {1U, 1U, 1U};

  auto err = cuTensorMapEncodeTiled(
    tmap,
    CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
    rank,
    const_cast<void *>(ptr),
    globalDim,
    globalStrides,
    boxDim,
    elementStrides,
    CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
    CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B,
    CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
    CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
  );
  check_cu(err);
}

// --------------------------
// NP base kernel (BLOCK_N=128, NUM_STAGES=6)
// --------------------------
namespace np_base {

constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;

constexpr int BLOCK_M = 128;
constexpr int BLOCK_N = 128;
constexpr int BLOCK_K = 256;
constexpr int NUM_STAGES = 6;

constexpr uint64_t EVICT_FIRST  = 0x12F0000000000000ULL;
constexpr uint64_t EVICT_LAST   = 0x14F0000000000000ULL;

__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void grouped_kernel(const DeviceBlob *blob, const Meta kmeta) {
  const Meta *meta = &kmeta;
  const int tid = threadIdx.x;
  const int bid = blockIdx.x;
  const int lane_id = tid % WARP_SIZE;
  const int warp_id = tid / WARP_SIZE;

  const int4 tile_s = reinterpret_cast<const int4 *>(meta->tiles_ptr)[bid];

  constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

  extern __shared__ __align__(1024) char smem_ptr[];
  const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
  constexpr int A_size = BLOCK_M * BLOCK_K / 2;
  constexpr int B_size = BLOCK_N * BLOCK_K / 2;
  constexpr int SFA_size = 128 * BLOCK_K / 16;
  constexpr int SFB_size = 128 * BLOCK_K / 16;
  constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;

  #pragma nv_diag_suppress static_var_with_dynamic_init
  __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
  const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
  const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
  const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;

  constexpr int SFA_tmem = BLOCK_N;
  constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K);

  if (warp_id == 0 && elect_sync()) {
    for (int i = 0; i < NUM_STAGES * 2 + 1; i++) mbarrier_init(tma_mbar_addr + i * 8, 1);
    asm volatile("fence.mbarrier_init.release.cluster;");
  } else if (warp_id == 1) {
    asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(BLOCK_N * 2));
  }
  // Only warps {0(init), 1(alloc), 4(TMA), 5(MMA)} must rendezvous here.
  if (warp_id == 0 || warp_id == 1 || warp_id == NUM_WARPS - 2 || warp_id == NUM_WARPS - 1) {
    asm volatile("bar.sync 2, %0;" :: "r"(BLOCK_M) : "memory");
  }


  const int group = tile_s.x;
  const int off_m = tile_s.y;
  const int off_n = tile_s.z;
  const int sfb_lane = tile_s.w;
  const int M = meta->M[group];
  const int N = meta->N[group];
  const int K = meta->K[group];

  const CUtensorMap *A_tmaps = blob->A;
  const CUtensorMap *B_tmaps = blob->B;
  const int a_slot = (int)meta->A_slot[group];
  const int b_slot = (int)meta->B_slot[group];
  const CUtensorMap *A_tmap = A_tmaps + group * TMAP_CACHE_CAP + a_slot;
  const CUtensorMap *B_tmap = B_tmaps + group * TMAP_CACHE_CAP + b_slot;
  if (warp_id == 0 && elect_sync()) {
    // (from nvfp4_dual_gemm/gaunernst.py) prefetch tensor maps early
    asm volatile("prefetch.tensormap [%0];" :: "l"(A_tmap) : "memory");
    asm volatile("prefetch.tensormap [%0];" :: "l"(B_tmap) : "memory");
  }
  const char *SFA_ptr = reinterpret_cast<const char *>(meta->SFA[group]);
  const char *SFB_ptr = reinterpret_cast<const char *>(meta->SFB[group]);
  half *C_ptr = reinterpret_cast<half *>(meta->C[group]);

  const int num_iters = K / BLOCK_K;
  const int rest_k = K / 64;
  uint64_t cache_A, cache_B;
  if (M > N) { cache_A = EVICT_FIRST; cache_B = EVICT_LAST; }
  else { cache_A = EVICT_LAST; cache_B = EVICT_FIRST; }

  if (warp_id == NUM_WARPS - 2 && elect_sync()) {
    const int tileA = off_m >> 7;
    const int tileB = off_n >> 7;
    const char *SFA_base = SFA_ptr + (tileA * rest_k) * 512;
    const char *SFB_base = SFB_ptr + (tileB * rest_k) * 512;

    auto issue_tma = [&](int iter_k, int stage_id) {
      const int mbar_addr = tma_mbar_addr + stage_id * 8;
      const int A_smem = smem + stage_id * STAGE_SIZE;
      const int B_smem = A_smem + A_size;
      const int SFA_smem = B_smem + B_size;
      const int SFB_smem = SFA_smem + SFA_size;

      tma_3d_gmem2smem(B_smem, B_tmap, 0, off_n, iter_k, mbar_addr, cache_B);
      tma_3d_gmem2smem(A_smem, A_tmap, 0, off_m, iter_k, mbar_addr, cache_A);

      const int sf_byte = iter_k << 11;
      const char *SFA_src = SFA_base + sf_byte;
      const char *SFB_src = SFB_base + sf_byte;
      tma_gmem2smem(SFB_smem, SFB_src, SFB_size, mbar_addr, cache_B);
      tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_A);

      asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
                  :: "r"(mbar_addr), "r"(STAGE_SIZE) : "memory");
    };

    for (int iter_k = 0; iter_k < NUM_STAGES && iter_k < num_iters; iter_k++) issue_tma(iter_k, iter_k);
    for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
      const int stage_id = iter_k % NUM_STAGES;
      const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
      mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
      issue_tma(iter_k, stage_id);
    }
  } else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
    constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)BLOCK_N >> 3U << 17U) | ((uint32_t)128 >> 7U << 27U);
    const int scaleA_base = SFA_tmem;
    const int scaleB_base = SFB_tmem + sfb_lane;

    for (int iter_k = 0; iter_k < num_iters; iter_k++) {
      const int stage_id = iter_k % NUM_STAGES;
      const int tma_phase = (iter_k / NUM_STAGES) % 2;
      mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);

      const int A_smem = smem + stage_id * STAGE_SIZE;
      const int B_smem = A_smem + A_size;
      const int SFA_smem = B_smem + B_size;
      const int SFB_smem = SFA_smem + SFA_size;

      auto make_desc_AB = [](int addr) -> uint64_t {
        const int SBO = 8 * 128;
        return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
      };
      auto make_desc_SF = [](int addr) -> uint64_t {
        const int SBO = 8 * 16;
        return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
      };

      constexpr uint64_t SF_desc = make_desc_SF(0);
      const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);
      const uint64_t SFB_desc = SF_desc + ((uint64_t)SFB_smem >> 4ULL);

      #pragma unroll
      for (int k = 0; k < BLOCK_K / MMA_K; k++) {
        const uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL);
        const uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);
        tcgen05_cp_nvfp4(SFA_tmem + k * 4, sfa_desc);
        tcgen05_cp_nvfp4(SFB_tmem + k * 4, sfb_desc);
      }

      #pragma unroll
      for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
        const uint64_t a_desc = make_desc_AB(A_smem + k2 * 32);
        const uint64_t b_desc = make_desc_AB(B_smem + k2 * 32);

        const int k_sf = k2;
        const int scale_A_tmem = scaleA_base + k_sf * 4;
        const int scale_B_tmem = scaleB_base + k_sf * 4;
        const int enable_input_d = (k2 == 0) ? iter_k : 1;
        tcgen05_mma_nvfp4(0, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
      }

      asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                  :: "r"(mma_mbar_addr + stage_id * 8) : "memory");
    }

    asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                :: "r"(mainloop_mbar_addr) : "memory");
  } else if (tid < BLOCK_M) {
    mbarrier_wait(mainloop_mbar_addr, 0);
    asm volatile("tcgen05.fence::after_thread_sync;");

    const bool full_tile = (off_m + BLOCK_M <= M) && (off_n + BLOCK_N <= N);
    if (full_tile) {
      #pragma unroll
      for (int m0 = 0; m0 < 32 / 16; m0++) {
        #pragma unroll
        for (int half_idx = 0; half_idx < 2; half_idx++) {
          #pragma unroll
          for (int i = 0; i < 64 / 8; i += 4) {
            float tmp[4 * 4];
            tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, half_idx * 64 + i * 8);
            asm volatile("tcgen05.wait::ld.sync.aligned;");
            const int row = off_m + warp_id * 32 + m0 * 16 + lane_id / 4;
            #pragma unroll
            for (int ii = 0; ii < 4; ii++) {
              const int col = off_n + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
              const float v00 = tmp[ii * 4 + 0];
              const float v01 = tmp[ii * 4 + 1];
              const float v10 = tmp[ii * 4 + 2];
              const float v11 = tmp[ii * 4 + 3];
              store_cs_half2(C_ptr + (row + 0) * N + col, v00, v01);
              store_cs_half2(C_ptr + (row + 8) * N + col, v10, v11);
            }
          }
        }
      }
    } else {
      #pragma unroll
      for (int m0 = 0; m0 < 32 / 16; m0++) {
        #pragma unroll
        for (int half_idx = 0; half_idx < 2; half_idx++) {
          #pragma unroll
          for (int i = 0; i < 64 / 8; i += 4) {
            float tmp[4 * 4];
            tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, half_idx * 64 + i * 8);
            asm volatile("tcgen05.wait::ld.sync.aligned;");
            const int row = off_m + warp_id * 32 + m0 * 16 + lane_id / 4;
            if (row >= M) continue;
            #pragma unroll
            for (int ii = 0; ii < 4; ii++) {
              const int col = off_n + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
              const float v00 = tmp[ii * 4 + 0];
              const float v01 = tmp[ii * 4 + 1];
              store_cs_half2(C_ptr + (row + 0) * N + col, v00, v01);
              if (row + 8 < M) {
                const float v10 = tmp[ii * 4 + 2];
                const float v11 = tmp[ii * 4 + 3];
                store_cs_half2(C_ptr + (row + 8) * N + col, v10, v11);
              }
            }
          }
        }
      }
    }

    asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
    if (warp_id == 0) asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 2));
  }
}

static void group_gemm(
  c10::List<at::Tensor> A_list,
  c10::List<at::Tensor> B_list,
  c10::List<at::Tensor> C_list,
  c10::List<at::Tensor> SFA_list,
  c10::List<at::Tensor> SFB_list,
  at::Tensor problem_sizes
) {
  const int64_t G = A_list.size();
  // Fast path: assume inputs satisfy task constraints (CPU int32 problem_sizes, 1..8 groups, all CUDA tensors).

  struct Cache {
    bool inited = false;
    int lastM[8] = {};
    int lastN[8] = {};
    int lastK[8] = {};
    int lastTilesN[8] = {};
    int A_size[8] = {};
    int A_next[8] = {};
    int B_size[8] = {};
    int B_next[8] = {};
    uint64_t A_ptr[8][TMAP_CACHE_CAP] = {};
    uint64_t B_ptr[8][TMAP_CACHE_CAP] = {};
    uint64_t A_ht_key[8][PTR_HT_CAP] = {};
    uint8_t  A_ht_val[8][PTR_HT_CAP] = {};
    uint64_t B_ht_key[8][PTR_HT_CAP] = {};
    uint8_t  B_ht_val[8][PTR_HT_CAP] = {};
    CUtensorMap A_tbl[8][TMAP_CACHE_CAP];
    CUtensorMap B_tbl[8][TMAP_CACHE_CAP];
    CUtensorMap A_template[8];
    CUtensorMap B_template[8];
    at::Tensor dBlob_u8;
    enum { META_RING = 256 };
    void *hMeta[META_RING] = {};
    int meta_idx = 0;
    at::Tensor dTiles;
    void *hTiles = nullptr;
    int tiles_cap = 0;
    int last_total_tiles = -1;
  };
  thread_local Cache cache;

  if (!cache.inited) {
    auto opts_u8  = at::TensorOptions().dtype(at::kByte).device(at::kCUDA);
    cache.dBlob_u8 = at::empty({(int64_t)sizeof(DeviceBlob)}, opts_u8);
    for (int b = 0; b < Cache::META_RING; b++) check_cuda(cudaHostAlloc(&cache.hMeta[b], sizeof(Meta), cudaHostAllocPortable));
    cache.inited = true;
  }

  uint8_t *blob_u8 = cache.dBlob_u8.data_ptr<uint8_t>();
  const size_t offA = offsetof(DeviceBlob, A);
  const size_t offB = offsetof(DeviceBlob, B);
  const size_t offM = offsetof(DeviceBlob, meta);

  void *hmeta_buf = cache.hMeta[cache.meta_idx];
  cache.meta_idx = (cache.meta_idx + 1) & (Cache::META_RING - 1);
  Meta *hmeta = reinterpret_cast<Meta *>(hmeta_buf);
  hmeta->offsets[0] = 0;
  hmeta->num_groups = (int)G;

  const int *ps_ptr = problem_sizes.data_ptr<int>();

  bool amap_dirty = false;
  bool bmap_dirty = false;
  bool tiles_dirty = false;

  for (int i = 0; i < (int)G; i++) {
    const int M = ps_ptr[i * 4 + 0];
    const int N = ps_ptr[i * 4 + 1];
    const int K = ps_ptr[i * 4 + 2];
    hmeta->M[i] = M; hmeta->N[i] = N; hmeta->K[i] = K;

    auto A = A_list.get(i);
    auto B = B_list.get(i);
    auto C = C_list.get(i);
    auto SFA = SFA_list.get(i);
    auto SFB = SFB_list.get(i);

    const uint64_t Ap = (uint64_t)A.data_ptr();
    const uint64_t Bp = (uint64_t)B.data_ptr();

    hmeta->C[i] = (uint64_t)C.data_ptr();
    hmeta->SFA[i] = (uint64_t)SFA.data_ptr();
    hmeta->SFB[i] = (uint64_t)SFB.data_ptr();

    const int tiles_m = (M + BLOCK_M - 1) / BLOCK_M;
    const int tiles_n = (N + BLOCK_N - 1) / BLOCK_N;
    hmeta->offsets[i + 1] = hmeta->offsets[i] + tiles_m * tiles_n;
    const bool shape_changed = (cache.lastM[i] != M) || (cache.lastN[i] != N) || (cache.lastK[i] != K);
    if (shape_changed || cache.lastTilesN[i] != tiles_n) {
      cache.lastTilesN[i] = tiles_n;
      tiles_dirty = true;
    }

    if (shape_changed) {
      cache.lastM[i] = M;
      cache.lastN[i] = N;
      cache.lastK[i] = K;
      cache.A_size[i] = 0; cache.A_next[i] = 0;
      cache.B_size[i] = 0; cache.B_next[i] = 0;
      for (int t = 0; t < PTR_HT_CAP; t++) { cache.A_ht_key[i][t] = 0; cache.B_ht_key[i][t] = 0; }
      init_AB_tmap(&cache.A_template[i], (const void*)Ap, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K);
      init_AB_tmap(&cache.B_template[i], (const void*)Bp, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K);
      // Seed slot 0 for both A/B with the current pointers.
      cache.A_size[i] = 1; cache.A_next[i] = 1;
      cache.A_ptr[i][0] = Ap;
      cache.A_tbl[i][0] = cache.A_template[i];
      cache.B_size[i] = 1; cache.B_next[i] = 1;
      cache.B_ptr[i][0] = Bp;
      cache.B_tbl[i][0] = cache.B_template[i];
      ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap, (uint8_t)0);
      ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp, (uint8_t)0);
      amap_dirty = true;
      bmap_dirty = true;
    }

    int a_slot = ht_find<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap);
    if (a_slot >= 0 && cache.A_ptr[i][a_slot] != Ap) {
      // Stale mapping (eviction/wrap): rebuild hash table for this group and retry once.
      for (int t = 0; t < PTR_HT_CAP; t++) cache.A_ht_key[i][t] = 0;
      for (int s = 0; s < cache.A_size[i]; s++) {
        const uint64_t p = cache.A_ptr[i][s];
        if (p) ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], p, (uint8_t)s);
      }
      a_slot = ht_find<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap);
    }
    if (a_slot < 0) {
      a_slot = cache.A_next[i];
      cache.A_next[i] = (cache.A_next[i] + 1) % TMAP_CACHE_CAP;
      if (cache.A_size[i] < TMAP_CACHE_CAP) cache.A_size[i]++;
      cache.A_ptr[i][a_slot] = Ap;
      cache.A_tbl[i][a_slot] = cache.A_template[i];
      check_cu(cuTensorMapReplaceAddress(&cache.A_tbl[i][a_slot], (void*)Ap));
      ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap, (uint8_t)a_slot);
      amap_dirty = true;
    }
    hmeta->A_slot[i] = (uint8_t)a_slot;

    int b_slot = ht_find<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp);
    if (b_slot >= 0 && cache.B_ptr[i][b_slot] != Bp) {
      for (int t = 0; t < PTR_HT_CAP; t++) cache.B_ht_key[i][t] = 0;
      for (int s = 0; s < cache.B_size[i]; s++) {
        const uint64_t p = cache.B_ptr[i][s];
        if (p) ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], p, (uint8_t)s);
      }
      b_slot = ht_find<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp);
    }
    if (b_slot < 0) {
      b_slot = cache.B_next[i];
      cache.B_next[i] = (cache.B_next[i] + 1) % TMAP_CACHE_CAP;
      if (cache.B_size[i] < TMAP_CACHE_CAP) cache.B_size[i]++;
      cache.B_ptr[i][b_slot] = Bp;
      cache.B_tbl[i][b_slot] = cache.B_template[i];
      check_cu(cuTensorMapReplaceAddress(&cache.B_tbl[i][b_slot], (void*)Bp));
      ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp, (uint8_t)b_slot);
      bmap_dirty = true;
    }
    hmeta->B_slot[i] = (uint8_t)b_slot;
  }

  const int total_tiles = hmeta->offsets[G];
  if (total_tiles == 0) return;
  if (total_tiles != cache.last_total_tiles) {
    cache.last_total_tiles = total_tiles;
    tiles_dirty = true;
  }

  if (tiles_dirty) {
    if (total_tiles > cache.tiles_cap) {
      auto opts_i32 = at::TensorOptions().dtype(at::kInt).device(at::kCUDA);
      cache.dTiles = at::empty({(int64_t)total_tiles, 4}, opts_i32);
      if (cache.hTiles) check_cuda(cudaFreeHost(cache.hTiles));
      check_cuda(cudaHostAlloc(&cache.hTiles, (size_t)total_tiles * sizeof(int4), cudaHostAllocPortable));
      cache.tiles_cap = total_tiles;
    }
    auto *tiles = reinterpret_cast<int4 *>(cache.hTiles);

    constexpr int GROUP_SIZE_M = 4;
    int order[8];
    for (int i = 0; i < (int)G; i++) order[i] = i;
    for (int i = 1; i < (int)G; i++) {
      const int key = order[i];
      const int keyK = hmeta->K[key];
      int j = i - 1;
      while (j >= 0 && hmeta->K[order[j]] < keyK) {
        order[j + 1] = order[j];
        j--;
      }
      order[j + 1] = key;
    }

    int t = 0;
    for (int oi = 0; oi < (int)G; oi++) {
      const int g = order[oi];
      const int M = hmeta->M[g];
      const int N = hmeta->N[g];
      const int tiles_m = (M + BLOCK_M - 1) / BLOCK_M;
      const int tiles_n = (N + BLOCK_N - 1) / BLOCK_N;

      for (int first_tm = 0; first_tm < tiles_m; first_tm += GROUP_SIZE_M) {
        const int group_size_m = min(GROUP_SIZE_M, tiles_m - first_tm);
        for (int tn = 0; tn < tiles_n; tn++) {
          for (int i = 0; i < group_size_m; i++) {
            const int tm = first_tm + i;
            const int off_m = tm * BLOCK_M;
            const int off_n = tn * BLOCK_N;
            const int sfb_lane = 0;
            tiles[t++] = make_int4(g, off_m, off_n, sfb_lane);
          }
        }
      }
    }
    check_cuda(cudaMemcpyAsync(cache.dTiles.data_ptr<int>(), tiles, (size_t)total_tiles * sizeof(int4), cudaMemcpyHostToDevice, 0));
  }

  hmeta->tiles_ptr = (uint64_t)cache.dTiles.data_ptr();
  hmeta->tiles_count = total_tiles;

  // Meta passed as kernel argument (constant memory) -- no H2D copy needed
  if (amap_dirty) check_cuda(cudaMemcpyAsync(blob_u8 + offA, &cache.A_tbl[0][0], (size_t)G * TMAP_CACHE_CAP * sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0));
  if (bmap_dirty) check_cuda(cudaMemcpyAsync(blob_u8 + offB, &cache.B_tbl[0][0], (size_t)G * TMAP_CACHE_CAP * sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0));

  dim3 grid(total_tiles, 1, 1);
  const int tb = BLOCK_M + 2 * WARP_SIZE;
  const int AB_size = (BLOCK_M + BLOCK_N) * (BLOCK_K / 2);
  const int SFAB_size = 128 * (BLOCK_K / 16) * 2;
  const int smem_size = (AB_size + SFAB_size) * NUM_STAGES;
  if (smem_size > 48'000) {
    static bool smem_attr_set = false;
    if (!smem_attr_set) {
      cudaFuncSetAttribute(grouped_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
      smem_attr_set = true;
    }
  }
  grouped_kernel<<<grid, tb, smem_size>>>((const DeviceBlob*)cache.dBlob_u8.data_ptr<uint8_t>(), *hmeta);
}

} // namespace np_base

// --------------------------
// NP mtile2 kernel (reuse B/SFB across 2 M tiles)
// --------------------------
namespace np_mtile2 {

constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;

constexpr int BLOCK_M = 128;
constexpr int BLOCK_N = 128;
constexpr int BLOCK_K = 256;
constexpr int NUM_STAGES = 4;

constexpr int ACCUM_STRIDE_TMEM = 128;
constexpr int SFA_tmem = 2 * ACCUM_STRIDE_TMEM;
constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K);
constexpr int TMEM_ALLOC = 512;

constexpr uint64_t EVICT_FIRST  = 0x12F0000000000000ULL;
constexpr uint64_t EVICT_LAST   = 0x14F0000000000000ULL;

__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void grouped_kernel_mtile2(const DeviceBlob *blob, const Meta kmeta) {
  const Meta *meta = &kmeta;
  const int tid = threadIdx.x;
  const int bid = blockIdx.x;
  const int lane_id = tid % WARP_SIZE;
  const int warp_id = tid / WARP_SIZE;

  const int4 tile_s = reinterpret_cast<const int4 *>(meta->tiles_ptr)[bid];

  constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

  extern __shared__ __align__(1024) char smem_ptr[];
  const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
  constexpr int A_tile = BLOCK_M * BLOCK_K / 2;
  constexpr int B_size = BLOCK_N * BLOCK_K / 2;
  constexpr int SFA_tile = 128 * BLOCK_K / 16;
  constexpr int SFB_size = 128 * BLOCK_K / 16;
  constexpr int STAGE_SIZE = A_tile * 2 + B_size + SFA_tile * 2 + SFB_size;

  #pragma nv_diag_suppress static_var_with_dynamic_init
  __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
  const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
  const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
  const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;

  if (warp_id == 0 && elect_sync()) {
    for (int i = 0; i < NUM_STAGES * 2 + 1; i++) mbarrier_init(tma_mbar_addr + i * 8, 1);
    asm volatile("fence.mbarrier_init.release.cluster;");
  } else if (warp_id == 1) {
    asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(TMEM_ALLOC));
  }
  // Only warps {0(init), 1(alloc), 4(TMA), 5(MMA)} must rendezvous here.
  if (warp_id == 0 || warp_id == 1 || warp_id == NUM_WARPS - 2 || warp_id == NUM_WARPS - 1) {
    asm volatile("bar.sync 2, %0;" :: "r"(BLOCK_M) : "memory");
  }

  const int group = tile_s.x;
  const int off_m = tile_s.y;
  const int off_n = tile_s.z;
  const int sfb_lane = tile_s.w;
  const int M = meta->M[group];
  const int N = meta->N[group];
  const int K = meta->K[group];
  const int off_m1 = off_m + BLOCK_M;
  const bool has_m1 = (off_m + BLOCK_M < M);

  const CUtensorMap *A_tmaps = blob->A;
  const CUtensorMap *B_tmaps = blob->B;
  const int a_slot = (int)meta->A_slot[group];
  const int b_slot = (int)meta->B_slot[group];
  const CUtensorMap *A_tmap = A_tmaps + group * TMAP_CACHE_CAP + a_slot;
  const CUtensorMap *B_tmap = B_tmaps + group * TMAP_CACHE_CAP + b_slot;
  if (warp_id == 0 && elect_sync()) {
    // (from nvfp4_dual_gemm/gaunernst.py) prefetch tensor maps early
    asm volatile("prefetch.tensormap [%0];" :: "l"(A_tmap) : "memory");
    asm volatile("prefetch.tensormap [%0];" :: "l"(B_tmap) : "memory");
  }
  const char *SFA_ptr = reinterpret_cast<const char *>(meta->SFA[group]);
  const char *SFB_ptr = reinterpret_cast<const char *>(meta->SFB[group]);
  half *C_ptr = reinterpret_cast<half *>(meta->C[group]);

  const int num_iters = K / BLOCK_K;
  const int rest_k = K / 64;
  uint64_t cache_A, cache_B;
  if (M > N) { cache_A = EVICT_FIRST; cache_B = EVICT_LAST; }
  else { cache_A = EVICT_LAST; cache_B = EVICT_FIRST; }

  if (warp_id == NUM_WARPS - 2 && elect_sync()) {
    const int tileA0 = off_m >> 7;
    const int tileA1 = off_m1 >> 7;
    const int tileB = off_n >> 7;
    const char *SFA_base0 = SFA_ptr + (tileA0 * rest_k) * 512;
    const char *SFA_base1 = SFA_ptr + (tileA1 * rest_k) * 512;
    const char *SFB_base = SFB_ptr + (tileB * rest_k) * 512;

    auto issue_tma = [&](int iter_k, int stage_id) {
      const int mbar_addr = tma_mbar_addr + stage_id * 8;
      const int A0_smem = smem + stage_id * STAGE_SIZE;
      const int A1_smem = A0_smem + A_tile;
      const int B_smem = A1_smem + A_tile;
      const int SFA0_smem = B_smem + B_size;
      const int SFA1_smem = SFA0_smem + SFA_tile;
      const int SFB_smem = SFA1_smem + SFA_tile;

      tma_3d_gmem2smem(B_smem, B_tmap, 0, off_n, iter_k, mbar_addr, cache_B);
      tma_3d_gmem2smem(A0_smem, A_tmap, 0, off_m, iter_k, mbar_addr, cache_A);

      const int sf_byte = iter_k << 11;
      int stage_bytes = A_tile + SFA_tile + B_size + SFB_size;
      if (has_m1) {
        tma_3d_gmem2smem(A1_smem, A_tmap, 0, off_m1, iter_k, mbar_addr, cache_A);
        stage_bytes += A_tile + SFA_tile;
      }
      // issue order like gaunernst: SFB before SFA
      tma_gmem2smem(SFB_smem, SFB_base + sf_byte, SFB_size, mbar_addr, cache_B);
      tma_gmem2smem(SFA0_smem, SFA_base0 + sf_byte, SFA_tile, mbar_addr, cache_A);
      if (has_m1) {
        tma_gmem2smem(SFA1_smem, SFA_base1 + sf_byte, SFA_tile, mbar_addr, cache_A);
      }

      asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
                  :: "r"(mbar_addr), "r"(stage_bytes) : "memory");
    };

    for (int iter_k = 0; iter_k < NUM_STAGES && iter_k < num_iters; iter_k++) issue_tma(iter_k, iter_k);
    for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
      const int stage_id = iter_k % NUM_STAGES;
      const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
      mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
      issue_tma(iter_k, stage_id);
    }
  } else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
    constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)BLOCK_N >> 3U << 17U) | ((uint32_t)128 >> 7U << 27U);
    const int scaleA_base = SFA_tmem;
    const int scaleB_base = SFB_tmem + sfb_lane;
    const int d_tmem0 = 0;
    const int d_tmem1 = ACCUM_STRIDE_TMEM;

    for (int iter_k = 0; iter_k < num_iters; iter_k++) {
      const int stage_id = iter_k % NUM_STAGES;
      const int tma_phase = (iter_k / NUM_STAGES) % 2;
      mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);

      const int A0_smem = smem + stage_id * STAGE_SIZE;
      const int A1_smem = A0_smem + A_tile;
      const int B_smem = A1_smem + A_tile;
      const int SFA0_smem = B_smem + B_size;
      const int SFA1_smem = SFA0_smem + SFA_tile;
      const int SFB_smem = SFA1_smem + SFA_tile;

      auto make_desc_AB = [](int addr) -> uint64_t {
        const int SBO = 8 * 128;
        return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
      };
      auto make_desc_SF = [](int addr) -> uint64_t {
        const int SBO = 8 * 16;
        return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
      };

      constexpr uint64_t SF_desc = make_desc_SF(0);
      const uint64_t SFA0_desc = SF_desc + ((uint64_t)SFA0_smem >> 4ULL);
      const uint64_t SFA1_desc = SF_desc + ((uint64_t)SFA1_smem >> 4ULL);
      const uint64_t SFB_desc = SF_desc + ((uint64_t)SFB_smem >> 4ULL);

      #pragma unroll
      for (int k = 0; k < BLOCK_K / MMA_K; k++) {
        const uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);
        tcgen05_cp_nvfp4(SFB_tmem + k * 4, sfb_desc);
      }

      #pragma unroll
      for (int k = 0; k < BLOCK_K / MMA_K; k++) {
        const uint64_t sfa_desc = SFA0_desc + (uint64_t)k * (512ULL >> 4ULL);
        tcgen05_cp_nvfp4(SFA_tmem + k * 4, sfa_desc);
      }

      #pragma unroll
      for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
        const uint64_t a_desc = make_desc_AB(A0_smem + k2 * 32);
        const uint64_t b_desc = make_desc_AB(B_smem + k2 * 32);
        const int k_sf = k2;
        const int scale_A_tmem = scaleA_base + k_sf * 4;
        const int scale_B_tmem = scaleB_base + k_sf * 4;
        const int enable_input_d = (k2 == 0) ? iter_k : 1;
        tcgen05_mma_nvfp4(d_tmem0, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
      }

      if (has_m1) {
        #pragma unroll
        for (int k = 0; k < BLOCK_K / MMA_K; k++) {
          const uint64_t sfa_desc = SFA1_desc + (uint64_t)k * (512ULL >> 4ULL);
          tcgen05_cp_nvfp4(SFA_tmem + k * 4, sfa_desc);
        }
        #pragma unroll
        for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
          const uint64_t a_desc = make_desc_AB(A1_smem + k2 * 32);
          const uint64_t b_desc = make_desc_AB(B_smem + k2 * 32);
          const int k_sf = k2;
          const int scale_A_tmem = scaleA_base + k_sf * 4;
          const int scale_B_tmem = scaleB_base + k_sf * 4;
          const int enable_input_d = (k2 == 0) ? iter_k : 1;
          tcgen05_mma_nvfp4(d_tmem1, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
        }
      }

      asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                  :: "r"(mma_mbar_addr + stage_id * 8) : "memory");
    }

    asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                :: "r"(mainloop_mbar_addr) : "memory");
  } else if (tid < BLOCK_M) {
    mbarrier_wait(mainloop_mbar_addr, 0);
    asm volatile("tcgen05.fence::after_thread_sync;");

    const int d_tmem0 = 0;
    const int d_tmem1 = ACCUM_STRIDE_TMEM;

    const bool full_tile0 = (off_m + BLOCK_M <= M) && (off_n + BLOCK_N <= N);
    if (full_tile0) {
      #pragma unroll
      for (int m0 = 0; m0 < 32 / 16; m0++) {
        #pragma unroll
        for (int half_idx = 0; half_idx < 2; half_idx++) {
          #pragma unroll
          for (int i = 0; i < 64 / 8; i += 4) {
            float tmp[4 * 4];
            tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, d_tmem0 + half_idx * 64 + i * 8);
            asm volatile("tcgen05.wait::ld.sync.aligned;");
            const int row = off_m + warp_id * 32 + m0 * 16 + lane_id / 4;
            #pragma unroll
            for (int ii = 0; ii < 4; ii++) {
              const int col = off_n + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
              const float v00 = tmp[ii * 4 + 0];
              const float v01 = tmp[ii * 4 + 1];
              const float v10 = tmp[ii * 4 + 2];
              const float v11 = tmp[ii * 4 + 3];
              store_cs_half2(C_ptr + (row + 0) * N + col, v00, v01);
              store_cs_half2(C_ptr + (row + 8) * N + col, v10, v11);
            }
          }
        }
      }
    } else {
      #pragma unroll
      for (int m0 = 0; m0 < 32 / 16; m0++) {
        #pragma unroll
        for (int half_idx = 0; half_idx < 2; half_idx++) {
          #pragma unroll
          for (int i = 0; i < 64 / 8; i += 4) {
            float tmp[4 * 4];
            tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, d_tmem0 + half_idx * 64 + i * 8);
            asm volatile("tcgen05.wait::ld.sync.aligned;");
            const int row = off_m + warp_id * 32 + m0 * 16 + lane_id / 4;
            if (row >= M) continue;
            #pragma unroll
            for (int ii = 0; ii < 4; ii++) {
              const int col = off_n + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
              const float v00 = tmp[ii * 4 + 0];
              const float v01 = tmp[ii * 4 + 1];
              store_cs_half2(C_ptr + (row + 0) * N + col, v00, v01);
              if (row + 8 < M) {
                const float v10 = tmp[ii * 4 + 2];
                const float v11 = tmp[ii * 4 + 3];
                store_cs_half2(C_ptr + (row + 8) * N + col, v10, v11);
              }
            }
          }
        }
      }
    }

    if (has_m1) {
      const int off_m1_base = off_m + BLOCK_M;
      const bool full_tile1 = (off_m1_base + BLOCK_M <= M) && (off_n + BLOCK_N <= N);
      if (full_tile1) {
        #pragma unroll
        for (int m0 = 0; m0 < 32 / 16; m0++) {
          #pragma unroll
          for (int half_idx = 0; half_idx < 2; half_idx++) {
            #pragma unroll
            for (int i = 0; i < 64 / 8; i += 4) {
              float tmp[4 * 4];
              tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, d_tmem1 + half_idx * 64 + i * 8);
              asm volatile("tcgen05.wait::ld.sync.aligned;");
              const int row = off_m1_base + warp_id * 32 + m0 * 16 + lane_id / 4;
              #pragma unroll
              for (int ii = 0; ii < 4; ii++) {
                const int col = off_n + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
                const float v00 = tmp[ii * 4 + 0];
                const float v01 = tmp[ii * 4 + 1];
                const float v10 = tmp[ii * 4 + 2];
                const float v11 = tmp[ii * 4 + 3];
                store_cs_half2(C_ptr + (row + 0) * N + col, v00, v01);
                store_cs_half2(C_ptr + (row + 8) * N + col, v10, v11);
              }
            }
          }
        }
      } else {
        #pragma unroll
        for (int m0 = 0; m0 < 32 / 16; m0++) {
          #pragma unroll
          for (int half_idx = 0; half_idx < 2; half_idx++) {
            #pragma unroll
            for (int i = 0; i < 64 / 8; i += 4) {
              float tmp[4 * 4];
              tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, d_tmem1 + half_idx * 64 + i * 8);
              asm volatile("tcgen05.wait::ld.sync.aligned;");
              const int row = off_m1_base + warp_id * 32 + m0 * 16 + lane_id / 4;
              if (row >= M) continue;
              #pragma unroll
              for (int ii = 0; ii < 4; ii++) {
                const int col = off_n + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
                const float v00 = tmp[ii * 4 + 0];
                const float v01 = tmp[ii * 4 + 1];
                store_cs_half2(C_ptr + (row + 0) * N + col, v00, v01);
                if (row + 8 < M) {
                  const float v10 = tmp[ii * 4 + 2];
                  const float v11 = tmp[ii * 4 + 3];
                  store_cs_half2(C_ptr + (row + 8) * N + col, v10, v11);
                }
              }
            }
          }
        }
      }
    }

    asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
    if (warp_id == 0) asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_ALLOC));
  }
}

static void group_gemm(
  c10::List<at::Tensor> A_list,
  c10::List<at::Tensor> B_list,
  c10::List<at::Tensor> C_list,
  c10::List<at::Tensor> SFA_list,
  c10::List<at::Tensor> SFB_list,
  at::Tensor problem_sizes
) {
  const int64_t G = A_list.size();
  // Fast path: assume inputs satisfy task constraints (CPU int32 problem_sizes, 1..8 groups, all CUDA tensors).

  struct Cache {
    bool inited = false;
    int lastM[8] = {};
    int lastN[8] = {};
    int lastK[8] = {};
    int lastTilesN[8] = {};
    int A_size[8] = {};
    int A_next[8] = {};
    int B_size[8] = {};
    int B_next[8] = {};
    uint64_t A_ptr[8][TMAP_CACHE_CAP] = {};
    uint64_t B_ptr[8][TMAP_CACHE_CAP] = {};
    uint64_t A_ht_key[8][PTR_HT_CAP] = {};
    uint8_t  A_ht_val[8][PTR_HT_CAP] = {};
    uint64_t B_ht_key[8][PTR_HT_CAP] = {};
    uint8_t  B_ht_val[8][PTR_HT_CAP] = {};
    CUtensorMap A_tbl[8][TMAP_CACHE_CAP];
    CUtensorMap B_tbl[8][TMAP_CACHE_CAP];
    CUtensorMap A_template[8];
    CUtensorMap B_template[8];
    at::Tensor dBlob_u8;
    enum { META_RING = 256 };
    void *hMeta[META_RING] = {};
    int meta_idx = 0;
    at::Tensor dTiles;
    void *hTiles = nullptr;
    int tiles_cap = 0;
    int last_total_tiles = -1;
  };
  thread_local Cache cache;

  if (!cache.inited) {
    auto opts_u8  = at::TensorOptions().dtype(at::kByte).device(at::kCUDA);
    cache.dBlob_u8 = at::empty({(int64_t)sizeof(DeviceBlob)}, opts_u8);
    for (int b = 0; b < Cache::META_RING; b++) check_cuda(cudaHostAlloc(&cache.hMeta[b], sizeof(Meta), cudaHostAllocPortable));
    cache.inited = true;
  }

  uint8_t *blob_u8 = cache.dBlob_u8.data_ptr<uint8_t>();
  const size_t offA = offsetof(DeviceBlob, A);
  const size_t offB = offsetof(DeviceBlob, B);
  const size_t offM = offsetof(DeviceBlob, meta);

  void *hmeta_buf = cache.hMeta[cache.meta_idx];
  cache.meta_idx = (cache.meta_idx + 1) & (Cache::META_RING - 1);
  Meta *hmeta = reinterpret_cast<Meta *>(hmeta_buf);
  hmeta->offsets[0] = 0;
  hmeta->num_groups = (int)G;

  const int *ps_ptr = problem_sizes.data_ptr<int>();

  bool amap_dirty = false;
  bool bmap_dirty = false;
  bool tiles_dirty = false;

  for (int i = 0; i < (int)G; i++) {
    const int M = ps_ptr[i * 4 + 0];
    const int N = ps_ptr[i * 4 + 1];
    const int K = ps_ptr[i * 4 + 2];
    hmeta->M[i] = M; hmeta->N[i] = N; hmeta->K[i] = K;

    auto A = A_list.get(i);
    auto B = B_list.get(i);
    auto C = C_list.get(i);
    auto SFA = SFA_list.get(i);
    auto SFB = SFB_list.get(i);

    const uint64_t Ap = (uint64_t)A.data_ptr();
    const uint64_t Bp = (uint64_t)B.data_ptr();

    hmeta->C[i] = (uint64_t)C.data_ptr();
    hmeta->SFA[i] = (uint64_t)SFA.data_ptr();
    hmeta->SFB[i] = (uint64_t)SFB.data_ptr();

    const int tiles_m = (M + BLOCK_M - 1) / BLOCK_M;
    const int tiles_n = (N + BLOCK_N - 1) / BLOCK_N;
    const int cluster_m = (tiles_m + 1) / 2;
    hmeta->offsets[i + 1] = hmeta->offsets[i] + cluster_m * tiles_n;
    const bool shape_changed = (cache.lastM[i] != M) || (cache.lastN[i] != N) || (cache.lastK[i] != K);
    if (shape_changed || cache.lastTilesN[i] != tiles_n) {
      cache.lastTilesN[i] = tiles_n;
      tiles_dirty = true;
    }

    if (shape_changed) {
      cache.lastM[i] = M;
      cache.lastN[i] = N;
      cache.lastK[i] = K;
      cache.A_size[i] = 0; cache.A_next[i] = 0;
      cache.B_size[i] = 0; cache.B_next[i] = 0;
      for (int t = 0; t < PTR_HT_CAP; t++) { cache.A_ht_key[i][t] = 0; cache.B_ht_key[i][t] = 0; }
      init_AB_tmap(&cache.A_template[i], (const void*)Ap, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K);
      init_AB_tmap(&cache.B_template[i], (const void*)Bp, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K);
      // Seed slot 0 for both A/B with the current pointers.
      cache.A_size[i] = 1; cache.A_next[i] = 1;
      cache.A_ptr[i][0] = Ap;
      cache.A_tbl[i][0] = cache.A_template[i];
      cache.B_size[i] = 1; cache.B_next[i] = 1;
      cache.B_ptr[i][0] = Bp;
      cache.B_tbl[i][0] = cache.B_template[i];
      ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap, (uint8_t)0);
      ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp, (uint8_t)0);
      amap_dirty = true;
      bmap_dirty = true;
    }

    int a_slot = ht_find<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap);
    if (a_slot >= 0 && cache.A_ptr[i][a_slot] != Ap) {
      for (int t = 0; t < PTR_HT_CAP; t++) cache.A_ht_key[i][t] = 0;
      for (int s = 0; s < cache.A_size[i]; s++) {
        const uint64_t p = cache.A_ptr[i][s];
        if (p) ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], p, (uint8_t)s);
      }
      a_slot = ht_find<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap);
    }
    if (a_slot < 0) {
      a_slot = cache.A_next[i];
      cache.A_next[i] = (cache.A_next[i] + 1) % TMAP_CACHE_CAP;
      if (cache.A_size[i] < TMAP_CACHE_CAP) cache.A_size[i]++;
      cache.A_ptr[i][a_slot] = Ap;
      cache.A_tbl[i][a_slot] = cache.A_template[i];
      check_cu(cuTensorMapReplaceAddress(&cache.A_tbl[i][a_slot], (void*)Ap));
      ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap, (uint8_t)a_slot);
      amap_dirty = true;
    }
    hmeta->A_slot[i] = (uint8_t)a_slot;

    int b_slot = ht_find<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp);
    if (b_slot >= 0 && cache.B_ptr[i][b_slot] != Bp) {
      for (int t = 0; t < PTR_HT_CAP; t++) cache.B_ht_key[i][t] = 0;
      for (int s = 0; s < cache.B_size[i]; s++) {
        const uint64_t p = cache.B_ptr[i][s];
        if (p) ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], p, (uint8_t)s);
      }
      b_slot = ht_find<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp);
    }
    if (b_slot < 0) {
      b_slot = cache.B_next[i];
      cache.B_next[i] = (cache.B_next[i] + 1) % TMAP_CACHE_CAP;
      if (cache.B_size[i] < TMAP_CACHE_CAP) cache.B_size[i]++;
      cache.B_ptr[i][b_slot] = Bp;
      cache.B_tbl[i][b_slot] = cache.B_template[i];
      check_cu(cuTensorMapReplaceAddress(&cache.B_tbl[i][b_slot], (void*)Bp));
      ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp, (uint8_t)b_slot);
      bmap_dirty = true;
    }
    hmeta->B_slot[i] = (uint8_t)b_slot;
  }

  const int total_tiles = hmeta->offsets[G];
  if (total_tiles == 0) return;
  if (total_tiles != cache.last_total_tiles) {
    cache.last_total_tiles = total_tiles;
    tiles_dirty = true;
  }

  if (tiles_dirty) {
    if (total_tiles > cache.tiles_cap) {
      auto opts_i32 = at::TensorOptions().dtype(at::kInt).device(at::kCUDA);
      cache.dTiles = at::empty({(int64_t)total_tiles, 4}, opts_i32);
      if (cache.hTiles) check_cuda(cudaFreeHost(cache.hTiles));
      check_cuda(cudaHostAlloc(&cache.hTiles, (size_t)total_tiles * sizeof(int4), cudaHostAllocPortable));
      cache.tiles_cap = total_tiles;
    }
    auto *tiles = reinterpret_cast<int4 *>(cache.hTiles);

    constexpr int GROUP_SIZE_M = 4;
    int order[8];
    for (int i = 0; i < (int)G; i++) order[i] = i;
    for (int i = 1; i < (int)G; i++) {
      const int key = order[i];
      const int keyK = hmeta->K[key];
      int j = i - 1;
      while (j >= 0 && hmeta->K[order[j]] < keyK) {
        order[j + 1] = order[j];
        j--;
      }
      order[j + 1] = key;
    }

    int t = 0;
    for (int oi = 0; oi < (int)G; oi++) {
      const int g = order[oi];
      const int M = hmeta->M[g];
      const int N = hmeta->N[g];
      const int tiles_m = (M + BLOCK_M - 1) / BLOCK_M;
      const int tiles_n = (N + BLOCK_N - 1) / BLOCK_N;

      for (int first_tm = 0; first_tm < tiles_m; first_tm += GROUP_SIZE_M) {
        const int group_size_m = min(GROUP_SIZE_M, tiles_m - first_tm);
        for (int tn = 0; tn < tiles_n; tn++) {
          for (int i = 0; i < group_size_m; i++) {
            const int tm = first_tm + i;
            if (tm & 1) continue;
            const int off_m = tm * BLOCK_M;
            const int off_n = tn * BLOCK_N;
            const int sfb_lane = 0;
            tiles[t++] = make_int4(g, off_m, off_n, sfb_lane);
          }
        }
      }
    }
    check_cuda(cudaMemcpyAsync(cache.dTiles.data_ptr<int>(), tiles, (size_t)total_tiles * sizeof(int4), cudaMemcpyHostToDevice, 0));
  }

  hmeta->tiles_ptr = (uint64_t)cache.dTiles.data_ptr();
  hmeta->tiles_count = total_tiles;

  // Meta passed as kernel argument (constant memory) -- no H2D copy needed
  if (amap_dirty) check_cuda(cudaMemcpyAsync(blob_u8 + offA, &cache.A_tbl[0][0], (size_t)G * TMAP_CACHE_CAP * sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0));
  if (bmap_dirty) check_cuda(cudaMemcpyAsync(blob_u8 + offB, &cache.B_tbl[0][0], (size_t)G * TMAP_CACHE_CAP * sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0));

  dim3 grid(total_tiles, 1, 1);
  const int tb = BLOCK_M + 2 * WARP_SIZE;
  const int AB_size = (2 * BLOCK_M + BLOCK_N) * (BLOCK_K / 2);
  const int SFAB_size = 128 * (BLOCK_K / 16) * 3;
  const int smem_size = (AB_size + SFAB_size) * NUM_STAGES;
  if (smem_size > 48'000) {
    static bool smem_attr_set = false;
    if (!smem_attr_set) {
      cudaFuncSetAttribute(grouped_kernel_mtile2, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
      smem_attr_set = true;
    }
  }
  grouped_kernel_mtile2<<<grid, tb, smem_size>>>((const DeviceBlob*)cache.dBlob_u8.data_ptr<uint8_t>(), *hmeta);
}

} // namespace np_mtile2

// --------------------------
// Persistent kernel (bench1)
// --------------------------
namespace persistent {

constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;

constexpr int BLOCK_M = 128;
constexpr int BLOCK_N = 128;
constexpr int BLOCK_K = 256;
constexpr int NUM_STAGES = 6;

constexpr int NUM_SMS_TARGET = 148;

constexpr int ACCUM_STRIDE_TMEM = 128;
constexpr int SCALE_BASE_TMEM   = 2 * ACCUM_STRIDE_TMEM;
constexpr int SFA_TMEM          = SCALE_BASE_TMEM;
constexpr int SFB_TMEM          = SFA_TMEM + 4 * (BLOCK_K / MMA_K);
constexpr int TMEM_ALLOC        = 512;

constexpr uint64_t EVICT_FIRST  = 0x12F0000000000000ULL;
constexpr uint64_t EVICT_LAST   = 0x14F0000000000000ULL;

__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void grouped_kernel_persistent_doublebuf_noreinit(const DeviceBlob *blob, const Meta kmeta) {
  const Meta *meta = &kmeta;
  const int tid = threadIdx.x;
  const int lane_id = tid % WARP_SIZE;
  const int warp_id = tid / WARP_SIZE;
  constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

  extern __shared__ __align__(1024) char smem_ptr[];
  const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
  constexpr int A_size = BLOCK_M * BLOCK_K / 2;
  constexpr int B_size = BLOCK_N * BLOCK_K / 2;
  constexpr int SFA_size = 128 * BLOCK_K / 16;
  constexpr int SFB_size = 128 * BLOCK_K / 16;
  constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;

  #pragma nv_diag_suppress static_var_with_dynamic_init
  __shared__ int64_t mbars[NUM_STAGES * 2 + 2];
  const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
  const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
  const int mainloop0_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;
  const int mainloop1_mbar_addr = mainloop0_mbar_addr + 8;

  if (warp_id == 1) {
    asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(TMEM_ALLOC));
  }

  if (warp_id == NUM_WARPS - 2 && elect_sync()) {
    for (int i = 0; i < NUM_STAGES * 2 + 2; i++) mbarrier_init(tma_mbar_addr + i * 8, 1);
    asm volatile("fence.mbarrier_init.release.cluster;");
  }
  // Only warps 1 (tmem alloc), 4 (mbarrier init), 5 (mma) must rendezvous here.
  if (warp_id == 1 || warp_id == NUM_WARPS - 2 || warp_id == NUM_WARPS - 1) {
    asm volatile("bar.sync 2, %0;" :: "r"(96) : "memory");
  }
  
  int iter = 0;
  int global_iter_base = 0;
  int4 tile_prev;

  for (int tile_id = (int)blockIdx.x; tile_id < meta->tiles_count; tile_id += (int)gridDim.x, iter++) {
    const int cur_buf = (iter & 1);
    const int cur_d_tmem = cur_buf * ACCUM_STRIDE_TMEM;
    const int cur_mainloop_mbar = (cur_buf == 0) ? mainloop0_mbar_addr : mainloop1_mbar_addr;

    if (iter && tid < BLOCK_M) {
      const int prev_iter = iter - 1;
      const int pbuf = (prev_iter & 1);
      const int prev_d_tmem = pbuf * ACCUM_STRIDE_TMEM;
      const int prev_mainloop_mbar = (pbuf == 0) ? mainloop0_mbar_addr : mainloop1_mbar_addr;
      const int prev_seq = (prev_iter >> 1);
      const int prev_phase = (prev_seq & 1);

      const int4 tile_p = tile_prev;
      const int group_p = tile_p.x;
      const int off_m_p = tile_p.y;
      const int off_n_p = tile_p.z;
      const int M_p = meta->M[group_p];
      const int N_p = meta->N[group_p];
      half *C_ptr_p = reinterpret_cast<half *>(meta->C[group_p]);

      mbarrier_wait(prev_mainloop_mbar, prev_phase);
      asm volatile("tcgen05.fence::after_thread_sync;");

      const bool full_tile = (off_m_p + BLOCK_M <= M_p) && (off_n_p + BLOCK_N <= N_p);
      if (full_tile) {
        #pragma unroll
        for (int m0 = 0; m0 < 32 / 16; m0++) {
          #pragma unroll
          for (int half_idx = 0; half_idx < 2; half_idx++) {
            #pragma unroll
            for (int i = 0; i < 64 / 8; i += 4) {
              float tmp[4 * 4];
              tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, prev_d_tmem + half_idx * 64 + i * 8);
              asm volatile("tcgen05.wait::ld.sync.aligned;");
              const int row = off_m_p + warp_id * 32 + m0 * 16 + lane_id / 4;
              #pragma unroll
              for (int ii = 0; ii < 4; ii++) {
                const int col = off_n_p + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
                store_cs_half2(C_ptr_p + (row + 0) * N_p + col, tmp[ii * 4 + 0], tmp[ii * 4 + 1]);
                store_cs_half2(C_ptr_p + (row + 8) * N_p + col, tmp[ii * 4 + 2], tmp[ii * 4 + 3]);
              }
            }
          }
        }
      } else {
        #pragma unroll
        for (int m0 = 0; m0 < 32 / 16; m0++) {
          #pragma unroll
          for (int half_idx = 0; half_idx < 2; half_idx++) {
            #pragma unroll
            for (int i = 0; i < 64 / 8; i += 4) {
              float tmp[4 * 4];
              tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, prev_d_tmem + half_idx * 64 + i * 8);
              asm volatile("tcgen05.wait::ld.sync.aligned;");
              const int row = off_m_p + warp_id * 32 + m0 * 16 + lane_id / 4;
              if (row >= M_p) continue;
              #pragma unroll
              for (int ii = 0; ii < 4; ii++) {
                const int col = off_n_p + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
                store_cs_half2(C_ptr_p + (row + 0) * N_p + col, tmp[ii * 4 + 0], tmp[ii * 4 + 1]);
                if (row + 8 < M_p) {
                  store_cs_half2(C_ptr_p + (row + 8) * N_p + col, tmp[ii * 4 + 2], tmp[ii * 4 + 3]);
                }
              }
            }
          }
        }
      }
    }

    const int4 tile_s = reinterpret_cast<const int4 *>(meta->tiles_ptr)[tile_id];
    const int group = tile_s.x;
    const int off_m = tile_s.y;
    const int off_n = tile_s.z;
    const int sfb_lane = tile_s.w;
    const int M = meta->M[group];
    const int N = meta->N[group];
    const int K = meta->K[group];

    const int a_slot = (int)meta->A_slot[group];
    const int b_slot = (int)meta->B_slot[group];
    const CUtensorMap *A_tmap = blob->A + group * TMAP_CACHE_CAP + a_slot;
    const CUtensorMap *B_tmap = blob->B + group * TMAP_CACHE_CAP + b_slot;
    if (warp_id == 0 && elect_sync()) {
      // (from nvfp4_dual_gemm/gaunernst.py) prefetch tensor maps early
      asm volatile("prefetch.tensormap [%0];" :: "l"(A_tmap) : "memory");
      asm volatile("prefetch.tensormap [%0];" :: "l"(B_tmap) : "memory");
    }
    const char *SFA_ptr = reinterpret_cast<const char *>(meta->SFA[group]);
    const char *SFB_ptr = reinterpret_cast<const char *>(meta->SFB[group]);

    const int num_iters = K / BLOCK_K;
    const int rest_k = K / 64;

    uint64_t cache_A, cache_B;
    if (M > N) { cache_A = EVICT_FIRST; cache_B = EVICT_LAST; }
    else { cache_A = EVICT_LAST; cache_B = EVICT_FIRST; }

    if (warp_id == NUM_WARPS - 2 && elect_sync()) {
      const int tileA = off_m >> 7;
      const int tileB = off_n >> 7;
      const char *SFA_base = SFA_ptr + (tileA * rest_k) * 512;
      const char *SFB_base = SFB_ptr + (tileB * rest_k) * 512;

      auto issue_tma = [&](int iter_k, int stage_id, int tma_phase) {
        const int mbar_addr = tma_mbar_addr + stage_id * 8;
        const int A_smem = smem + stage_id * STAGE_SIZE;
        const int B_smem = A_smem + A_size;
        const int SFA_smem = B_smem + B_size;
        const int SFB_smem = SFA_smem + SFA_size;

        tma_3d_gmem2smem(B_smem, B_tmap, 0, off_n, iter_k, mbar_addr, cache_B);
        tma_3d_gmem2smem(A_smem, A_tmap, 0, off_m, iter_k, mbar_addr, cache_A);

        const int sf_byte = iter_k << 11;
        tma_gmem2smem(SFB_smem, SFB_base + sf_byte, SFB_size, mbar_addr, cache_B);
        tma_gmem2smem(SFA_smem, SFA_base + sf_byte, SFA_size, mbar_addr, cache_A);

        asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
                    :: "r"(mbar_addr), "r"(STAGE_SIZE) : "memory");
      };

      for (int iter_k = 0; iter_k < num_iters; iter_k++) {
        const int giter = global_iter_base + iter_k;
        const int stage_id = giter % NUM_STAGES;
        const int group_phase = giter / NUM_STAGES;
        const int tma_phase = (group_phase & 1);
        if (giter >= NUM_STAGES) {
          const int mma_phase = ((group_phase - 1) & 1);
          mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
        }
        issue_tma(iter_k, stage_id, tma_phase);
      }
    } else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
      constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)BLOCK_N >> 3U << 17U) | ((uint32_t)128 >> 7U << 27U);
      const int scaleA_base = SFA_TMEM;
      const int scaleB_base = SFB_TMEM + sfb_lane;

      for (int iter_k = 0; iter_k < num_iters; iter_k++) {
        const int giter = global_iter_base + iter_k;
        const int stage_id = giter % NUM_STAGES;
        const int tma_phase = (giter / NUM_STAGES) & 1;

        mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);

        const int A_smem = smem + stage_id * STAGE_SIZE;
        const int B_smem = A_smem + A_size;
        const int SFA_smem = B_smem + B_size;
        const int SFB_smem = SFA_smem + SFA_size;

        auto make_desc_AB = [](int addr) -> uint64_t {
          const int SBO = 8 * 128;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
        };
        auto make_desc_SF = [](int addr) -> uint64_t {
          const int SBO = 8 * 16;
          return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
        };

        constexpr uint64_t SF_desc = make_desc_SF(0);
        const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);
        const uint64_t SFB_desc = SF_desc + ((uint64_t)SFB_smem >> 4ULL);

        #pragma unroll
        for (int k = 0; k < BLOCK_K / MMA_K; k++) {
          tcgen05_cp_nvfp4(SFA_TMEM + k * 4, SFA_desc + (uint64_t)k * (512ULL >> 4ULL));
          tcgen05_cp_nvfp4(SFB_TMEM + k * 4, SFB_desc + (uint64_t)k * (512ULL >> 4ULL));
        }

        #pragma unroll
        for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
          const uint64_t a_desc = make_desc_AB(A_smem + k2 * 32);
          const uint64_t b_desc = make_desc_AB(B_smem + k2 * 32);
          const int enable_input_d = (k2 == 0) ? iter_k : 1;
          tcgen05_mma_nvfp4(cur_d_tmem, a_desc, b_desc, i_desc, scaleA_base + k2 * 4, scaleB_base + k2 * 4, enable_input_d);
        }

        asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                    :: "r"(mma_mbar_addr + stage_id * 8) : "memory");
      }

      asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
                  :: "r"(cur_mainloop_mbar) : "memory");
    }
    tile_prev = tile_s;
    global_iter_base += num_iters;
  }

  if (iter && tid < BLOCK_M) {
    const int prev_iter = iter - 1;
    const int pbuf = (prev_iter & 1);
    const int prev_d_tmem = pbuf * ACCUM_STRIDE_TMEM;
    const int prev_mainloop_mbar = (pbuf == 0) ? mainloop0_mbar_addr : mainloop1_mbar_addr;
    const int prev_seq = (prev_iter >> 1);
    const int prev_phase = (prev_seq & 1);

    const int4 tile_p = tile_prev;
    const int group_p = tile_p.x;
    const int off_m_p = tile_p.y;
    const int off_n_p = tile_p.z;
    const int M_p = meta->M[group_p];
    const int N_p = meta->N[group_p];
    half *C_ptr_p = reinterpret_cast<half *>(meta->C[group_p]);

    mbarrier_wait(prev_mainloop_mbar, prev_phase);
    asm volatile("tcgen05.fence::after_thread_sync;");

    const bool full_tile = (off_m_p + BLOCK_M <= M_p) && (off_n_p + BLOCK_N <= N_p);
    if (full_tile) {
      #pragma unroll
      for (int m0 = 0; m0 < 32 / 16; m0++) {
        #pragma unroll
        for (int half_idx = 0; half_idx < 2; half_idx++) {
          #pragma unroll
          for (int i = 0; i < 64 / 8; i += 4) {
            float tmp[4 * 4];
            tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, prev_d_tmem + half_idx * 64 + i * 8);
            asm volatile("tcgen05.wait::ld.sync.aligned;");
            const int row = off_m_p + warp_id * 32 + m0 * 16 + lane_id / 4;
            #pragma unroll
            for (int ii = 0; ii < 4; ii++) {
              const int col = off_n_p + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
              store_cs_half2(C_ptr_p + (row + 0) * N_p + col, tmp[ii * 4 + 0], tmp[ii * 4 + 1]);
              store_cs_half2(C_ptr_p + (row + 8) * N_p + col, tmp[ii * 4 + 2], tmp[ii * 4 + 3]);
            }
          }
        }
      }
    } else {
      #pragma unroll
      for (int m0 = 0; m0 < 32 / 16; m0++) {
        #pragma unroll
        for (int half_idx = 0; half_idx < 2; half_idx++) {
          #pragma unroll
          for (int i = 0; i < 64 / 8; i += 4) {
            float tmp[4 * 4];
            tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, prev_d_tmem + half_idx * 64 + i * 8);
            asm volatile("tcgen05.wait::ld.sync.aligned;");
            const int row = off_m_p + warp_id * 32 + m0 * 16 + lane_id / 4;
            if (row >= M_p) continue;
            #pragma unroll
            for (int ii = 0; ii < 4; ii++) {
              const int col = off_n_p + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
              store_cs_half2(C_ptr_p + (row + 0) * N_p + col, tmp[ii * 4 + 0], tmp[ii * 4 + 1]);
              if (row + 8 < M_p) {
                store_cs_half2(C_ptr_p + (row + 8) * N_p + col, tmp[ii * 4 + 2], tmp[ii * 4 + 3]);
              }
            }
          }
        }
      }
    }
  }

  if (tid < BLOCK_M) {
    asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
  }
  if (warp_id == 0) {
    asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_ALLOC));
  }
}

static void group_gemm(
  c10::List<at::Tensor> A_list,
  c10::List<at::Tensor> B_list,
  c10::List<at::Tensor> C_list,
  c10::List<at::Tensor> SFA_list,
  c10::List<at::Tensor> SFB_list,
  at::Tensor problem_sizes
) {
  const int64_t G = A_list.size();
  // Fast path: assume inputs satisfy task constraints (CPU int32 problem_sizes, 1..8 groups, all CUDA tensors).

  struct Cache {
    bool inited = false;
    int lastM[8] = {};
    int lastN[8] = {};
    int lastK[8] = {};
    int lastTilesN[8] = {};
    int A_size[8] = {};
    int A_next[8] = {};
    int B_size[8] = {};
    int B_next[8] = {};
    uint64_t A_ptr[8][TMAP_CACHE_CAP] = {};
    uint64_t B_ptr[8][TMAP_CACHE_CAP] = {};
    uint64_t A_ht_key[8][PTR_HT_CAP] = {};
    uint8_t  A_ht_val[8][PTR_HT_CAP] = {};
    uint64_t B_ht_key[8][PTR_HT_CAP] = {};
    uint8_t  B_ht_val[8][PTR_HT_CAP] = {};
    CUtensorMap A_tbl[8][TMAP_CACHE_CAP];
    CUtensorMap B_tbl[8][TMAP_CACHE_CAP];
    CUtensorMap A_template[8];
    CUtensorMap B_template[8];
    at::Tensor dBlob_u8;
    enum { META_RING = 256 };
    void *hMeta[META_RING] = {};
    int meta_idx = 0;
    at::Tensor dTiles;
    void *hTiles = nullptr;
    int tiles_cap = 0;
    int last_total_tiles = -1;
  };
  thread_local Cache cache;

  if (!cache.inited) {
    auto opts_u8  = at::TensorOptions().dtype(at::kByte).device(at::kCUDA);
    cache.dBlob_u8 = at::empty({(int64_t)sizeof(DeviceBlob)}, opts_u8);
    for (int b = 0; b < Cache::META_RING; b++) check_cuda(cudaHostAlloc(&cache.hMeta[b], sizeof(Meta), cudaHostAllocPortable));
    cache.inited = true;
  }

  uint8_t *blob_u8 = cache.dBlob_u8.data_ptr<uint8_t>();
  const size_t offA = offsetof(DeviceBlob, A);
  const size_t offB = offsetof(DeviceBlob, B);
  const size_t offM = offsetof(DeviceBlob, meta);

  void *hmeta_buf = cache.hMeta[cache.meta_idx];
  cache.meta_idx = (cache.meta_idx + 1) & (Cache::META_RING - 1);
  Meta *hmeta = reinterpret_cast<Meta *>(hmeta_buf);
  hmeta->offsets[0] = 0;
  hmeta->num_groups = (int)G;

  const int *ps_ptr = problem_sizes.data_ptr<int>();

  bool amap_dirty = false;
  bool bmap_dirty = false;
  bool tiles_dirty = false;

  for (int i = 0; i < (int)G; i++) {
    const int M = ps_ptr[i * 4 + 0];
    const int N = ps_ptr[i * 4 + 1];
    const int K = ps_ptr[i * 4 + 2];
    hmeta->M[i] = M; hmeta->N[i] = N; hmeta->K[i] = K;

    auto A = A_list.get(i);
    auto B = B_list.get(i);
    auto C = C_list.get(i);
    auto SFA = SFA_list.get(i);
    auto SFB = SFB_list.get(i);

    const uint64_t Ap = (uint64_t)A.data_ptr();
    const uint64_t Bp = (uint64_t)B.data_ptr();

    hmeta->C[i] = (uint64_t)C.data_ptr();
    hmeta->SFA[i] = (uint64_t)SFA.data_ptr();
    hmeta->SFB[i] = (uint64_t)SFB.data_ptr();

    const int tiles_m = (M + BLOCK_M - 1) / BLOCK_M;
    const int tiles_n = (N + BLOCK_N - 1) / BLOCK_N;
    hmeta->offsets[i + 1] = hmeta->offsets[i] + tiles_m * tiles_n;
    const bool shape_changed = (cache.lastM[i] != M) || (cache.lastN[i] != N) || (cache.lastK[i] != K);
    if (shape_changed || cache.lastTilesN[i] != tiles_n) {
      cache.lastTilesN[i] = tiles_n;
      tiles_dirty = true;
    }

    if (shape_changed) {
      cache.lastM[i] = M;
      cache.lastN[i] = N;
      cache.lastK[i] = K;
      cache.A_size[i] = 0; cache.A_next[i] = 0;
      cache.B_size[i] = 0; cache.B_next[i] = 0;
      for (int t = 0; t < PTR_HT_CAP; t++) { cache.A_ht_key[i][t] = 0; cache.B_ht_key[i][t] = 0; }
      init_AB_tmap(&cache.A_template[i], (const void*)Ap, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K);
      init_AB_tmap(&cache.B_template[i], (const void*)Bp, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K);
      // Seed slot 0 for both A/B with the current pointers.
      cache.A_size[i] = 1; cache.A_next[i] = 1;
      cache.A_ptr[i][0] = Ap;
      cache.A_tbl[i][0] = cache.A_template[i];
      cache.B_size[i] = 1; cache.B_next[i] = 1;
      cache.B_ptr[i][0] = Bp;
      cache.B_tbl[i][0] = cache.B_template[i];
      ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap, (uint8_t)0);
      ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp, (uint8_t)0);
      amap_dirty = true;
      bmap_dirty = true;
    }

    int a_slot = ht_find<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap);
    if (a_slot >= 0 && cache.A_ptr[i][a_slot] != Ap) {
      for (int t = 0; t < PTR_HT_CAP; t++) cache.A_ht_key[i][t] = 0;
      for (int s = 0; s < cache.A_size[i]; s++) {
        const uint64_t p = cache.A_ptr[i][s];
        if (p) ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], p, (uint8_t)s);
      }
      a_slot = ht_find<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap);
    }
    if (a_slot < 0) {
      a_slot = cache.A_next[i];
      cache.A_next[i] = (cache.A_next[i] + 1) % TMAP_CACHE_CAP;
      if (cache.A_size[i] < TMAP_CACHE_CAP) cache.A_size[i]++;
      cache.A_ptr[i][a_slot] = Ap;
      cache.A_tbl[i][a_slot] = cache.A_template[i];
      check_cu(cuTensorMapReplaceAddress(&cache.A_tbl[i][a_slot], (void*)Ap));
      ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap, (uint8_t)a_slot);
      amap_dirty = true;
    }
    hmeta->A_slot[i] = (uint8_t)a_slot;

    int b_slot = ht_find<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp);
    if (b_slot >= 0 && cache.B_ptr[i][b_slot] != Bp) {
      for (int t = 0; t < PTR_HT_CAP; t++) cache.B_ht_key[i][t] = 0;
      for (int s = 0; s < cache.B_size[i]; s++) {
        const uint64_t p = cache.B_ptr[i][s];
        if (p) ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], p, (uint8_t)s);
      }
      b_slot = ht_find<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp);
    }
    if (b_slot < 0) {
      b_slot = cache.B_next[i];
      cache.B_next[i] = (cache.B_next[i] + 1) % TMAP_CACHE_CAP;
      if (cache.B_size[i] < TMAP_CACHE_CAP) cache.B_size[i]++;
      cache.B_ptr[i][b_slot] = Bp;
      cache.B_tbl[i][b_slot] = cache.B_template[i];
      check_cu(cuTensorMapReplaceAddress(&cache.B_tbl[i][b_slot], (void*)Bp));
      ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp, (uint8_t)b_slot);
      bmap_dirty = true;
    }
    hmeta->B_slot[i] = (uint8_t)b_slot;
  }

  const int total_tiles = hmeta->offsets[G];
  if (total_tiles == 0) return;
  if (total_tiles != cache.last_total_tiles) {
    cache.last_total_tiles = total_tiles;
    tiles_dirty = true;
  }

  if (tiles_dirty) {
    if (total_tiles > cache.tiles_cap) {
      auto opts_i32 = at::TensorOptions().dtype(at::kInt).device(at::kCUDA);
      cache.dTiles = at::empty({(int64_t)total_tiles, 4}, opts_i32);
      if (cache.hTiles) check_cuda(cudaFreeHost(cache.hTiles));
      check_cuda(cudaHostAlloc(&cache.hTiles, (size_t)total_tiles * sizeof(int4), cudaHostAllocPortable));
      cache.tiles_cap = total_tiles;
    }
    auto *tiles = reinterpret_cast<int4 *>(cache.hTiles);

    constexpr int GROUP_SIZE_M = 4;
    int order[8];
    for (int i = 0; i < (int)G; i++) order[i] = i;
    for (int i = 1; i < (int)G; i++) {
      const int key = order[i];
      const int keyK = hmeta->K[key];
      int j = i - 1;
      while (j >= 0 && hmeta->K[order[j]] < keyK) {
        order[j + 1] = order[j];
        j--;
      }
      order[j + 1] = key;
    }

    int t = 0;
    for (int oi = 0; oi < (int)G; oi++) {
      const int g = order[oi];
      const int M = hmeta->M[g];
      const int N = hmeta->N[g];
      const int tiles_m = (M + BLOCK_M - 1) / BLOCK_M;
      const int tiles_n = (N + BLOCK_N - 1) / BLOCK_N;

      for (int first_tm = 0; first_tm < tiles_m; first_tm += GROUP_SIZE_M) {
        const int group_size_m = min(GROUP_SIZE_M, tiles_m - first_tm);
        for (int tn = 0; tn < tiles_n; tn++) {
          for (int i = 0; i < group_size_m; i++) {
            const int tm = first_tm + i;
            const int off_m = tm * BLOCK_M;
            const int off_n = tn * BLOCK_N;
            const int sfb_lane = 0;
            tiles[t++] = make_int4(g, off_m, off_n, sfb_lane);
          }
        }
      }
    }
    check_cuda(cudaMemcpyAsync(cache.dTiles.data_ptr<int>(), tiles, (size_t)total_tiles * sizeof(int4), cudaMemcpyHostToDevice, 0));
  }

  hmeta->tiles_ptr = (uint64_t)cache.dTiles.data_ptr();
  hmeta->tiles_count = total_tiles;

  // Meta passed as kernel argument (constant memory) -- no H2D copy needed
  if (amap_dirty) check_cuda(cudaMemcpyAsync(blob_u8 + offA, &cache.A_tbl[0][0], (size_t)G * TMAP_CACHE_CAP * sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0));
  if (bmap_dirty) check_cuda(cudaMemcpyAsync(blob_u8 + offB, &cache.B_tbl[0][0], (size_t)G * TMAP_CACHE_CAP * sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0));

  int grid_x = NUM_SMS_TARGET;
  if (grid_x > total_tiles) grid_x = total_tiles;
  dim3 grid(grid_x, 1, 1);

  const int tb = BLOCK_M + 2 * WARP_SIZE;
  const int AB_size = (BLOCK_M + BLOCK_N) * (BLOCK_K / 2);
  const int SFAB_size = 128 * (BLOCK_K / 16) * 2;
  const int smem_size = (AB_size + SFAB_size) * NUM_STAGES;
  if (smem_size > 48'000) {
    static bool smem_attr_set = false;
    if (!smem_attr_set) {
      cudaFuncSetAttribute(grouped_kernel_persistent_doublebuf_noreinit, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
      smem_attr_set = true;
    }
  }
  grouped_kernel_persistent_doublebuf_noreinit<<<grid, tb, smem_size>>>((const DeviceBlob*)cache.dBlob_u8.data_ptr<uint8_t>(), *hmeta);
}

} // namespace persistent

// --------------------------
// C++ dispatch
// --------------------------
static void dispatch_group_gemm(
  c10::List<at::Tensor> A_list,
  c10::List<at::Tensor> B_list,
  c10::List<at::Tensor> C_list,
  c10::List<at::Tensor> SFA_list,
  c10::List<at::Tensor> SFB_list,
  at::Tensor problem_sizes
) {
  const int64_t G = A_list.size();
  const int *ps_ptr = problem_sizes.data_ptr<int>();
  const int N0 = ps_ptr[1];
  const int K0 = ps_ptr[2];
  const int L0 = ps_ptr[3];

  // Route all G==8 cases to persistent kernel (enough tiles to fill 148 SMs)
  if (G == 8 && L0 == 1) {
    persistent::group_gemm(A_list, B_list, C_list, SFA_list, SFB_list, problem_sizes);
    return;
  }
  np_base::group_gemm(A_list, B_list, C_list, SFA_list, SFB_list, problem_sizes);
}

TORCH_LIBRARY(nvfp4_group_gemm_variant_combined, m) {
  m.def("dispatch_group_gemm(Tensor[] A, Tensor[] B, Tensor[] C, Tensor[] SFA, Tensor[] SFB, Tensor problem_sizes) -> ()");
  m.impl("dispatch_group_gemm", &dispatch_group_gemm);
}
"""


LIB_NAME = "nvfp4_group_gemm_variant_combined"
EXT_NAME = "nvfp4_group_gemm_variant_combined_ext"

_EXT: torch.nn.Module | None = None

# Cache CPU-side problem_sizes tensors to reduce per-call overhead.
_Key = Tuple[Tuple[int, int, int, int], ...]
_HOST_PS: Dict[_Key, torch.Tensor] = {}


def _cpu_problem_sizes(problem_sizes: List[tuple[int, int, int, int]]) -> torch.Tensor:
    sig: _Key = tuple(tuple(sz) for sz in problem_sizes)
    cached = _HOST_PS.get(sig)
    if cached is None:
        cached = torch.tensor(problem_sizes, dtype=torch.int32, device="cpu")
        _HOST_PS[sig] = cached
    return cached


def _maybe_build() -> torch.nn.Module:
    global _EXT
    mod = _EXT
    if mod is None:
        mod = load_inline(
            name=EXT_NAME,
            cpp_sources=CPP_SRC,
            cuda_sources=[CUDA_SRC],
            functions=None,
            extra_cuda_cflags=[
                "-O3",
                "--use_fast_math",
                "--expt-relaxed-constexpr",
                "--extra-device-vectorization",
                "--relocatable-device-code=false",
                "-gencode=arch=compute_100a,code=sm_100a",
                "-Xptxas=-v",
                "-lineinfo",
            ],
            extra_ldflags=["-lcuda"],
            with_cuda=True,
            verbose=False,
        )
        _EXT = mod
    return mod


def _split_inputs(data: input_t):
    abc_pack, _sf_cpu, sf_pack, dims = data
    g = len(dims)
    a = []
    b = []
    c = []
    sfa = []
    sfb = []
    for i in range(g):
        ai, bi, ci = abc_pack[i]
        sfa_i, sfb_i = sf_pack[i]
        a.append(ai)
        b.append(bi)
        c.append(ci)
        sfa.append(sfa_i)
        sfb.append(sfb_i)
    return a, b, c, sfa, sfb, dims


def custom_kernel(data: input_t) -> output_t:
    a_list, b_list, c_list, sfa_list, sfb_list, problem_sizes = _split_inputs(data)
    _maybe_build()
    ps = _cpu_problem_sizes(problem_sizes)
    ops = getattr(torch.ops, LIB_NAME)
    ops.dispatch_group_gemm(a_list, b_list, c_list, sfa_list, sfb_list, ps)
    return c_list

scrolls · 1993 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 486954.

⋯ 12 unchanged lines
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {}
"""
- # Clean, fully-materialized single-file CUDA source for the current best kernel:
- # - Epilogue granularity combo: persistent ldnum=2, np_base ldnum=4, np_mtile2 ldnum=4
- # - TMA issue order like gaunernst: B -> A -> SFB -> SFA (all kernels)
CUDA_SRC = r"""
#include <torch/types.h>
#include <cuda.h>
⋯ 136 unchanged lines
tcgen05_ld<num * 4, SHAPE::_16x256b, num>(tmp, row, col);
}
+ __device__ __forceinline__ void store_cs_half2(half *ptr, float f0, float f1) {
+ half2 val = __floats2half2_rn(f0, f1);
+ asm volatile("st.cs.b32 [%0], %1;" :: "l"(ptr), "r"(*(uint32_t*)&val) : "memory");
+ }
+
static void check_cu(CUresult err) {
if (err == CUDA_SUCCESS) return;
const char *error_msg_ptr = nullptr;
cuGetErrorString(err, &error_msg_ptr);
- TORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", (error_msg_ptr ? error_msg_ptr : "unknown"));
+ TORCH_CHECK(false, "cuTensorMap error: ", (error_msg_ptr ? error_msg_ptr : "unknown"));
}
static void check_cuda(cudaError_t err) {
⋯ 2 unchanged lines
}
// Shared meta types
- constexpr int TMAP_CACHE_CAP = 32;
+ constexpr int TMAP_CACHE_CAP = 64;
+ constexpr int PTR_HT_CAP = 128;
+ static inline uint32_t hash_u64(uint64_t x) {
+ // MurmurHash3 finalizer
+ x ^= x >> 33;
+ x *= 0xff51afd7ed558ccdULL;
+ x ^= x >> 33;
+ x *= 0xc4ceb9fe1a85ec53ULL;
+ x ^= x >> 33;
+ return (uint32_t)x;
+ }
+
+ template <int HT_CAP>
+ static inline int ht_find(const uint64_t *keys, const uint8_t *vals, uint64_t key) {
+ const uint32_t mask = (uint32_t)HT_CAP - 1U;
+ uint32_t idx = hash_u64(key) & mask;
+ #pragma unroll 1
+ for (int p = 0; p < HT_CAP; p++) {
+ const uint64_t k = keys[idx];
+ if (k == key) return (int)vals[idx];
+ if (k == 0) return -1;
+ idx = (idx + 1U) & mask;
+ }
+ return -1;
+ }
+
+ template <int HT_CAP>
+ static inline void ht_insert(uint64_t *keys, uint8_t *vals, uint64_t key, uint8_t val) {
+ const uint32_t mask = (uint32_t)HT_CAP - 1U;
+ uint32_t idx = hash_u64(key) & mask;
+ #pragma unroll 1
+ for (int p = 0; p < HT_CAP; p++) {
+ const uint64_t k = keys[idx];
+ if (k == 0 || k == key) {
+ keys[idx] = key;
+ vals[idx] = val;
+ return;
+ }
+ idx = (idx + 1U) & mask;
+ }
+ }
+
struct __align__(16) Meta {
uint64_t C[8];
uint64_t SFA[8];
⋯ 61 unchanged lines
constexpr uint64_t EVICT_LAST = 0x14F0000000000000ULL;
__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
- void grouped_kernel(const DeviceBlob *blob) {
- const Meta *meta = &blob->meta;
+ void grouped_kernel(const DeviceBlob *blob, const Meta kmeta) {
+ const Meta *meta = &kmeta;
const int tid = threadIdx.x;
const int bid = blockIdx.x;
const int lane_id = tid % WARP_SIZE;
⋯ 26 unchanged lines
} else if (warp_id == 1) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(BLOCK_N * 2));
}
- if (tid < 64) asm volatile("bar.sync 2, %0;" :: "r"(64) : "memory");
+ // Only warps {0(init), 1(alloc), 4(TMA), 5(MMA)} must rendezvous here.
+ if (warp_id == 0 || warp_id == 1 || warp_id == NUM_WARPS - 2 || warp_id == NUM_WARPS - 1) {
+ asm volatile("bar.sync 2, %0;" :: "r"(BLOCK_M) : "memory");
+ }
const int group = tile_s.x;
⋯ 135 unchanged lines
const float v01 = tmp[ii * 4 + 1];
const float v10 = tmp[ii * 4 + 2];
const float v11 = tmp[ii * 4 + 3];
- reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __floats2half2_rn(v00, v01);
- reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __floats2half2_rn(v10, v11);
+ store_cs_half2(C_ptr + (row + 0) * N + col, v00, v01);
+ store_cs_half2(C_ptr + (row + 8) * N + col, v10, v11);
}
}
}
⋯ 15 unchanged lines
const int col = off_n + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
const float v00 = tmp[ii * 4 + 0];
const float v01 = tmp[ii * 4 + 1];
- reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __floats2half2_rn(v00, v01);
+ store_cs_half2(C_ptr + (row + 0) * N + col, v00, v01);
if (row + 8 < M) {
const float v10 = tmp[ii * 4 + 2];
const float v11 = tmp[ii * 4 + 3];
- reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __floats2half2_rn(v10, v11);
+ store_cs_half2(C_ptr + (row + 8) * N + col, v10, v11);
}
}
}
⋯ 15 unchanged lines
at::Tensor problem_sizes
) {
const int64_t G = A_list.size();
- TORCH_CHECK(G == B_list.size() && G == C_list.size() && G == SFA_list.size() && G == SFB_list.size(), "group list sizes mismatch");
- TORCH_CHECK(problem_sizes.device().is_cpu(), "problem_sizes must be a CPU tensor");
- TORCH_CHECK(problem_sizes.scalar_type() == at::kInt && problem_sizes.dim() == 2 && problem_sizes.size(0) == G && problem_sizes.size(1) == 4,
- "problem_sizes must be int32 CPU tensor of shape [G,4]");
- TORCH_CHECK(G >= 1 && G <= 8, "expected 1..8 groups");
+ // Fast path: assume inputs satisfy task constraints (CPU int32 problem_sizes, 1..8 groups, all CUDA tensors).
struct Cache {
bool inited = false;
⋯ 1 unchanged lines
int lastN[8] = {};
int lastK[8] = {};
int lastTilesN[8] = {};
- struct TMapEntry {
- uint64_t ptr = 0;
- int dim0 = 0;
- int dim1 = 0;
- CUtensorMap map;
- };
- struct TMapRing {
- int size = 0;
- int next = 0;
- TMapEntry e[TMAP_CACHE_CAP];
- };
- TMapRing A_cache[8];
- TMapRing B_cache[8];
+ int A_size[8] = {};
+ int A_next[8] = {};
+ int B_size[8] = {};
+ int B_next[8] = {};
+ uint64_t A_ptr[8][TMAP_CACHE_CAP] = {};
+ uint64_t B_ptr[8][TMAP_CACHE_CAP] = {};
+ uint64_t A_ht_key[8][PTR_HT_CAP] = {};
+ uint8_t A_ht_val[8][PTR_HT_CAP] = {};
+ uint64_t B_ht_key[8][PTR_HT_CAP] = {};
+ uint8_t B_ht_val[8][PTR_HT_CAP] = {};
+ CUtensorMap A_tbl[8][TMAP_CACHE_CAP];
+ CUtensorMap B_tbl[8][TMAP_CACHE_CAP];
+ CUtensorMap A_template[8];
+ CUtensorMap B_template[8];
at::Tensor dBlob_u8;
enum { META_RING = 256 };
void *hMeta[META_RING] = {};
⋯ 23 unchanged lines
hmeta->offsets[0] = 0;
hmeta->num_groups = (int)G;
- auto ps = problem_sizes.contiguous();
- const int *ps_ptr = ps.data_ptr<int>();
+ const int *ps_ptr = problem_sizes.data_ptr<int>();
+ bool amap_dirty = false;
+ bool bmap_dirty = false;
bool tiles_dirty = false;
for (int i = 0; i < (int)G; i++) {
⋯ 8 unchanged lines
auto SFA = SFA_list.get(i);
auto SFB = SFB_list.get(i);
- TORCH_CHECK(A.is_cuda() && B.is_cuda() && C.is_cuda() && SFA.is_cuda() && SFB.is_cuda(), "all tensors must be CUDA");
const uint64_t Ap = (uint64_t)A.data_ptr();
const uint64_t Bp = (uint64_t)B.data_ptr();
⋯ 9 unchanged lines
cache.lastTilesN[i] = tiles_n;
tiles_dirty = true;
}
+
if (shape_changed) {
cache.lastM[i] = M;
cache.lastN[i] = N;
cache.lastK[i] = K;
- cache.A_cache[i].size = 0; cache.A_cache[i].next = 0;
- cache.B_cache[i].size = 0; cache.B_cache[i].next = 0;
+ cache.A_size[i] = 0; cache.A_next[i] = 0;
+ cache.B_size[i] = 0; cache.B_next[i] = 0;
+ for (int t = 0; t < PTR_HT_CAP; t++) { cache.A_ht_key[i][t] = 0; cache.B_ht_key[i][t] = 0; }
+ init_AB_tmap(&cache.A_template[i], (const void*)Ap, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K);
+ init_AB_tmap(&cache.B_template[i], (const void*)Bp, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K);
+ // Seed slot 0 for both A/B with the current pointers.
+ cache.A_size[i] = 1; cache.A_next[i] = 1;
+ cache.A_ptr[i][0] = Ap;
+ cache.A_tbl[i][0] = cache.A_template[i];
+ cache.B_size[i] = 1; cache.B_next[i] = 1;
+ cache.B_ptr[i][0] = Bp;
+ cache.B_tbl[i][0] = cache.B_template[i];
+ ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap, (uint8_t)0);
+ ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp, (uint8_t)0);
+ amap_dirty = true;
+ bmap_dirty = true;
}
- int a_slot = -1;
- auto &ac = cache.A_cache[i];
- for (int e = 0; e < ac.size; e++) {
- const auto &ent = ac.e[e];
- if (ent.ptr == Ap && ent.dim0 == M && ent.dim1 == K) { a_slot = e; break; }
+ int a_slot = ht_find<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap);
+ if (a_slot >= 0 && cache.A_ptr[i][a_slot] != Ap) {
+ // Stale mapping (eviction/wrap): rebuild hash table for this group and retry once.
+ for (int t = 0; t < PTR_HT_CAP; t++) cache.A_ht_key[i][t] = 0;
+ for (int s = 0; s < cache.A_size[i]; s++) {
+ const uint64_t p = cache.A_ptr[i][s];
+ if (p) ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], p, (uint8_t)s);
+ }
+ a_slot = ht_find<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap);
}
if (a_slot < 0) {
- a_slot = ac.next;
- auto &ent = ac.e[a_slot];
- init_AB_tmap(&ent.map, (const void*)Ap, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K);
- ent.ptr = Ap; ent.dim0 = M; ent.dim1 = K;
- ac.next = (ac.next + 1) % TMAP_CACHE_CAP;
- if (ac.size < TMAP_CACHE_CAP) ac.size++;
- check_cuda(cudaMemcpyAsync(
- blob_u8 + offA + (size_t)(i * TMAP_CACHE_CAP + a_slot) * sizeof(CUtensorMap),
- &ent.map, sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0
- ));
+ a_slot = cache.A_next[i];
+ cache.A_next[i] = (cache.A_next[i] + 1) % TMAP_CACHE_CAP;
+ if (cache.A_size[i] < TMAP_CACHE_CAP) cache.A_size[i]++;
+ cache.A_ptr[i][a_slot] = Ap;
+ cache.A_tbl[i][a_slot] = cache.A_template[i];
+ check_cu(cuTensorMapReplaceAddress(&cache.A_tbl[i][a_slot], (void*)Ap));
+ ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap, (uint8_t)a_slot);
+ amap_dirty = true;
}
hmeta->A_slot[i] = (uint8_t)a_slot;
- int b_slot = -1;
- auto &bc = cache.B_cache[i];
- for (int e = 0; e < bc.size; e++) {
- const auto &ent = bc.e[e];
- if (ent.ptr == Bp && ent.dim0 == N && ent.dim1 == K) { b_slot = e; break; }
+ int b_slot = ht_find<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp);
+ if (b_slot >= 0 && cache.B_ptr[i][b_slot] != Bp) {
+ for (int t = 0; t < PTR_HT_CAP; t++) cache.B_ht_key[i][t] = 0;
+ for (int s = 0; s < cache.B_size[i]; s++) {
+ const uint64_t p = cache.B_ptr[i][s];
+ if (p) ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], p, (uint8_t)s);
+ }
+ b_slot = ht_find<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp);
}
if (b_slot < 0) {
- b_slot = bc.next;
- auto &ent = bc.e[b_slot];
- init_AB_tmap(&ent.map, (const void*)Bp, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K);
- ent.ptr = Bp; ent.dim0 = N; ent.dim1 = K;
- bc.next = (bc.next + 1) % TMAP_CACHE_CAP;
- if (bc.size < TMAP_CACHE_CAP) bc.size++;
- check_cuda(cudaMemcpyAsync(
- blob_u8 + offB + (size_t)(i * TMAP_CACHE_CAP + b_slot) * sizeof(CUtensorMap),
- &ent.map, sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0
- ));
+ b_slot = cache.B_next[i];
+ cache.B_next[i] = (cache.B_next[i] + 1) % TMAP_CACHE_CAP;
+ if (cache.B_size[i] < TMAP_CACHE_CAP) cache.B_size[i]++;
+ cache.B_ptr[i][b_slot] = Bp;
+ cache.B_tbl[i][b_slot] = cache.B_template[i];
+ check_cu(cuTensorMapReplaceAddress(&cache.B_tbl[i][b_slot], (void*)Bp));
+ ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp, (uint8_t)b_slot);
+ bmap_dirty = true;
}
hmeta->B_slot[i] = (uint8_t)b_slot;
}
⋯ 56 unchanged lines
hmeta->tiles_ptr = (uint64_t)cache.dTiles.data_ptr();
hmeta->tiles_count = total_tiles;
- check_cuda(cudaMemcpyAsync(blob_u8 + offM, hmeta, sizeof(Meta), cudaMemcpyHostToDevice, 0));
+ // Meta passed as kernel argument (constant memory) -- no H2D copy needed
+ if (amap_dirty) check_cuda(cudaMemcpyAsync(blob_u8 + offA, &cache.A_tbl[0][0], (size_t)G * TMAP_CACHE_CAP * sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0));
+ if (bmap_dirty) check_cuda(cudaMemcpyAsync(blob_u8 + offB, &cache.B_tbl[0][0], (size_t)G * TMAP_CACHE_CAP * sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0));
dim3 grid(total_tiles, 1, 1);
const int tb = BLOCK_M + 2 * WARP_SIZE;
⋯ 7 unchanged lines
smem_attr_set = true;
}
}
- grouped_kernel<<<grid, tb, smem_size>>>((const DeviceBlob*)cache.dBlob_u8.data_ptr<uint8_t>());
+ grouped_kernel<<<grid, tb, smem_size>>>((const DeviceBlob*)cache.dBlob_u8.data_ptr<uint8_t>(), *hmeta);
}
} // namespace np_base
⋯ 20 unchanged lines
constexpr uint64_t EVICT_LAST = 0x14F0000000000000ULL;
__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
- void grouped_kernel_mtile2(const DeviceBlob *blob) {
- const Meta *meta = &blob->meta;
+ void grouped_kernel_mtile2(const DeviceBlob *blob, const Meta kmeta) {
+ const Meta *meta = &kmeta;
const int tid = threadIdx.x;
const int bid = blockIdx.x;
const int lane_id = tid % WARP_SIZE;
⋯ 23 unchanged lines
} else if (warp_id == 1) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(TMEM_ALLOC));
}
- if (tid < 64) asm volatile("bar.sync 2, %0;" :: "r"(64) : "memory");
+ // Only warps {0(init), 1(alloc), 4(TMA), 5(MMA)} must rendezvous here.
+ if (warp_id == 0 || warp_id == 1 || warp_id == NUM_WARPS - 2 || warp_id == NUM_WARPS - 1) {
+ asm volatile("bar.sync 2, %0;" :: "r"(BLOCK_M) : "memory");
+ }
const int group = tile_s.x;
const int off_m = tile_s.y;
⋯ 176 unchanged lines
const float v01 = tmp[ii * 4 + 1];
const float v10 = tmp[ii * 4 + 2];
const float v11 = tmp[ii * 4 + 3];
- reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __floats2half2_rn(v00, v01);
- reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __floats2half2_rn(v10, v11);
+ store_cs_half2(C_ptr + (row + 0) * N + col, v00, v01);
+ store_cs_half2(C_ptr + (row + 8) * N + col, v10, v11);
}
}
}
⋯ 15 unchanged lines
const int col = off_n + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
const float v00 = tmp[ii * 4 + 0];
const float v01 = tmp[ii * 4 + 1];
- reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __floats2half2_rn(v00, v01);
+ store_cs_half2(C_ptr + (row + 0) * N + col, v00, v01);
if (row + 8 < M) {
const float v10 = tmp[ii * 4 + 2];
const float v11 = tmp[ii * 4 + 3];
- reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __floats2half2_rn(v10, v11);
+ store_cs_half2(C_ptr + (row + 8) * N + col, v10, v11);
}
}
}
⋯ 22 unchanged lines
const float v01 = tmp[ii * 4 + 1];
const float v10 = tmp[ii * 4 + 2];
const float v11 = tmp[ii * 4 + 3];
- reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __floats2half2_rn(v00, v01);
- reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __floats2half2_rn(v10, v11);
+ store_cs_half2(C_ptr + (row + 0) * N + col, v00, v01);
+ store_cs_half2(C_ptr + (row + 8) * N + col, v10, v11);
}
}
}
⋯ 15 unchanged lines
const int col = off_n + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
const float v00 = tmp[ii * 4 + 0];
const float v01 = tmp[ii * 4 + 1];
- reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __floats2half2_rn(v00, v01);
+ store_cs_half2(C_ptr + (row + 0) * N + col, v00, v01);
if (row + 8 < M) {
const float v10 = tmp[ii * 4 + 2];
const float v11 = tmp[ii * 4 + 3];
- reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __floats2half2_rn(v10, v11);
+ store_cs_half2(C_ptr + (row + 8) * N + col, v10, v11);
}
}
}
⋯ 16 unchanged lines
at::Tensor problem_sizes
) {
const int64_t G = A_list.size();
- TORCH_CHECK(G == B_list.size() && G == C_list.size() && G == SFA_list.size() && G == SFB_list.size(), "group list sizes mismatch");
- TORCH_CHECK(problem_sizes.device().is_cpu(), "problem_sizes must be a CPU tensor");
- TORCH_CHECK(problem_sizes.scalar_type() == at::kInt && problem_sizes.dim() == 2 && problem_sizes.size(0) == G && problem_sizes.size(1) == 4,
- "problem_sizes must be int32 CPU tensor of shape [G,4]");
- TORCH_CHECK(G >= 1 && G <= 8, "expected 1..8 groups");
+ // Fast path: assume inputs satisfy task constraints (CPU int32 problem_sizes, 1..8 groups, all CUDA tensors).
struct Cache {
bool inited = false;
⋯ 1 unchanged lines
int lastN[8] = {};
int lastK[8] = {};
int lastTilesN[8] = {};
- struct TMapEntry {
- uint64_t ptr = 0;
- int dim0 = 0;
- int dim1 = 0;
- CUtensorMap map;
- };
- struct TMapRing {
- int size = 0;
- int next = 0;
- TMapEntry e[TMAP_CACHE_CAP];
- };
- TMapRing A_cache[8];
- TMapRing B_cache[8];
+ int A_size[8] = {};
+ int A_next[8] = {};
+ int B_size[8] = {};
+ int B_next[8] = {};
+ uint64_t A_ptr[8][TMAP_CACHE_CAP] = {};
+ uint64_t B_ptr[8][TMAP_CACHE_CAP] = {};
+ uint64_t A_ht_key[8][PTR_HT_CAP] = {};
+ uint8_t A_ht_val[8][PTR_HT_CAP] = {};
+ uint64_t B_ht_key[8][PTR_HT_CAP] = {};
+ uint8_t B_ht_val[8][PTR_HT_CAP] = {};
+ CUtensorMap A_tbl[8][TMAP_CACHE_CAP];
+ CUtensorMap B_tbl[8][TMAP_CACHE_CAP];
+ CUtensorMap A_template[8];
+ CUtensorMap B_template[8];
at::Tensor dBlob_u8;
enum { META_RING = 256 };
void *hMeta[META_RING] = {};
⋯ 23 unchanged lines
hmeta->offsets[0] = 0;
hmeta->num_groups = (int)G;
- auto ps = problem_sizes.contiguous();
- const int *ps_ptr = ps.data_ptr<int>();
+ const int *ps_ptr = problem_sizes.data_ptr<int>();
+ bool amap_dirty = false;
+ bool bmap_dirty = false;
bool tiles_dirty = false;
for (int i = 0; i < (int)G; i++) {
⋯ 8 unchanged lines
auto SFA = SFA_list.get(i);
auto SFB = SFB_list.get(i);
- TORCH_CHECK(A.is_cuda() && B.is_cuda() && C.is_cuda() && SFA.is_cuda() && SFB.is_cuda(), "all tensors must be CUDA");
const uint64_t Ap = (uint64_t)A.data_ptr();
const uint64_t Bp = (uint64_t)B.data_ptr();
⋯ 10 unchanged lines
cache.lastTilesN[i] = tiles_n;
tiles_dirty = true;
}
+
if (shape_changed) {
cache.lastM[i] = M;
cache.lastN[i] = N;
cache.lastK[i] = K;
- cache.A_cache[i].size = 0; cache.A_cache[i].next = 0;
- cache.B_cache[i].size = 0; cache.B_cache[i].next = 0;
+ cache.A_size[i] = 0; cache.A_next[i] = 0;
+ cache.B_size[i] = 0; cache.B_next[i] = 0;
+ for (int t = 0; t < PTR_HT_CAP; t++) { cache.A_ht_key[i][t] = 0; cache.B_ht_key[i][t] = 0; }
+ init_AB_tmap(&cache.A_template[i], (const void*)Ap, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K);
+ init_AB_tmap(&cache.B_template[i], (const void*)Bp, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K);
+ // Seed slot 0 for both A/B with the current pointers.
+ cache.A_size[i] = 1; cache.A_next[i] = 1;
+ cache.A_ptr[i][0] = Ap;
+ cache.A_tbl[i][0] = cache.A_template[i];
+ cache.B_size[i] = 1; cache.B_next[i] = 1;
+ cache.B_ptr[i][0] = Bp;
+ cache.B_tbl[i][0] = cache.B_template[i];
+ ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap, (uint8_t)0);
+ ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp, (uint8_t)0);
+ amap_dirty = true;
+ bmap_dirty = true;
}
- int a_slot = -1;
- auto &ac = cache.A_cache[i];
- for (int e = 0; e < ac.size; e++) {
- const auto &ent = ac.e[e];
- if (ent.ptr == Ap && ent.dim0 == M && ent.dim1 == K) { a_slot = e; break; }
+ int a_slot = ht_find<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap);
+ if (a_slot >= 0 && cache.A_ptr[i][a_slot] != Ap) {
+ for (int t = 0; t < PTR_HT_CAP; t++) cache.A_ht_key[i][t] = 0;
+ for (int s = 0; s < cache.A_size[i]; s++) {
+ const uint64_t p = cache.A_ptr[i][s];
+ if (p) ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], p, (uint8_t)s);
+ }
+ a_slot = ht_find<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap);
}
if (a_slot < 0) {
- a_slot = ac.next;
- auto &ent = ac.e[a_slot];
- init_AB_tmap(&ent.map, (const void*)Ap, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K);
- ent.ptr = Ap; ent.dim0 = M; ent.dim1 = K;
- ac.next = (ac.next + 1) % TMAP_CACHE_CAP;
- if (ac.size < TMAP_CACHE_CAP) ac.size++;
- check_cuda(cudaMemcpyAsync(
- blob_u8 + offA + (size_t)(i * TMAP_CACHE_CAP + a_slot) * sizeof(CUtensorMap),
- &ent.map, sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0
- ));
+ a_slot = cache.A_next[i];
+ cache.A_next[i] = (cache.A_next[i] + 1) % TMAP_CACHE_CAP;
+ if (cache.A_size[i] < TMAP_CACHE_CAP) cache.A_size[i]++;
+ cache.A_ptr[i][a_slot] = Ap;
+ cache.A_tbl[i][a_slot] = cache.A_template[i];
+ check_cu(cuTensorMapReplaceAddress(&cache.A_tbl[i][a_slot], (void*)Ap));
+ ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap, (uint8_t)a_slot);
+ amap_dirty = true;
}
hmeta->A_slot[i] = (uint8_t)a_slot;
- int b_slot = -1;
- auto &bc = cache.B_cache[i];
- for (int e = 0; e < bc.size; e++) {
- const auto &ent = bc.e[e];
- if (ent.ptr == Bp && ent.dim0 == N && ent.dim1 == K) { b_slot = e; break; }
+ int b_slot = ht_find<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp);
+ if (b_slot >= 0 && cache.B_ptr[i][b_slot] != Bp) {
+ for (int t = 0; t < PTR_HT_CAP; t++) cache.B_ht_key[i][t] = 0;
+ for (int s = 0; s < cache.B_size[i]; s++) {
+ const uint64_t p = cache.B_ptr[i][s];
+ if (p) ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], p, (uint8_t)s);
+ }
+ b_slot = ht_find<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp);
}
if (b_slot < 0) {
- b_slot = bc.next;
- auto &ent = bc.e[b_slot];
- init_AB_tmap(&ent.map, (const void*)Bp, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K);
- ent.ptr = Bp; ent.dim0 = N; ent.dim1 = K;
- bc.next = (bc.next + 1) % TMAP_CACHE_CAP;
- if (bc.size < TMAP_CACHE_CAP) bc.size++;
- check_cuda(cudaMemcpyAsync(
- blob_u8 + offB + (size_t)(i * TMAP_CACHE_CAP + b_slot) * sizeof(CUtensorMap),
- &ent.map, sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0
- ));
+ b_slot = cache.B_next[i];
+ cache.B_next[i] = (cache.B_next[i] + 1) % TMAP_CACHE_CAP;
+ if (cache.B_size[i] < TMAP_CACHE_CAP) cache.B_size[i]++;
+ cache.B_ptr[i][b_slot] = Bp;
+ cache.B_tbl[i][b_slot] = cache.B_template[i];
+ check_cu(cuTensorMapReplaceAddress(&cache.B_tbl[i][b_slot], (void*)Bp));
+ ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp, (uint8_t)b_slot);
+ bmap_dirty = true;
}
hmeta->B_slot[i] = (uint8_t)b_slot;
}
⋯ 57 unchanged lines
hmeta->tiles_ptr = (uint64_t)cache.dTiles.data_ptr();
hmeta->tiles_count = total_tiles;
- check_cuda(cudaMemcpyAsync(blob_u8 + offM, hmeta, sizeof(Meta), cudaMemcpyHostToDevice, 0));
+ // Meta passed as kernel argument (constant memory) -- no H2D copy needed
+ if (amap_dirty) check_cuda(cudaMemcpyAsync(blob_u8 + offA, &cache.A_tbl[0][0], (size_t)G * TMAP_CACHE_CAP * sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0));
+ if (bmap_dirty) check_cuda(cudaMemcpyAsync(blob_u8 + offB, &cache.B_tbl[0][0], (size_t)G * TMAP_CACHE_CAP * sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0));
dim3 grid(total_tiles, 1, 1);
const int tb = BLOCK_M + 2 * WARP_SIZE;
⋯ 7 unchanged lines
smem_attr_set = true;
}
}
- grouped_kernel_mtile2<<<grid, tb, smem_size>>>((const DeviceBlob*)cache.dBlob_u8.data_ptr<uint8_t>());
+ grouped_kernel_mtile2<<<grid, tb, smem_size>>>((const DeviceBlob*)cache.dBlob_u8.data_ptr<uint8_t>(), *hmeta);
}
} // namespace np_mtile2
⋯ 9 unchanged lines
constexpr int BLOCK_M = 128;
constexpr int BLOCK_N = 128;
constexpr int BLOCK_K = 256;
- constexpr int NUM_STAGES = 4;
+ constexpr int NUM_STAGES = 6;
- constexpr int NUM_SMS_TARGET = 128;
+ constexpr int NUM_SMS_TARGET = 148;
constexpr int ACCUM_STRIDE_TMEM = 128;
constexpr int SCALE_BASE_TMEM = 2 * ACCUM_STRIDE_TMEM;
⋯ 5 unchanged lines
constexpr uint64_t EVICT_LAST = 0x14F0000000000000ULL;
__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
- void grouped_kernel_persistent_doublebuf_noreinit(const DeviceBlob *blob) {
- const Meta *meta = &blob->meta;
+ void grouped_kernel_persistent_doublebuf_noreinit(const DeviceBlob *blob, const Meta kmeta) {
+ const Meta *meta = &kmeta;
const int tid = threadIdx.x;
const int lane_id = tid % WARP_SIZE;
const int warp_id = tid / WARP_SIZE;
⋯ 22 unchanged lines
for (int i = 0; i < NUM_STAGES * 2 + 2; i++) mbarrier_init(tma_mbar_addr + i * 8, 1);
asm volatile("fence.mbarrier_init.release.cluster;");
}
- // asm volatile("bar.sync 1, %0;" :: "r"(64) : "memory");
- __syncthreads();
- __shared__ int4 prev_tile_s;
- __shared__ int had_any;
- if (tid == 0) had_any = 0;
- if (tid < BLOCK_M) {
- asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
+ // Only warps 1 (tmem alloc), 4 (mbarrier init), 5 (mma) must rendezvous here.
+ if (warp_id == 1 || warp_id == NUM_WARPS - 2 || warp_id == NUM_WARPS - 1) {
+ asm volatile("bar.sync 2, %0;" :: "r"(96) : "memory");
}
int iter = 0;
int global_iter_base = 0;
+ int4 tile_prev;
for (int tile_id = (int)blockIdx.x; tile_id < meta->tiles_count; tile_id += (int)gridDim.x, iter++) {
const int cur_buf = (iter & 1);
const int cur_d_tmem = cur_buf * ACCUM_STRIDE_TMEM;
const int cur_mainloop_mbar = (cur_buf == 0) ? mainloop0_mbar_addr : mainloop1_mbar_addr;
- if (had_any && tid < BLOCK_M) {
+ if (iter && tid < BLOCK_M) {
const int prev_iter = iter - 1;
const int pbuf = (prev_iter & 1);
const int prev_d_tmem = pbuf * ACCUM_STRIDE_TMEM;
⋯ 1 unchanged lines
const int prev_seq = (prev_iter >> 1);
const int prev_phase = (prev_seq & 1);
- const int4 tile_p = prev_tile_s;
+ const int4 tile_p = tile_prev;
const int group_p = tile_p.x;
const int off_m_p = tile_p.y;
const int off_n_p = tile_p.z;
⋯ 11 unchanged lines
#pragma unroll
for (int half_idx = 0; half_idx < 2; half_idx++) {
#pragma unroll
- for (int i = 0; i < 64 / 8; i += 2) {
- float tmp[2 * 4];
- tcgen05_ld_16x256b<2>(tmp, warp_id * 32 + m0 * 16, prev_d_tmem + half_idx * 64 + i * 8);
+ for (int i = 0; i < 64 / 8; i += 4) {
+ float tmp[4 * 4];
+ tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, prev_d_tmem + half_idx * 64 + i * 8);
asm volatile("tcgen05.wait::ld.sync.aligned;");
const int row = off_m_p + warp_id * 32 + m0 * 16 + lane_id / 4;
#pragma unroll
- for (int ii = 0; ii < 2; ii++) {
+ for (int ii = 0; ii < 4; ii++) {
const int col = off_n_p + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
- reinterpret_cast<half2 *>(C_ptr_p + (row + 0) * N_p + col)[0] = __floats2half2_rn(tmp[ii * 4 + 0], tmp[ii * 4 + 1]);
- reinterpret_cast<half2 *>(C_ptr_p + (row + 8) * N_p + col)[0] = __floats2half2_rn(tmp[ii * 4 + 2], tmp[ii * 4 + 3]);
+ store_cs_half2(C_ptr_p + (row + 0) * N_p + col, tmp[ii * 4 + 0], tmp[ii * 4 + 1]);
+ store_cs_half2(C_ptr_p + (row + 8) * N_p + col, tmp[ii * 4 + 2], tmp[ii * 4 + 3]);
}
}
}
⋯ 4 unchanged lines
#pragma unroll
for (int half_idx = 0; half_idx < 2; half_idx++) {
#pragma unroll
- for (int i = 0; i < 64 / 8; i += 2) {
- float tmp[2 * 4];
- tcgen05_ld_16x256b<2>(tmp, warp_id * 32 + m0 * 16, prev_d_tmem + half_idx * 64 + i * 8);
+ for (int i = 0; i < 64 / 8; i += 4) {
+ float tmp[4 * 4];
+ tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, prev_d_tmem + half_idx * 64 + i * 8);
asm volatile("tcgen05.wait::ld.sync.aligned;");
const int row = off_m_p + warp_id * 32 + m0 * 16 + lane_id / 4;
if (row >= M_p) continue;
#pragma unroll
- for (int ii = 0; ii < 2; ii++) {
+ for (int ii = 0; ii < 4; ii++) {
const int col = off_n_p + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
- reinterpret_cast<half2 *>(C_ptr_p + (row + 0) * N_p + col)[0] = __floats2half2_rn(tmp[ii * 4 + 0], tmp[ii * 4 + 1]);
+ store_cs_half2(C_ptr_p + (row + 0) * N_p + col, tmp[ii * 4 + 0], tmp[ii * 4 + 1]);
if (row + 8 < M_p) {
- reinterpret_cast<half2 *>(C_ptr_p + (row + 8) * N_p + col)[0] = __floats2half2_rn(tmp[ii * 4 + 2], tmp[ii * 4 + 3]);
+ store_cs_half2(C_ptr_p + (row + 8) * N_p + col, tmp[ii * 4 + 2], tmp[ii * 4 + 3]);
}
}
}
⋯ 56 unchanged lines
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
const int giter = global_iter_base + iter_k;
- const int stage_id = giter & (NUM_STAGES - 1);
- const int group_phase = (giter >> 2);
+ const int stage_id = giter % NUM_STAGES;
+ const int group_phase = giter / NUM_STAGES;
const int tma_phase = (group_phase & 1);
if (giter >= NUM_STAGES) {
const int mma_phase = ((group_phase - 1) & 1);
⋯ 8 unchanged lines
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
const int giter = global_iter_base + iter_k;
- const int stage_id = giter & (NUM_STAGES - 1);
- const int tma_phase = ((giter >> 2) & 1);
+ const int stage_id = giter % NUM_STAGES;
+ const int tma_phase = (giter / NUM_STAGES) & 1;
mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);
⋯ 36 unchanged lines
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(cur_mainloop_mbar) : "memory");
}
-
- if (tid == 0) {
- prev_tile_s = tile_s;
- had_any = 1;
- }
- if (tid < BLOCK_M) {
- asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
- }
+ tile_prev = tile_s;
global_iter_base += num_iters;
}
- if (had_any && tid < BLOCK_M) {
+ if (iter && tid < BLOCK_M) {
const int prev_iter = iter - 1;
const int pbuf = (prev_iter & 1);
const int prev_d_tmem = pbuf * ACCUM_STRIDE_TMEM;
⋯ 1 unchanged lines
const int prev_seq = (prev_iter >> 1);
const int prev_phase = (prev_seq & 1);
- const int4 tile_p = prev_tile_s;
+ const int4 tile_p = tile_prev;
const int group_p = tile_p.x;
const int off_m_p = tile_p.y;
const int off_n_p = tile_p.z;
⋯ 11 unchanged lines
#pragma unroll
for (int half_idx = 0; half_idx < 2; half_idx++) {
#pragma unroll
- for (int i = 0; i < 64 / 8; i += 2) {
- float tmp[2 * 4];
- tcgen05_ld_16x256b<2>(tmp, warp_id * 32 + m0 * 16, prev_d_tmem + half_idx * 64 + i * 8);
+ for (int i = 0; i < 64 / 8; i += 4) {
+ float tmp[4 * 4];
+ tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, prev_d_tmem + half_idx * 64 + i * 8);
asm volatile("tcgen05.wait::ld.sync.aligned;");
const int row = off_m_p + warp_id * 32 + m0 * 16 + lane_id / 4;
#pragma unroll
- for (int ii = 0; ii < 2; ii++) {
+ for (int ii = 0; ii < 4; ii++) {
const int col = off_n_p + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
- reinterpret_cast<half2 *>(C_ptr_p + (row + 0) * N_p + col)[0] = __floats2half2_rn(tmp[ii * 4 + 0], tmp[ii * 4 + 1]);
- reinterpret_cast<half2 *>(C_ptr_p + (row + 8) * N_p + col)[0] = __floats2half2_rn(tmp[ii * 4 + 2], tmp[ii * 4 + 3]);
+ store_cs_half2(C_ptr_p + (row + 0) * N_p + col, tmp[ii * 4 + 0], tmp[ii * 4 + 1]);
+ store_cs_half2(C_ptr_p + (row + 8) * N_p + col, tmp[ii * 4 + 2], tmp[ii * 4 + 3]);
}
}
}
⋯ 4 unchanged lines
#pragma unroll
for (int half_idx = 0; half_idx < 2; half_idx++) {
#pragma unroll
- for (int i = 0; i < 64 / 8; i += 2) {
- float tmp[2 * 4];
- tcgen05_ld_16x256b<2>(tmp, warp_id * 32 + m0 * 16, prev_d_tmem + half_idx * 64 + i * 8);
+ for (int i = 0; i < 64 / 8; i += 4) {
+ float tmp[4 * 4];
+ tcgen05_ld_16x256b<4>(tmp, warp_id * 32 + m0 * 16, prev_d_tmem + half_idx * 64 + i * 8);
asm volatile("tcgen05.wait::ld.sync.aligned;");
const int row = off_m_p + warp_id * 32 + m0 * 16 + lane_id / 4;
if (row >= M_p) continue;
#pragma unroll
- for (int ii = 0; ii < 2; ii++) {
+ for (int ii = 0; ii < 4; ii++) {
const int col = off_n_p + half_idx * 64 + (i + ii) * 8 + (lane_id % 4) * 2;
- reinterpret_cast<half2 *>(C_ptr_p + (row + 0) * N_p + col)[0] = __floats2half2_rn(tmp[ii * 4 + 0], tmp[ii * 4 + 1]);
+ store_cs_half2(C_ptr_p + (row + 0) * N_p + col, tmp[ii * 4 + 0], tmp[ii * 4 + 1]);
if (row + 8 < M_p) {
- reinterpret_cast<half2 *>(C_ptr_p + (row + 8) * N_p + col)[0] = __floats2half2_rn(tmp[ii * 4 + 2], tmp[ii * 4 + 3]);
+ store_cs_half2(C_ptr_p + (row + 8) * N_p + col, tmp[ii * 4 + 2], tmp[ii * 4 + 3]);
}
}
}
⋯ 19 unchanged lines
at::Tensor problem_sizes
) {
const int64_t G = A_list.size();
- TORCH_CHECK(G == B_list.size() && G == C_list.size() && G == SFA_list.size() && G == SFB_list.size(), "group list sizes mismatch");
- TORCH_CHECK(problem_sizes.device().is_cpu(), "problem_sizes must be a CPU tensor");
- TORCH_CHECK(problem_sizes.scalar_type() == at::kInt && problem_sizes.dim() == 2 && problem_sizes.size(0) == G && problem_sizes.size(1) == 4,
- "problem_sizes must be int32 CPU tensor of shape [G,4]");
- TORCH_CHECK(G >= 1 && G <= 8, "expected 1..8 groups");
+ // Fast path: assume inputs satisfy task constraints (CPU int32 problem_sizes, 1..8 groups, all CUDA tensors).
struct Cache {
bool inited = false;
⋯ 1 unchanged lines
int lastN[8] = {};
int lastK[8] = {};
int lastTilesN[8] = {};
- struct TMapEntry {
- uint64_t ptr = 0;
- int dim0 = 0;
- int dim1 = 0;
- CUtensorMap map;
- };
- struct TMapRing {
- int size = 0;
- int next = 0;
- TMapEntry e[TMAP_CACHE_CAP];
- };
- TMapRing A_cache[8];
- TMapRing B_cache[8];
+ int A_size[8] = {};
+ int A_next[8] = {};
+ int B_size[8] = {};
+ int B_next[8] = {};
+ uint64_t A_ptr[8][TMAP_CACHE_CAP] = {};
+ uint64_t B_ptr[8][TMAP_CACHE_CAP] = {};
+ uint64_t A_ht_key[8][PTR_HT_CAP] = {};
+ uint8_t A_ht_val[8][PTR_HT_CAP] = {};
+ uint64_t B_ht_key[8][PTR_HT_CAP] = {};
+ uint8_t B_ht_val[8][PTR_HT_CAP] = {};
+ CUtensorMap A_tbl[8][TMAP_CACHE_CAP];
+ CUtensorMap B_tbl[8][TMAP_CACHE_CAP];
+ CUtensorMap A_template[8];
+ CUtensorMap B_template[8];
at::Tensor dBlob_u8;
enum { META_RING = 256 };
void *hMeta[META_RING] = {};
⋯ 23 unchanged lines
hmeta->offsets[0] = 0;
hmeta->num_groups = (int)G;
- auto ps = problem_sizes.contiguous();
- const int *ps_ptr = ps.data_ptr<int>();
+ const int *ps_ptr = problem_sizes.data_ptr<int>();
+ bool amap_dirty = false;
+ bool bmap_dirty = false;
bool tiles_dirty = false;
for (int i = 0; i < (int)G; i++) {
⋯ 8 unchanged lines
auto SFA = SFA_list.get(i);
auto SFB = SFB_list.get(i);
- TORCH_CHECK(A.is_cuda() && B.is_cuda() && C.is_cuda() && SFA.is_cuda() && SFB.is_cuda(), "all tensors must be CUDA");
const uint64_t Ap = (uint64_t)A.data_ptr();
const uint64_t Bp = (uint64_t)B.data_ptr();
⋯ 9 unchanged lines
cache.lastTilesN[i] = tiles_n;
tiles_dirty = true;
}
+
if (shape_changed) {
cache.lastM[i] = M;
cache.lastN[i] = N;
cache.lastK[i] = K;
- cache.A_cache[i].size = 0; cache.A_cache[i].next = 0;
- cache.B_cache[i].size = 0; cache.B_cache[i].next = 0;
+ cache.A_size[i] = 0; cache.A_next[i] = 0;
+ cache.B_size[i] = 0; cache.B_next[i] = 0;
+ for (int t = 0; t < PTR_HT_CAP; t++) { cache.A_ht_key[i][t] = 0; cache.B_ht_key[i][t] = 0; }
+ init_AB_tmap(&cache.A_template[i], (const void*)Ap, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K);
+ init_AB_tmap(&cache.B_template[i], (const void*)Bp, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K);
+ // Seed slot 0 for both A/B with the current pointers.
+ cache.A_size[i] = 1; cache.A_next[i] = 1;
+ cache.A_ptr[i][0] = Ap;
+ cache.A_tbl[i][0] = cache.A_template[i];
+ cache.B_size[i] = 1; cache.B_next[i] = 1;
+ cache.B_ptr[i][0] = Bp;
+ cache.B_tbl[i][0] = cache.B_template[i];
+ ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap, (uint8_t)0);
+ ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp, (uint8_t)0);
+ amap_dirty = true;
+ bmap_dirty = true;
}
- int a_slot = -1;
- auto &ac = cache.A_cache[i];
- for (int e = 0; e < ac.size; e++) {
- const auto &ent = ac.e[e];
- if (ent.ptr == Ap && ent.dim0 == M && ent.dim1 == K) { a_slot = e; break; }
+ int a_slot = ht_find<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap);
+ if (a_slot >= 0 && cache.A_ptr[i][a_slot] != Ap) {
+ for (int t = 0; t < PTR_HT_CAP; t++) cache.A_ht_key[i][t] = 0;
+ for (int s = 0; s < cache.A_size[i]; s++) {
+ const uint64_t p = cache.A_ptr[i][s];
+ if (p) ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], p, (uint8_t)s);
+ }
+ a_slot = ht_find<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap);
}
if (a_slot < 0) {
- a_slot = ac.next;
- auto &ent = ac.e[a_slot];
- init_AB_tmap(&ent.map, (const void*)Ap, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K);
- ent.ptr = Ap; ent.dim0 = M; ent.dim1 = K;
- ac.next = (ac.next + 1) % TMAP_CACHE_CAP;
- if (ac.size < TMAP_CACHE_CAP) ac.size++;
- check_cuda(cudaMemcpyAsync(
- blob_u8 + offA + (size_t)(i * TMAP_CACHE_CAP + a_slot) * sizeof(CUtensorMap),
- &ent.map, sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0
- ));
+ a_slot = cache.A_next[i];
+ cache.A_next[i] = (cache.A_next[i] + 1) % TMAP_CACHE_CAP;
+ if (cache.A_size[i] < TMAP_CACHE_CAP) cache.A_size[i]++;
+ cache.A_ptr[i][a_slot] = Ap;
+ cache.A_tbl[i][a_slot] = cache.A_template[i];
+ check_cu(cuTensorMapReplaceAddress(&cache.A_tbl[i][a_slot], (void*)Ap));
+ ht_insert<PTR_HT_CAP>(cache.A_ht_key[i], cache.A_ht_val[i], Ap, (uint8_t)a_slot);
+ amap_dirty = true;
}
hmeta->A_slot[i] = (uint8_t)a_slot;
- int b_slot = -1;
- auto &bc = cache.B_cache[i];
- for (int e = 0; e < bc.size; e++) {
- const auto &ent = bc.e[e];
- if (ent.ptr == Bp && ent.dim0 == N && ent.dim1 == K) { b_slot = e; break; }
+ int b_slot = ht_find<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp);
+ if (b_slot >= 0 && cache.B_ptr[i][b_slot] != Bp) {
+ for (int t = 0; t < PTR_HT_CAP; t++) cache.B_ht_key[i][t] = 0;
+ for (int s = 0; s < cache.B_size[i]; s++) {
+ const uint64_t p = cache.B_ptr[i][s];
+ if (p) ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], p, (uint8_t)s);
+ }
+ b_slot = ht_find<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp);
}
if (b_slot < 0) {
- b_slot = bc.next;
- auto &ent = bc.e[b_slot];
- init_AB_tmap(&ent.map, (const void*)Bp, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K);
- ent.ptr = Bp; ent.dim0 = N; ent.dim1 = K;
- bc.next = (bc.next + 1) % TMAP_CACHE_CAP;
- if (bc.size < TMAP_CACHE_CAP) bc.size++;
- check_cuda(cudaMemcpyAsync(
- blob_u8 + offB + (size_t)(i * TMAP_CACHE_CAP + b_slot) * sizeof(CUtensorMap),
- &ent.map, sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0
- ));
+ b_slot = cache.B_next[i];
+ cache.B_next[i] = (cache.B_next[i] + 1) % TMAP_CACHE_CAP;
+ if (cache.B_size[i] < TMAP_CACHE_CAP) cache.B_size[i]++;
+ cache.B_ptr[i][b_slot] = Bp;
+ cache.B_tbl[i][b_slot] = cache.B_template[i];
+ check_cu(cuTensorMapReplaceAddress(&cache.B_tbl[i][b_slot], (void*)Bp));
+ ht_insert<PTR_HT_CAP>(cache.B_ht_key[i], cache.B_ht_val[i], Bp, (uint8_t)b_slot);
+ bmap_dirty = true;
}
hmeta->B_slot[i] = (uint8_t)b_slot;
}
⋯ 56 unchanged lines
hmeta->tiles_ptr = (uint64_t)cache.dTiles.data_ptr();
hmeta->tiles_count = total_tiles;
- check_cuda(cudaMemcpyAsync(blob_u8 + offM, hmeta, sizeof(Meta), cudaMemcpyHostToDevice, 0));
+ // Meta passed as kernel argument (constant memory) -- no H2D copy needed
+ if (amap_dirty) check_cuda(cudaMemcpyAsync(blob_u8 + offA, &cache.A_tbl[0][0], (size_t)G * TMAP_CACHE_CAP * sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0));
+ if (bmap_dirty) check_cuda(cudaMemcpyAsync(blob_u8 + offB, &cache.B_tbl[0][0], (size_t)G * TMAP_CACHE_CAP * sizeof(CUtensorMap), cudaMemcpyHostToDevice, 0));
int grid_x = NUM_SMS_TARGET;
if (grid_x > total_tiles) grid_x = total_tiles;
⋯ 10 unchanged lines
smem_attr_set = true;
}
}
- grouped_kernel_persistent_doublebuf_noreinit<<<grid, tb, smem_size>>>((const DeviceBlob*)cache.dBlob_u8.data_ptr<uint8_t>());
+ grouped_kernel_persistent_doublebuf_noreinit<<<grid, tb, smem_size>>>((const DeviceBlob*)cache.dBlob_u8.data_ptr<uint8_t>(), *hmeta);
}
} // namespace persistent
⋯ 10 unchanged lines
at::Tensor problem_sizes
) {
const int64_t G = A_list.size();
- auto ps = problem_sizes.contiguous();
- const int *ps_ptr = ps.data_ptr<int>();
+ const int *ps_ptr = problem_sizes.data_ptr<int>();
const int N0 = ps_ptr[1];
const int K0 = ps_ptr[2];
const int L0 = ps_ptr[3];
- if (G == 8 && N0 == 7168 && K0 == 2048 && L0 == 1) {
+ // Route all G==8 cases to persistent kernel (enough tiles to fill 148 SMs)
+ if (G == 8 && L0 == 1) {
persistent::group_gemm(A_list, B_list, C_list, SFA_list, SFB_list, problem_sizes);
return;
}
- if (G == 8 && N0 == 4096 && K0 == 7168 && L0 == 1) {
- np_mtile2::group_gemm(A_list, B_list, C_list, SFA_list, SFB_list, problem_sizes);
- return;
- }
np_base::group_gemm(A_list, B_list, C_list, SFA_list, SFB_list, problem_sizes);
}
- TORCH_LIBRARY(nvfp4_group_gemm_ext_onefile_cppdispatch_v1_submission_candidate_gpt_epilogue_combo_v24_tmaorder_v27, m) {
+ TORCH_LIBRARY(nvfp4_group_gemm_variant_combined, m) {
m.def("dispatch_group_gemm(Tensor[] A, Tensor[] B, Tensor[] C, Tensor[] SFA, Tensor[] SFB, Tensor problem_sizes) -> ()");
m.impl("dispatch_group_gemm", &dispatch_group_gemm);
}
"""
- LIB_NAME = "nvfp4_group_gemm_ext_onefile_cppdispatch_v1_submission_candidate_gpt_epilogue_combo_v24_tmaorder_v27"
- EXT_NAME = "nvfp4_group_gemm_ext_mod_submission_candidate_gpt_epilogue_combo_v24_tmaorder_v27_clean"
+ LIB_NAME = "nvfp4_group_gemm_variant_combined"
+ EXT_NAME = "nvfp4_group_gemm_variant_combined_ext"
_EXT: torch.nn.Module | None = None
scrolls · 1017 diff lines total

Best evidence level for this revision: reported

JSON