Skip to content
KernelIndex
Search⌘K

submission 486943

dxyz · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-486943?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
90.4µs
#86 of 145
2026-02-08

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:697d43eb7c40c18ae93b0ff6b37fa8822ffe815073e0ca6b3bfa99117ff0e215
license declaredunknown
license concludedunknown
authorsdxyz
imported2026-08-15

Techniques

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

cluster__cluster_dims__(CTA_GROUP, 1, 1)
fp4PyTorch reference implementation of NVFP4 block-scaled group GEMM.
fused-epilogueWaitEpilogue,
mbarriervoid mbarrier_init(int mbar_addr, int count) {
shared-memoryauto asmem_layout = make_layout(make_shape(_256{}, BLOCK_M, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
tcgen05asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
tmaasm volatile("cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;"
vector-width = half2reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 0], tmp[i * 4 + 1]});

Kernel source

submission_v1.py3661 lines
from task import input_t, output_t
from reference import ref_kernel

import torch
from torch.utils.cpp_extension import load_inline

COMMON_CU = r"""
#include <stdio.h>

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

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

#include "cutlass/cutlass.h"

#include "cute/tensor.hpp"

constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;  // 32 bytes
constexpr int SF_BLOCK_SIZE = 16;

#define DIVUP(a, b) (((a) + (b) - 1) / (b))

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

enum ProfilerTag {
  Setup = 0,
  IssueTMA,
  IssueMMA,
  WaitTMA,
  WaitMMA,
  WaitMainloop,
  WaitEpilogue,
  Epilogue,
};

__device__ inline
int64_t globaltimer() {
  int64_t t;
  asm volatile("mov.u64 %0, %globaltimer;" : "=l"(t) :: "memory");
  return t;
}

struct Profiler {
  int64_t *data_ptr_;
  int sm_id_;
  int cnt_;

  __device__
  void init(int num_entries, int64_t *data_ptr, int bid) {
    data_ptr_ = data_ptr + bid * (1 + num_entries * 4);
    asm volatile("mov.u32 %0, %smid;\n" : "=r"(sm_id_));
    cnt_ = 0;
  }

  __device__
  void start(ProfilerTag tag) {
    data_ptr_[1 + cnt_ * 4 + 0] = sm_id_;
    data_ptr_[1 + cnt_ * 4 + 1] = tag;
    data_ptr_[1 + cnt_ * 4 + 2] = globaltimer();
  }

  __device__
  void stop() {
    data_ptr_[1 + cnt_ * 4 + 3] = globaltimer() - data_ptr_[1 + cnt_ * 4 + 2];
    cnt_ += 1;
  }

  __device__
  void flush() {
    data_ptr_[0] = cnt_;
  }
};

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

// https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cute/arch/cluster_sm90.hpp#L180
__device__
uint32_t elect_sync() {
  uint32_t pred = 0;
  asm volatile(
    "{\n\t"
    ".reg .pred %%px;\n\t"
    "elect.sync _|%%px, %1;\n\t"
    "@%%px mov.s32 %0, 1;\n\t"
    "}"
    : "+r"(pred)
    : "r"(0xFFFFFFFF)
  );
  return pred;
}

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

// https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cutlass/arch/barrier.h#L408
__device__
void mbarrier_wait(int mbar_addr, int phase) {
  uint32_t ticks = 0x989680;  // this is optional
  asm volatile(
    "{\n\t"
    ".reg .pred P1;\n\t"
    "LAB_WAIT:\n\t"
    "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t"
    "@P1 bra.uni DONE;\n\t"
    "bra.uni LAB_WAIT;\n\t"
    "DONE:\n\t"
    "}"
    :: "r"(mbar_addr), "r"(phase), "r"(ticks)
  );
}

template <int CTA_GROUP = 1>
__device__ inline
void tcgen05_commit_mcast(int mbar_addr, uint16_t cta_mask) {
  asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
              :: "r"(mbar_addr), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
}

// See https://docs.nvidia.com/cuda/parallel-thread-execution/#data-movement-and-conversion-instructions-bulk-copy
// cp.async.bulk is a non-blocking instruction which initiates an asynchronous bulk-copy operation
// from the location specified by source address operand srcMem 
// to the location specified by destination address operand dstMem.
__device__ inline
void tma_gmem2smem(int dst, const void *src, int size, int mbar_addr, uint64_t cache_policy) {
  asm volatile("cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;"
              :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy));
}

template <int CTA_GROUP = 1>
__device__ inline
void tma_1d_gmem2smem(int dst, const void *tmap_ptr, int x, int mbar_addr, uint64_t cache_policy) {
  asm volatile("cp.async.bulk.tensor.1d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::%5.L2::cache_hint "
              "[%0], [%1, {%2}], [%3], %4;"
              :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(mbar_addr), "l"(cache_policy), "n"(CTA_GROUP)
              : "memory");
}

template <int CTA_GROUP = 1>
__device__ inline
void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy) {
  asm volatile("cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::%7.L2::cache_hint "
              "[%0], [%1, {%2, %3, %4}], [%5], %6;"
              :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "l"(cache_policy), "n"(CTA_GROUP)
              : "memory");
}
// See https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-instructions-tcgen05-cp
// Instruction tcgen05.cp initiates an asynchronous copy operation from shared memory 
// to the location specified by the address operand taddr in the Tensor Memory.
template <int CTA_GROUP = 1>
__device__ inline
void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
  // .32x128b corresponds to lane x size. So we move 32 lanes and 128 bits = 128 / 8 bytes = 16 bytes which is 4 columns
  // x4 means the data is multicast into 4 warps and each warp gets 1/4 of the data.
  // .warpx4 populates data across 32-lane groups (lane in the sense of tmem).
  // Some of the .shape qualifiers require certain .multicast qualifiers.

  // .64x128b requires .warpx2::02_13 or .warpx2::01_23
  // .32x128b requires .warpx4
  //
  // See https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-data-movement-shape
  // 32x128b means 32 lanes, 128b == 16 bytes spans 16 / 4 == 4 columns
  asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc), "n"(CTA_GROUP));
}


struct COLLECTOR_USAGE {
  static constexpr char NONE[]      = "";
  static constexpr char A_FILL[]    = ".collector::a::fill";
  static constexpr char A_USE[]     = ".collector::a::use";
  static constexpr char A_LASTUSE[] = ".collector::a::lastuse";
  static constexpr char A_DISCARD[] = ".collector::a::discard";
};

template <int CTA_GROUP = 1, const char *collector_usage = COLLECTOR_USAGE::NONE>
__device__ inline
void tcgen05_mma_nvfp4(
  int d_tmem,
  uint64_t a_desc,
  uint64_t b_desc,
  uint32_t i_desc,
  int scale_A_tmem,
  int scale_B_tmem,
  int enable_input_d
) {
  asm volatile(
    "{\n\t"
    ".reg .pred p;\n\t"  // predicate register enable-input-d
    "setp.ne.b32 p, %6, 0;\n\t"
    "tcgen05.mma.cta_group::%7.kind::mxf4nvf4.block_scale.block16%8 [%0], %1, %2, %3, [%4], [%5], p;\n\t"
    "}"
    :: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
       "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d),
       "n"(CTA_GROUP), "C"(collector_usage)
  );
}

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

struct NUM {
  static constexpr char x4[]  = ".x4";
  static constexpr char x8[]  = ".x8";
  static constexpr char x16[] = ".x16";
  static constexpr char x32[] = ".x32";
  static constexpr char x64[] = ".x64";
  static constexpr char x128[] = ".x128";
};

template <const char *SHAPE, const char *NUM>
__device__ inline
void tcgen05_ld_16regs(float *tmp, int row, int col) {
  asm volatile("tcgen05.ld.sync.aligned%17%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"((row << 16) | col), "C"(SHAPE), "C"(NUM));
}

template <const char *SHAPE, const char *NUM>
__device__ inline
void tcgen05_ld_32regs(float *tmp, int row, int col) {
  asm volatile("tcgen05.ld.sync.aligned%33%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"((row << 16) | col), "C"(SHAPE), "C"(NUM));
}

template <const char *SHAPE, const char *NUM>
__device__ inline
void tcgen05_ld_64regs(float *tmp, int row, int col) {
  asm volatile("tcgen05.ld.sync.aligned%65%66.b32 "
              "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7, "
              "  %8,  %9, %10, %11, %12, %13, %14, %15, "
              " %16, %17, %18, %19, %20, %21, %22, %23, "
              " %24, %25, %26, %27, %28, %29, %30, %31, "
              " %32, %33, %34, %35, %36, %37, %38, %39, "
              " %40, %41, %42, %43, %44, %45, %46, %47, "
              " %48, %49, %50, %51, %52, %53, %54, %55, "
              " %56, %57, %58, %59, %60, %61, %62, %63}, [%64];"
              : "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
                "=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
                "=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
                "=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31]),
                "=f"(tmp[32]), "=f"(tmp[33]), "=f"(tmp[34]), "=f"(tmp[35]), "=f"(tmp[36]), "=f"(tmp[37]), "=f"(tmp[38]), "=f"(tmp[39]),
                "=f"(tmp[40]), "=f"(tmp[41]), "=f"(tmp[42]), "=f"(tmp[43]), "=f"(tmp[44]), "=f"(tmp[45]), "=f"(tmp[46]), "=f"(tmp[47]),
                "=f"(tmp[48]), "=f"(tmp[49]), "=f"(tmp[50]), "=f"(tmp[51]), "=f"(tmp[52]), "=f"(tmp[53]), "=f"(tmp[54]), "=f"(tmp[55]),
                "=f"(tmp[56]), "=f"(tmp[57]), "=f"(tmp[58]), "=f"(tmp[59]), "=f"(tmp[60]), "=f"(tmp[61]), "=f"(tmp[62]), "=f"(tmp[63])
              : "r"((row << 16) | col), "C"(SHAPE), "C"(NUM));
}

template <const char *SHAPE, const char *NUM>
__device__ inline
void tcgen05_ld_128regs(float *tmp, int row, int col) {
  asm volatile("tcgen05.ld.sync.aligned%129%130.b32 "
              "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7, "
              "  %8,  %9, %10, %11, %12, %13, %14, %15, "
              " %16, %17, %18, %19, %20, %21, %22, %23, "
              " %24, %25, %26, %27, %28, %29, %30, %31, "
              " %32, %33, %34, %35, %36, %37, %38, %39, "
              " %40, %41, %42, %43, %44, %45, %46, %47, "
              " %48, %49, %50, %51, %52, %53, %54, %55, "
              " %56, %57, %58, %59, %60, %61, %62, %63, "
              " %64, %65, %66, %67, %68, %69, %70, %71, "
              " %72, %73, %74, %75, %76, %77, %78, %79, "
              " %80, %81, %82, %83, %84, %85, %86, %87, "
              " %88, %89, %90, %91, %92, %93, %94, %95, "
              " %96, %97, %98, %99,%100,%101,%102,%103, "
              "%104,%105,%106,%107,%108,%109,%110,%111, "
              "%112,%113,%114,%115,%116,%117,%118,%119, "
              "%120,%121,%122,%123,%124,%125,%126,%127}, [%128];"
              : "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
                "=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
                "=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
                "=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31]),
                "=f"(tmp[32]), "=f"(tmp[33]), "=f"(tmp[34]), "=f"(tmp[35]), "=f"(tmp[36]), "=f"(tmp[37]), "=f"(tmp[38]), "=f"(tmp[39]),
                "=f"(tmp[40]), "=f"(tmp[41]), "=f"(tmp[42]), "=f"(tmp[43]), "=f"(tmp[44]), "=f"(tmp[45]), "=f"(tmp[46]), "=f"(tmp[47]),
                "=f"(tmp[48]), "=f"(tmp[49]), "=f"(tmp[50]), "=f"(tmp[51]), "=f"(tmp[52]), "=f"(tmp[53]), "=f"(tmp[54]), "=f"(tmp[55]),
                "=f"(tmp[56]), "=f"(tmp[57]), "=f"(tmp[58]), "=f"(tmp[59]), "=f"(tmp[60]), "=f"(tmp[61]), "=f"(tmp[62]), "=f"(tmp[63]),
                "=f"(tmp[64]), "=f"(tmp[65]), "=f"(tmp[66]), "=f"(tmp[67]), "=f"(tmp[68]), "=f"(tmp[69]), "=f"(tmp[70]), "=f"(tmp[71]),
                "=f"(tmp[72]), "=f"(tmp[73]), "=f"(tmp[74]), "=f"(tmp[75]), "=f"(tmp[76]), "=f"(tmp[77]), "=f"(tmp[78]), "=f"(tmp[79]),
                "=f"(tmp[80]), "=f"(tmp[81]), "=f"(tmp[82]), "=f"(tmp[83]), "=f"(tmp[84]), "=f"(tmp[85]), "=f"(tmp[86]), "=f"(tmp[87]),
                "=f"(tmp[88]), "=f"(tmp[89]), "=f"(tmp[90]), "=f"(tmp[91]), "=f"(tmp[92]), "=f"(tmp[93]), "=f"(tmp[94]), "=f"(tmp[95]),
                "=f"(tmp[96]), "=f"(tmp[97]), "=f"(tmp[98]), "=f"(tmp[99]), "=f"(tmp[100]),"=f"(tmp[101]),"=f"(tmp[102]),"=f"(tmp[103]),
                "=f"(tmp[104]),"=f"(tmp[105]),"=f"(tmp[106]),"=f"(tmp[107]),"=f"(tmp[108]),"=f"(tmp[109]),"=f"(tmp[110]),"=f"(tmp[111]),
                "=f"(tmp[112]),"=f"(tmp[113]),"=f"(tmp[114]),"=f"(tmp[115]),"=f"(tmp[116]),"=f"(tmp[117]),"=f"(tmp[118]),"=f"(tmp[119]),
                "=f"(tmp[120]),"=f"(tmp[121]),"=f"(tmp[122]),"=f"(tmp[123]),"=f"(tmp[124]),"=f"(tmp[125]),"=f"(tmp[126]),"=f"(tmp[127])
              : "r"((row << 16) | col), "C"(SHAPE), "C"(NUM));
}

__device__ inline void tcgen05_ld_32x32bx32(float *tmp, int row, int col) { tcgen05_ld_32regs<SHAPE::_32x32b, NUM::x32>(tmp, row, col); }
__device__ inline void tcgen05_ld_32x32bx64(float *tmp, int row, int col) { tcgen05_ld_64regs<SHAPE::_32x32b, NUM::x64>(tmp, row, col); }
__device__ inline void tcgen05_ld_32x32bx128(float *tmp, int row, int col) { tcgen05_ld_128regs<SHAPE::_32x32b, NUM::x128>(tmp, row, col); }

__device__ inline void tcgen05_ld_16x128bx8(float *tmp, int row, int col) { tcgen05_ld_16regs<SHAPE::_16x128b, NUM::x8>(tmp, row, col); }
__device__ inline void tcgen05_ld_16x128bx16(float *tmp, int row, int col) { tcgen05_ld_32regs<SHAPE::_16x128b, NUM::x16>(tmp, row, col); }
__device__ inline void tcgen05_ld_16x128bx32(float *tmp, int row, int col) { tcgen05_ld_64regs<SHAPE::_16x128b, NUM::x32>(tmp, row, col); }

__device__ inline void tcgen05_ld_16x256bx4(float *tmp, int row, int col) { tcgen05_ld_16regs<SHAPE::_16x256b, NUM::x4>(tmp, row, col); }
__device__ inline void tcgen05_ld_16x256bx8(float *tmp, int row, int col) { tcgen05_ld_32regs<SHAPE::_16x256b, NUM::x8>(tmp, row, col); }
__device__ inline void tcgen05_ld_16x256bx16(float *tmp, int row, int col) { tcgen05_ld_64regs<SHAPE::_16x256b, NUM::x16>(tmp, row, col); }

void check_cu(CUresult err) {
  if (err == CUDA_SUCCESS) return;
  const char *error_msg_ptr;
  if (cuGetErrorString(err, &error_msg_ptr) != CUDA_SUCCESS)
    error_msg_ptr = "unable to get error string";
  TORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", error_msg_ptr);
}

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

void init_AB_tmap(
  CUtensorMap *tmap,
  const char *ptr,
  uint64_t global_height, uint64_t global_width,
  uint32_t shared_height, uint32_t shared_width
) {
  constexpr uint32_t rank = 3;

  // Tile shape (boxDim)            Global stride	Comment
  // [BLOCK_M, BLOCK_K]	            [K, 1]	      Original layout of threadblock tile in global memory.
  // [BLOCK_M, BLOCK_K / 8, 8]	    [K, 8, 1]	    “Unflatten” the last dim.
  // [BLOCK_K / 8, BLOCK_M, 8]	    [8, K, 1]	    Swap the first 2 dims.
  // 8 if for bf16, 256 for nvfp4: we need contiguous strips of 256 elements

  // The order needs to be reversed as cuTensorMapEncodeTiled assumes the tensors are
  // col major.

  // For the smem layout see https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-shared-memory-layout-swizzling
  // See also https://docs.nvidia.com/cuda/parallel-thread-execution/#asynchronous-warpgroup-level-matrix-shared-memory-layout
  //
  // For example, 128B MN major swizzle atom would have a shape of (8*(128/32)) x 8 = 32x8 for tf32 tensor core inputs.
  // So for 4 bits we have (8 * 128 / 4) x 8 = 256 x 8
  // the swizzling pattern is given in the tables.

  uint64_t globalDim[rank]       = {256, global_height, global_width / 256};
  uint64_t globalStrides[rank-1] = {global_width / 2, 128};  // in bytes
  uint32_t boxDim[rank]          = {256, shared_height, shared_width / 256};
  uint32_t elementStrides[rank]  = {1, 1, 1}; // When all elements of elementStrides array is one, boxDim specifies the number of elements to load.

  auto err = cuTensorMapEncodeTiled(
    tmap,
    CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
    rank,
    (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);
}

void init_SF_tmap(CUtensorMap *tmap, const char *ptr, uint64_t global_size, uint32_t shared_size) {
  // use int64 as dtype, hence divide sizes by 8
  constexpr uint32_t rank = 1;
  uint64_t globalDim[rank]       = {global_size / 8};
  uint64_t globalStrides[rank-1] = {};  // in bytes
  uint32_t boxDim[rank]          = {shared_size / 8};
  uint32_t elementStrides[rank]  = {1};

  auto err = cuTensorMapEncodeTiled(
    tmap,
    CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_INT64,
    rank,
    (void *)ptr,
    globalDim,
    globalStrides,
    boxDim,
    elementStrides,
    CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
    CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_NONE,
    CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
    CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
  );
  //check_cu(err);
}

"""

CUDA_SRC_0 = r"""
using namespace cute;

constexpr int MAX_TMEM_COLS = 512;
constexpr int NUM_SMS = 148;

constexpr int NUM_BLOCKS0 = 64;
constexpr int NUM_BLOCKS1 = 96;
constexpr int NUM_BLOCKS2 = 64;
constexpr int NUM_BLOCKS3 = 64;
constexpr int NUM_BLOCKS4 = 32;
constexpr int NUM_BLOCKS5 = 128;
constexpr int NUM_BLOCKS6 = 64;
constexpr int NUM_BLOCKS7 = 96;

constexpr int GEMM_BLOCK_END0 = NUM_BLOCKS0;
constexpr int GEMM_BLOCK_END1 = GEMM_BLOCK_END0 + NUM_BLOCKS1;
constexpr int GEMM_BLOCK_END2 = GEMM_BLOCK_END1 + NUM_BLOCKS2;
constexpr int GEMM_BLOCK_END3 = GEMM_BLOCK_END2 + NUM_BLOCKS3;
constexpr int GEMM_BLOCK_END4 = GEMM_BLOCK_END3 + NUM_BLOCKS4;
constexpr int GEMM_BLOCK_END5 = GEMM_BLOCK_END4 + NUM_BLOCKS5;
constexpr int GEMM_BLOCK_END6 = GEMM_BLOCK_END5 + NUM_BLOCKS6;
constexpr int GEMM_BLOCK_END7 = GEMM_BLOCK_END6 + NUM_BLOCKS7;

constexpr int NUM_BLOCKS = GEMM_BLOCK_END7;

template<int PROBLEM_SIZE>
__constant__ CUtensorMap A_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ CUtensorMap B_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFA_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFB_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ half* C_ptrs[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ int Ms[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ int Ns[PROBLEM_SIZE];

template <
  int K,
  int BLOCK_M,
  int BLOCK_N,
  int BLOCK_K,
  int NUM_STAGES,
  class SFALayout,
  class SFBLayout,
  bool SWAP_AB,
  int PROBLEM_SIZE,
  int CTA_GROUP
>
__global__
__cluster_dims__(CTA_GROUP, 1, 1)
__launch_bounds__(BLOCK_M + 2 * WARP_SIZE)  // __launch_bounds__(MAX_THREADS_PER_BLOCK, MIN_BLOCKS_PER_MP)
void multi_gemm_kernel_p0() {
  // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-matrix-shape
  static_assert(BLOCK_M == 128); // There doesnt seem to be any code for tiling block_M

  // This layout is used as TMA assumes col major, this is a colasced canonical layout.
  auto asmem_layout = make_layout(make_shape(_256{}, BLOCK_M, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
  // Each member of the cluster will get a cta_rank specific offset
  auto bsmem_layout = make_layout(make_shape(_256{}, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_N / CTA_GROUP));

  // an SF atom is ((32, 4), (16, 4)), the shmem layout is ((128 * SF_ATOMS_PER_BLOCK_K, 4)) for TMA purposes
  // Since we are using CU_TENSOR_MAP_SWIZZLE_NONE, the TMA atom has shape  (8, 128bits) or (8, 16) bytes, so SBO below is 8 * 16
  constexpr int SF_TILE_ROWS = 128;
  constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
  constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns
  constexpr int SF_ROWS_PER_TMEM_COL = 32; // each tmem col has 32 rows - since we are using tcgen05.cp.cta_group::1.32x128b.warpx4

  using Problem = std::tuple<const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, half*, int, int, SFALayout, SFBLayout>;

  auto get_problem = [&](int problem_id) -> Problem {
  
    auto make_sf_layout = [](int rest) {
      constexpr int rest_K = K / 16 / 4;
      return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)), 
                                make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
    };

    const CUtensorMap* A_tmap = &A_tmaps<PROBLEM_SIZE>[problem_id];
    const CUtensorMap* B_tmap = &B_tmaps<PROBLEM_SIZE>[problem_id];
    const CUtensorMap* SFA_tmap = &SFA_tmaps<PROBLEM_SIZE>[problem_id];
    const CUtensorMap* SFB_tmap = &SFB_tmaps<PROBLEM_SIZE>[problem_id];

    half* C = C_ptrs<PROBLEM_SIZE>[problem_id];

    int M = Ms<PROBLEM_SIZE>[problem_id];
    int N = Ns<PROBLEM_SIZE>[problem_id];

    const int rest_M = DIVUP(M, SF_TILE_ROWS);
    const int rest_N = DIVUP(N, SF_TILE_ROWS);
    SFALayout sfa_layout = make_sf_layout(rest_M);
    SFBLayout sfb_layout = make_sf_layout(rest_N);

    return {A_tmap, B_tmap, SFA_tmap, SFB_tmap, C, M, N, sfa_layout, sfb_layout};
  };

  // CTA rank in a cluster
  int cta_rank;
  asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));

  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;

  // set up smem
  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; // BLOCK_M * BLOCK_K elements, each element is 4 bits so divide by 2 to get bytes.
  constexpr int B_size = BLOCK_N * BLOCK_K / 2 / CTA_GROUP; // each CTA only loads half of B
  constexpr int SFA_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K;  // always copy one atom - this is ATOM_SIZE_IN_BYTES * ATOMS_PER_BLOCK_K
  constexpr int SFB_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K;
  constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;

  // set up mbarriers and tmem
  // we have NUM_STAGES mbars for TMA
  //         NUM_STAGES mbars for MMA
  //                  1 mbar  for mainloop
  #pragma nv_diag_suppress static_var_with_dynamic_init
  __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
  __shared__ int tmem_addr[1];  // tmem address is 32-bit

  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;
  const int epilogue_mbar_addr = mainloop_mbar_addr + 2 * 8;

  // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
  // each MMA consumes:
  // - (128, 64) of A -> (128, 4) of SFA -> reshaped as (32, 4', 4) -> 4 tmem columns
  constexpr int SFA_tmem = BLOCK_N * 2; // Double buffer the mma output.
  constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K); // we need BLOCK_K / MMA_K many scale factors for the "tiledMMA"

  if (warp_id == 0 && elect_sync()) {
    for (int prefetch_idx = 0; prefetch_idx < 8; ++prefetch_idx) {
      asm volatile("prefetch.tensormap [%0];" :: "l"(&A_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
      asm volatile("prefetch.tensormap [%0];" :: "l"(&B_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
      asm volatile("prefetch.tensormap [%0];" :: "l"(&SFA_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
      asm volatile("prefetch.tensormap [%0];" :: "l"(&SFB_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
    }
  }  else if (warp_id == 1 && elect_sync()) {
    for (int i = 0; i < NUM_STAGES; i++) {
      mbarrier_init(tma_mbar_addr + i * 8, CTA_GROUP); // one thread in each cluster CTA reports.
      mbarrier_init(mma_mbar_addr + i * 8, 1);
    }
    mbarrier_init(mainloop_mbar_addr, 1); // barrier for first tmem mma buffer
    mbarrier_init(mainloop_mbar_addr + 8, 1); // barrier for second tmem mma buffer

    mbarrier_init(epilogue_mbar_addr, 4 * CTA_GROUP);
    mbarrier_init(epilogue_mbar_addr + 8, 4 * CTA_GROUP);

    asm volatile("fence.mbarrier_init.release.cluster;");  // visible to async proxy
  }
  
  if constexpr (CTA_GROUP > 1) {
    // visible to all threads in a cluster
    asm volatile("barrier.cluster.arrive.relaxed.aligned;");
    asm volatile("barrier.cluster.wait.acquire.aligned;");
  }
  else {
    // visible to all threads in a threadblock
    __syncthreads();
  }

  constexpr int num_iters = K / BLOCK_K;
  const int bid = blockIdx.x;

  auto scheduler = [&](int global_cluster_id) -> std::tuple<int, int, int, int> {
    int problem_id = 0;
    int local_cluster_id = global_cluster_id;

    if (global_cluster_id < GEMM_BLOCK_END0 / CTA_GROUP) {
      problem_id = 0;
      local_cluster_id = global_cluster_id;
    } else if (global_cluster_id < GEMM_BLOCK_END1 / CTA_GROUP) {
      problem_id = 1;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END0 / CTA_GROUP;
    } else if (global_cluster_id < GEMM_BLOCK_END2 / CTA_GROUP) {
      problem_id = 2;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END1 / CTA_GROUP;
    } else if (global_cluster_id < GEMM_BLOCK_END3 / CTA_GROUP) {
      problem_id = 3;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END2 / CTA_GROUP;
    } else if (global_cluster_id < GEMM_BLOCK_END4 / CTA_GROUP) {
      problem_id = 4;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END3 / CTA_GROUP;
    } else if (global_cluster_id < GEMM_BLOCK_END5 / CTA_GROUP) {
      problem_id = 5;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END4 / CTA_GROUP;
    } else if (global_cluster_id < GEMM_BLOCK_END6 / CTA_GROUP) {
      problem_id = 6;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END5 / CTA_GROUP;
    } else if (global_cluster_id < GEMM_BLOCK_END7 / CTA_GROUP) {
      problem_id = 7;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END6 / CTA_GROUP;
    }

    int N = Ns<PROBLEM_SIZE>[problem_id];

    const int grid_n = DIVUP(N, BLOCK_N) ;
    // The cluster Ids are arranged in a (grid_m / 2 , grid_n) matrix in a row major layout. 
    const int bid_m = local_cluster_id / grid_n * CTA_GROUP + cta_rank;
    const int bid_n = local_cluster_id % grid_n;

    const int off_m = bid_m * BLOCK_M;
    const int cluster_off_n = bid_n * BLOCK_N; // N offset for this cluster
    const int off_n = cluster_off_n + cta_rank * (BLOCK_N / CTA_GROUP); // N offset for this CTA
    return {off_m, off_n, cluster_off_n, problem_id};
  };

  // warp-specialization
  if (warp_id == NUM_WARPS - 2 && elect_sync()) {
    // TMA warp

    auto issue_tma = [&](int iter_k, int stage_id, int off_m, int off_n, int problem_id) {
      const int mbar_addr = (tma_mbar_addr + stage_id * 8) & 0xFEFFFFFF;  // CTA0's barrier
      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 problem = get_problem(problem_id);

      const CUtensorMap* A_tmap = std::get<0>(problem);
      const CUtensorMap* B_tmap = std::get<1>(problem);    

      const int M = std::get<5>(problem);
      const int N = std::get<6>(problem);

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

      // issue TMA
      const int off_k = iter_k * BLOCK_K;
      tma_3d_gmem2smem<CTA_GROUP>(A_smem, A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
      tma_3d_gmem2smem<CTA_GROUP>(B_smem, B_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);

      const int rest_m = off_m / 32 / 4; // M atom is (32, 4)
      const int rest_n = off_n / 32 / 4; // N atom is (32, 4)

      const int rest_k = off_k / 16 / 4; // The SF atom is ((32, 4), (16, 4))

      const CUtensorMap* SFA_tmap = std::get<2>(problem);
      const CUtensorMap* SFB_tmap = std::get<3>(problem);

      SFALayout sfa_layout = std::get<7>(problem);
      SFALayout sfb_layout = std::get<8>(problem);
  
      // Divide by 8 since underlying type is INT64
      int sfa_offset = crd2idx(make_coord(make_coord(_0{}, rest_m), make_coord(_0{}, rest_k)), sfa_layout) / 8;
      tma_1d_gmem2smem<CTA_GROUP>(SFA_smem, SFA_tmap, sfa_offset, mbar_addr, cache_A);

      int sfb_offset = crd2idx(make_coord(make_coord(_0{}, rest_n), make_coord(_0{}, rest_k)), sfb_layout) / 8;
      tma_1d_gmem2smem<CTA_GROUP>(SFB_smem, SFB_tmap, sfb_offset, mbar_addr, cache_B);

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

    int mma_phase = 1; // Start with phase 1 as there are no pending MMA operations
    int tma_pipeline_stage = 0;
    for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
      auto [off_m, off_n, dummy, problem_id] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id

      for (int iter_k = 0; iter_k < num_iters; iter_k++) {
        // wait MMA
        // NB when debugging, it will crash w/o this barrier. You cannot keep issuing TMAs w/o the
        // signal the mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase); has completed.
        mbarrier_wait(mma_mbar_addr + tma_pipeline_stage * 8, mma_phase);
        issue_tma(iter_k, tma_pipeline_stage, off_m, off_n, problem_id);

        tma_pipeline_stage = (tma_pipeline_stage + 1) % NUM_STAGES;

        if (tma_pipeline_stage == 0) {
          mma_phase ^= 1;
        }
      }
    }
  }
  else if (warp_id == NUM_WARPS - 1) {
    // allocate tmem
    const int addr = static_cast<int>(__cvta_generic_to_shared(tmem_addr));
    asm volatile("tcgen05.alloc.cta_group::%2.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(addr), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
    if (cta_rank == 0 && elect_sync()) {
      // MMA warp
      // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-instruction-descriptor
      // fp4 MMA doesn't support MMA_M=64. Hence, we will use MMA_M=128 and ignore the rest.
      constexpr int MMA_N = BLOCK_N;
      constexpr int MMA_M = 128 * CTA_GROUP;
      constexpr int16_t cta_mask = (1 << CTA_GROUP) - 1;


      constexpr uint32_t i_desc = (1U << 7U)   // atype=E2M1
                                | (1U << 10U)  // btype=E2M1
                                | ((uint32_t)MMA_N >> 3U << 17U)
                                | ((uint32_t)MMA_M >> 7U << 27U)
                                ;

      int tma_phase = 0;
      int mma_pipeline_stage = 0;
      int epilogue_phase = 1;
      int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result

      for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
        mbarrier_wait(epilogue_mbar_addr + 8 * mainloop_stage, epilogue_phase);

        auto [off_m, off_n, dummy, dummy1] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id

        // Given a k slice of the scale factors which is an atom, the scale factors are stored in groups of four columns in tmem
        // These will be aligned to rows 0, SF_TILE_ROWS, SF_TILE_ROWS * 2, ...
        // And the sf 0..31 will be in index_col=0, 32..63 index_col =1 etc
        // Given an arbitrary row_offset in M or N, in order to find the tmem column offset,
        // Find the residue at the SF tile row granularity, (row_offset % SF_TILE_ROWS) and
        // since each tmem col stores 32 SF, move over by the desired number of columns
        // sf_tmem_offset = (row_offset % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;

        const int M_sf_tmem_offset = (off_m % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
        const int N_sf_tmem_offset = (off_n % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
          // wait TMA
          mbarrier_wait(tma_mbar_addr + mma_pipeline_stage * 8, tma_phase);

          const int A_smem = smem + mma_pipeline_stage * 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;

          // set up shared memory descriptors for A and B
          // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-shared-memory-descriptor
          // 128-byte swizzling. LBO is implied to be 1.
          auto make_desc_AB = [](int addr) -> uint64_t {
            // GMMA::Layout_K_SW128_Atom<cutlass::float_e2m1_t>{} is (_8,_256):(_256,_1) so to get to the start of the next row is 256 / 2 = 128 bytes
            // so to skip 8 rows, we need 8 * 128 bytes.
            const int SBO = 8 * 128; 
            return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
          };
          // no swizzling
          auto make_desc_SF = [](int addr) -> uint64_t {
            // Since we are using no swizzling, the atom has size (8, 16 = 128 / 8). These atoms
            // are tiled (SF_ATOMS_PER_BLOCK_K = 4, 1). So the SBO is 8 * 16
            const int SBO = 8 * 16;
            return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
          };

          // tcgen05.cp -> tcgen05.mma should be pipelined correctly per PTX doc
          // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-memory-consistency-model-pipelined-instructions
          // cutlass issues all of smem->tmem BEFORE mma
          // https://github.com/NVIDIA/cutlass/blob/v4.3.2/include/cutlass/gemm/collective/sm100_blockscaled_mma_warpspecialized.hpp#L1013-L1016
          constexpr uint64_t SF_desc = make_desc_SF(0);
          // The below should really be SF_desc + matrix-descriptor-encode(x) , where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
          // but (x & 0x3FFFF) >> 4 simplifies to x >> 4.
          const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);
          const uint64_t SFB_desc = SF_desc + ((uint64_t)SFB_smem >> 4ULL);

          for (int k = 0; k < BLOCK_K / MMA_K; k++) {
            // The below should really be SFA_desc + (uint64_t)k * matrix-descriptor-encode(512ULL); where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
            // but (x & 0x3FFFF) >> 4 simplifies to x >> 4. So this is SFA_desc + (uint64_t)k * (512ULL >> 4ULL)
            // So we are copying 512 bytes chunks , so adjust the shmem desc accordingly.
            uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL);  // 4 columns, 512 bytes of 128x4 / 32x4x4
            uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);
            // tmem addresses
            // Tensor Memory addresses are 32-bit wide and specify two components.

            // Lane index
            // Column index
            // The layout is as follows:

            // 31 16
            // 15 0
            // Lane index
            // Column index

            // We are using tcgen05.cp.cta_group::1.32x128b.warpx4
            // See https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
            // so the slices ((32, 4), (16, 4)) slices of the scale factors are in the 32 rows (lanes) of tmem

            tcgen05_cp_nvfp4<CTA_GROUP>(SFA_tmem + k * 4, sfa_desc);
            tcgen05_cp_nvfp4<CTA_GROUP>(SFB_tmem + k * 4, sfb_desc);
          }

          // Compare to tiled_mma
          // k1 selects the (BLOCK_M, 256) tile.
          // k2 selects the (BLOCK_M, 64) tile, whose rows are swizzled.
          // NOTE: this doesn't work with BLOCK_N=32, since apparently tcgen05.mma requires SFB_tmem
          // to have 2-column (8-byte) alignment (looks like not documented).
          // HOWEVER: works with a cluster with BLOCK_N = 64 - maybe only the issuing thread has the alignment
          // requirement.
          for (int k1 = 0; k1 < BLOCK_K / 256; k1++) { // BLOCK_K == 256 so this loop only executes once
            for (int k2 = 0; k2 < 256 / MMA_K; k2++) {  // this is the inner loop needed since MMA_K is in general < 256 = BLOCK_K
              // The layout in shmem is the result of the TMA: {256, BLOCK_M, BLOCK_K / 256};
              // Notice that in TMA, it’s implied that the shared memory destination has the natural contiguous layout of the given shape
              // (we don’t specify shared memory stride anywhere).
              // since cuTensorMapEncodeTiled assumes the tensors are col major cute::LayoutLeft
              // The layout is then (256, BLOCK_M, BLOCK_K / 256) : (1, 256, 256 * BLOCK_M) elements 
              // https://gau-nernst.github.io/tcgen05/#decipher-tcgen05
              //
              // Thus given k1 and k2, the coords are (k2 * MMA_K, 0, k1)
              // so the offset in elements is k1 * 256 * BLOCK_M + k2 * MMA_K = k1 * 256 * BLOCK_M + k2 * 64  (MMA_K = 64)
              // so the offset in bytes = k1 * 128 * BLOCK_M + k2 * 32
              // which is the expression below.
      
              // uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
              // crd2idx(make_coord(k2 * MMA_K, 0, k1), asmem_layout) is in elements, divide by 2 to get bytes
              uint64_t a_desc = make_desc_AB(A_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), asmem_layout) / 2);
              uint64_t b_desc = make_desc_AB(B_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), bsmem_layout) / 2);

              // tmem addresses
              // Tensor Memory addresses are 32-bit wide and specify two components.

              // Lane index
              // Column index
              // The layout is as follows:

              // 31 16
              // 15 0
              // Lane index
              // Column index

              int k_sf = k1 * 4 + k2; // the scale factors are in blocks of 4. This is the block idx
              // k_sf is mulitplied by 4 to get us to the actual column.

              const int scale_A_tmem = SFA_tmem + k_sf * 4 + M_sf_tmem_offset;
              const int scale_B_tmem = SFB_tmem + k_sf * 4 + N_sf_tmem_offset;

              const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
              tcgen05_mma_nvfp4<CTA_GROUP>(/*dst tmem*/ BLOCK_N * mainloop_stage, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
            } // for k2
          } // for k1

          // signal MMA done
          // @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
          // tcgen05_commit_mcast does not work if CTA_GROUP = 1
          // tcgen05_commit_mcast<CTA_GROUP>(mma_mbar_addr + mma_pipeline_stage * 8, cta_mask);  // signal MMA done
          // @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
          asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
                      :: "r"(mma_mbar_addr + mma_pipeline_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");

          mma_pipeline_stage = (mma_pipeline_stage + 1) % NUM_STAGES;

          if (mma_pipeline_stage == 0) {
            tma_phase ^= 1;
          }
        }

        // signal mainloop done
        // tcgen05_commit_mcast<CTA_GROUP>(mainloop_mbar_addr + mainloop_stage * 8, cta_mask);  // signal mainloop done
        asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
                    :: "r"(mainloop_mbar_addr + mainloop_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
        mainloop_stage = (mainloop_stage + 1) % 2; 

        if (mainloop_stage == 0) {
          epilogue_phase ^= 1;
        }
      }  // for bid
    } // cta_rank == 0 && elect_sync()
  } // warp_id == NUM_WARPS - 1
  else if (tid < BLOCK_M) {
    // epilogue warps
    int mainloop_phase = 0;
    int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result
    
    for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
      auto result = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
      int off_m = std::get<0>(result);
      int cluster_off_n = std::get<2>(result);
      int problem_id = std::get<3>(result);

      auto problem = get_problem(problem_id);

      const int M = std::get<5>(problem);
      const int N = std::get<6>(problem);

      half* C_ptr = std::get<4>(problem);

      // wait mainloop
      mbarrier_wait(mainloop_mbar_addr + mainloop_stage * 8, mainloop_phase);
      asm volatile("tcgen05.fence::after_thread_sync;");

      int tmem_buffer_offset = BLOCK_N * mainloop_stage;

      auto epilogue_M_major = [&]() {
        // C is M-major
        constexpr int WIDTH = std::min(BLOCK_N, 64);  // using 128 might be slower

        for (int n = 0; n < (BLOCK_N + WIDTH - 1)/ WIDTH; n++) {
          float tmp[WIDTH];  // if WIDTH=128, we are using 128 registers here
          if constexpr (WIDTH == 128) tcgen05_ld_32x32bx128(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
          if constexpr (WIDTH == 64) tcgen05_ld_32x32bx64(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
          if constexpr (WIDTH == 32) tcgen05_ld_32x32bx32(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
          asm volatile("tcgen05.wait::ld.sync.aligned;");

          for (int i = 0; i < WIDTH; i++) {
            // NB row and col are transposed.
            const int row = cluster_off_n + n * WIDTH + i;
            const int col = off_m + tid;
            if (row < N && col < M) {
              C_ptr[row * M + col] = __float2half(tmp[i]);
            }
          }
        }
      };
      auto epilogue_N_major = [&]() {
        // C is N-major
        for (int m = 0; m < 32 / 16; m++) {
          float tmp[BLOCK_N / 2];
          if constexpr (BLOCK_N == 128) tcgen05_ld_16x256bx16(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
          if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
          if constexpr (BLOCK_N == 32) tcgen05_ld_16x256bx4(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
          asm volatile("tcgen05.wait::ld.sync.aligned;");

          for (int i = 0; i < (BLOCK_N + 7)/ 8; i++) {
            // TODO replace this with cute tensors.
            const int row = off_m + warp_id * 32 + m * 16 + lane_id / 4;
            const int col = cluster_off_n + i * 8 + (lane_id % 4) * 2;
            if (row < M) {
              reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 0], tmp[i * 4 + 1]});
            }
            if (row + 8 < M) {
              reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 2], tmp[i * 4 + 3]});
            }
          }
        }
      };

      if constexpr (SWAP_AB) {
        epilogue_M_major();
      } else {
        epilogue_N_major();
      }
      if (elect_sync()) {
        // // signal Epilogue done
        // constexpr int16_t cta_mask = 3;
        // tcgen05_commit_mcast<CTA_GROUP>(epilogue_mbar_addr + 8 * mainloop_stage, cta_mask);  // signal Epilogue done
        const int mbar_addr = (epilogue_mbar_addr + mainloop_stage * 8) & 0xFEFFFFFF;
        asm volatile("mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];" :: "r"(mbar_addr) : "memory");
      }
      
      mainloop_stage = (mainloop_stage + 1) % 2; 
      if (mainloop_stage == 0) {
        mainloop_phase ^= 1;
      }
    } // for bid
  } //epilogue warps

  // All warps.
  if constexpr (CTA_GROUP > 1) {
    // Dont exit w/o the peer CTA.
    asm volatile("barrier.cluster.arrive.release.aligned;");
    asm volatile("barrier.cluster.wait.acquire.aligned;");
  } else {
    __syncthreads();  // all threads finish reading data from tmem
  }
  // deallocate tmem. tmem address should be 0.
  if (warp_id == 0) {
    asm volatile("tcgen05.dealloc.cta_group::%2.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
  }
}

template <
  int K,
  int BLOCK_M,
  int BLOCK_N,
  int BLOCK_K,
  int NUM_STAGES,
  bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p0(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs) {

  constexpr int CTA_GROUP = 2;
  static_assert(BLOCK_K % 256 == 0);

  constexpr int rest_K = K / 16 / 4;
  constexpr int grid = NUM_SMS;
  constexpr int tb_size = BLOCK_M + 2 * WARP_SIZE;
  constexpr int AB_size = (BLOCK_M + BLOCK_N / CTA_GROUP) * (BLOCK_K / 2);
  constexpr int SFAB_size = 128 * (BLOCK_K / 16) * 2;
  constexpr int smem_size = (AB_size + SFAB_size) * NUM_STAGES;

  // This layout is used as TMA assumes col major, this is a colasced canonical layout.
  auto layout_smem_A = make_layout(make_shape(256, BLOCK_M, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
  // Each member of the cluster will get a cta_rank specific offset
  auto layout_smem_B = make_layout(make_shape(256, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_N / CTA_GROUP));

  constexpr int SF_TILE_ROWS = 128;
  constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
  constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns

  auto make_sf_layout = [&](int rest) {
    return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)), 
                              make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
  };
  
  // This is for the type only
  auto layout_sfa = make_sf_layout(16); 
  auto layout_sfb = make_sf_layout(16); 

  auto populate_kernel_config = [&ABC, &SFAB, &outputs](CUtensorMap* A_tmaps, CUtensorMap* B_tmaps, CUtensorMap* SFA_tmaps, CUtensorMap* SFB_tmaps,
    half** C_ptrs, int* Ms, int* Ns, auto problem_size_const) {
    constexpr int PROBLEM_SIZE = decltype(problem_size_const)::value;

    for (int workItem = 0; workItem < PROBLEM_SIZE; ++workItem) {
      std::tuple<at::Tensor, at::Tensor, at::Tensor> abc_tuple = ABC[workItem]; 
      auto A = std::get<0>(abc_tuple);
      auto B = std::get<1>(abc_tuple);

      std::tuple<at::Tensor, at::Tensor> sfab_tuple = SFAB[workItem]; 
      auto SFA = std::get<0>(sfab_tuple);
      auto SFB = std::get<1>(sfab_tuple);

      at::Tensor C = outputs[workItem];

      const int M = A.size(0);
      const int N = B.size(0);

      auto A_ptr   = reinterpret_cast<const char *>(A.data_ptr());
      auto B_ptr  = reinterpret_cast<const char *>(B.data_ptr());
      auto SFA_ptr = reinterpret_cast<const char *>(SFA.data_ptr());
      auto SFB_ptr = reinterpret_cast<const char *>(SFB.data_ptr());
      C_ptrs[workItem]   = reinterpret_cast<half *>(C.data_ptr());

      int new_M = M;
      int new_N = N;

      if constexpr (SWAP_AB) {
        std::swap(A_ptr, B_ptr);
        std::swap(SFA_ptr, SFB_ptr);
        std::swap(new_M, new_N);
      }

      Ms[workItem] = new_M;
      Ns[workItem] = new_N;

      init_AB_tmap(&A_tmaps[workItem], A_ptr, new_M, K, BLOCK_M, BLOCK_K);
      init_AB_tmap(&B_tmaps[workItem], B_ptr, new_N, K, BLOCK_N / CTA_GROUP, BLOCK_K);

      // CUtensorMap SFA_tmap0, SFB_tmap0;
      // SF Atom is ((32, 4), (16, 4))
      const int rest_M = DIVUP(new_M, SF_TILE_ROWS);
      const int rest_N = DIVUP(new_N, SF_TILE_ROWS);

      auto new_M_padded = SF_TILE_ROWS * rest_M; // 128 is the number of rows in the atom.
      auto new_N_padded = SF_TILE_ROWS * rest_N;

      init_SF_tmap(&SFA_tmaps[workItem], SFA_ptr, new_M_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
      init_SF_tmap(&SFB_tmaps[workItem], SFB_ptr, new_N_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
    }
  };

  constexpr int PROBLEM_SIZE = 8;
  CUtensorMap host_A_tmaps[PROBLEM_SIZE];
  CUtensorMap host_B_tmaps[PROBLEM_SIZE];
  CUtensorMap host_SFA_tmaps[PROBLEM_SIZE];
  CUtensorMap host_SFB_tmaps[PROBLEM_SIZE];
  half* host_C_ptrs[PROBLEM_SIZE];
  int host_Ms[PROBLEM_SIZE];
  int host_Ns[PROBLEM_SIZE];

  populate_kernel_config(host_A_tmaps, host_B_tmaps, host_SFA_tmaps, host_SFB_tmaps, host_C_ptrs, host_Ms, host_Ns, std::integral_constant<int, PROBLEM_SIZE>{});

  cudaMemcpyToSymbol(A_tmaps<PROBLEM_SIZE>, host_A_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(B_tmaps<PROBLEM_SIZE>, host_B_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(SFA_tmaps<PROBLEM_SIZE>, host_SFA_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(SFB_tmaps<PROBLEM_SIZE>, host_SFB_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(C_ptrs<PROBLEM_SIZE>, host_C_ptrs, PROBLEM_SIZE * sizeof(half*));
  cudaMemcpyToSymbol(Ms<PROBLEM_SIZE>, host_Ms, PROBLEM_SIZE * sizeof(int));
  cudaMemcpyToSymbol(Ns<PROBLEM_SIZE>, host_Ns, PROBLEM_SIZE * sizeof(int));

  auto this_kernel = multi_gemm_kernel_p0<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES,
                            decltype(layout_sfa), decltype(layout_sfb),
                            SWAP_AB, PROBLEM_SIZE, CTA_GROUP>;

  if (smem_size > 48'000)
    cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);

  this_kernel<<<grid, tb_size, smem_size>>>();
  return {outputs[0], outputs[1], outputs[2], outputs[3], outputs[4], outputs[5], outputs[6], outputs[7]};

}


template std::vector<at::Tensor> launch_gemm_cstr_p0<7168, 128, 64, 256, 9, true>(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs);
"""

CPP_SRC_0 = r"""

#include <vector>

#include <torch/script.h>

template <
  int K,
  int BLOCK_M,
  int BLOCK_N,
  int BLOCK_K,
  int NUM_STAGES,
  bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p0(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs);

std::vector<at::Tensor> gemm_cstr_p0(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs) {
    auto C = launch_gemm_cstr_p0<7168, 128, 64, 256, 9, true>(ABC, SFAB, outputs);

  return C;
}

TORCH_LIBRARY(blockscaled_cstr_p0, m) {
  m.def("gemm_cstr_p0((Tensor, Tensor, Tensor)[] ABC, (Tensor, Tensor)[] SFAB, Tensor[] outputs) -> Tensor[]");
  m.impl("gemm_cstr_p0", &gemm_cstr_p0);
}
"""

CUDA_SRC_1  = r"""
using namespace cute;

constexpr int MAX_TMEM_COLS = 512;
constexpr int NUM_SMS = 148;

constexpr int NUM_BLOCKS0 = 56;
constexpr int NUM_BLOCKS1 = 112;
constexpr int NUM_BLOCKS2 = 168;
constexpr int NUM_BLOCKS3 = 112;
constexpr int NUM_BLOCKS4 = 168;
constexpr int NUM_BLOCKS5 = 168;
constexpr int NUM_BLOCKS6 = 224;
constexpr int NUM_BLOCKS7 = 168;

constexpr int GEMM_BLOCK_END0 = NUM_BLOCKS0;
constexpr int GEMM_BLOCK_END1 = GEMM_BLOCK_END0 + NUM_BLOCKS1;
constexpr int GEMM_BLOCK_END2 = GEMM_BLOCK_END1 + NUM_BLOCKS2;
constexpr int GEMM_BLOCK_END3 = GEMM_BLOCK_END2 + NUM_BLOCKS3;
constexpr int GEMM_BLOCK_END4 = GEMM_BLOCK_END3 + NUM_BLOCKS4;
constexpr int GEMM_BLOCK_END5 = GEMM_BLOCK_END4 + NUM_BLOCKS5;
constexpr int GEMM_BLOCK_END6 = GEMM_BLOCK_END5 + NUM_BLOCKS6;
constexpr int GEMM_BLOCK_END7 = GEMM_BLOCK_END6 + NUM_BLOCKS7;

constexpr int NUM_BLOCKS = GEMM_BLOCK_END7;

template<int PROBLEM_SIZE>
__constant__ CUtensorMap A_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ CUtensorMap B_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFA_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFB_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ half* C_ptrs[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ int Ms[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ int Ns[PROBLEM_SIZE];

template <
  int K,
  int BLOCK_M,
  int BLOCK_N,
  int BLOCK_K,
  int NUM_STAGES,
  class SFALayout,
  class SFBLayout,
  bool SWAP_AB,
  int PROBLEM_SIZE,
  int CTA_GROUP
>
__global__
__cluster_dims__(CTA_GROUP, 1, 1)
__launch_bounds__(BLOCK_M + 2 * WARP_SIZE)  // __launch_bounds__(MAX_THREADS_PER_BLOCK, MIN_BLOCKS_PER_MP)
void multi_gemm_kernel_p1() {
  // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-matrix-shape
  static_assert(BLOCK_M == 128); // There doesnt seem to be any code for tiling block_M

  // This layout is used as TMA assumes col major, this is a colasced canonical layout.
  auto asmem_layout = make_layout(make_shape(_256{}, BLOCK_M, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
  // Each member of the cluster will get a cta_rank specific offset
  auto bsmem_layout = make_layout(make_shape(_256{}, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_N / CTA_GROUP));

  // an SF atom is ((32, 4), (16, 4)), the shmem layout is ((128 * SF_ATOMS_PER_BLOCK_K, 4)) for TMA purposes
  // Since we are using CU_TENSOR_MAP_SWIZZLE_NONE, the TMA atom has shape  (8, 128bits) or (8, 16) bytes, so SBO below is 8 * 16
  constexpr int SF_TILE_ROWS = 128;
  constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
  constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns
  constexpr int SF_ROWS_PER_TMEM_COL = 32; // each tmem col has 32 rows - since we are using tcgen05.cp.cta_group::1.32x128b.warpx4

  using Problem = std::tuple<const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, half*, int, int, SFALayout, SFBLayout>;

  auto get_problem = [&](int problem_id) -> Problem {
  
    auto make_sf_layout = [](int rest) {
      constexpr int rest_K = K / 16 / 4;
      return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)), 
                                make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
    };

    const CUtensorMap* A_tmap = &A_tmaps<PROBLEM_SIZE>[problem_id];
    const CUtensorMap* B_tmap = &B_tmaps<PROBLEM_SIZE>[problem_id];
    const CUtensorMap* SFA_tmap = &SFA_tmaps<PROBLEM_SIZE>[problem_id];
    const CUtensorMap* SFB_tmap = &SFB_tmaps<PROBLEM_SIZE>[problem_id];

    half* C = C_ptrs<PROBLEM_SIZE>[problem_id];

    int M = Ms<PROBLEM_SIZE>[problem_id];
    int N = Ns<PROBLEM_SIZE>[problem_id];

    const int rest_M = DIVUP(M, SF_TILE_ROWS);
    const int rest_N = DIVUP(N, SF_TILE_ROWS);
    SFALayout sfa_layout = make_sf_layout(rest_M);
    SFBLayout sfb_layout = make_sf_layout(rest_N);

    return {A_tmap, B_tmap, SFA_tmap, SFB_tmap, C, M, N, sfa_layout, sfb_layout};
  };

  // CTA rank in a cluster
  int cta_rank;
  asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));

  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;

  // set up smem
  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; // BLOCK_M * BLOCK_K elements, each element is 4 bits so divide by 2 to get bytes.
  constexpr int B_size = BLOCK_N * BLOCK_K / 2 / CTA_GROUP; // each CTA only loads half of B
  constexpr int SFA_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K;  // always copy one atom - this is ATOM_SIZE_IN_BYTES * ATOMS_PER_BLOCK_K
  constexpr int SFB_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K;
  constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;

  // set up mbarriers and tmem
  // we have NUM_STAGES mbars for TMA
  //         NUM_STAGES mbars for MMA
  //                  1 mbar  for mainloop
  #pragma nv_diag_suppress static_var_with_dynamic_init
  __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
  __shared__ int tmem_addr[1];  // tmem address is 32-bit

  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;
  const int epilogue_mbar_addr = mainloop_mbar_addr + 2 * 8;

  // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
  // each MMA consumes:
  // - (128, 64) of A -> (128, 4) of SFA -> reshaped as (32, 4', 4) -> 4 tmem columns
  constexpr int SFA_tmem = BLOCK_N * 2; // Double buffer the mma output.
  constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K); // we need BLOCK_K / MMA_K many scale factors for the "tiledMMA"

  if (warp_id == 0 && elect_sync()) {
    for (int prefetch_idx = 0; prefetch_idx < 8; ++prefetch_idx) {
      asm volatile("prefetch.tensormap [%0];" :: "l"(&A_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
      asm volatile("prefetch.tensormap [%0];" :: "l"(&B_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
      asm volatile("prefetch.tensormap [%0];" :: "l"(&SFA_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
      asm volatile("prefetch.tensormap [%0];" :: "l"(&SFB_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
    }
  }  else if (warp_id == 1 && elect_sync()) {
    for (int i = 0; i < NUM_STAGES; i++) {
      mbarrier_init(tma_mbar_addr + i * 8, CTA_GROUP); // one thread in each cluster CTA reports.
      mbarrier_init(mma_mbar_addr + i * 8, 1);
    }
    mbarrier_init(mainloop_mbar_addr, 1); // barrier for first tmem mma buffer
    mbarrier_init(mainloop_mbar_addr + 8, 1); // barrier for second tmem mma buffer

    mbarrier_init(epilogue_mbar_addr, 4 * CTA_GROUP);
    mbarrier_init(epilogue_mbar_addr + 8, 4 * CTA_GROUP);

    asm volatile("fence.mbarrier_init.release.cluster;");  // visible to async proxy
  }
  
  if constexpr (CTA_GROUP > 1) {
    // visible to all threads in a cluster
    asm volatile("barrier.cluster.arrive.relaxed.aligned;");
    asm volatile("barrier.cluster.wait.acquire.aligned;");
  }
  else {
    // visible to all threads in a threadblock
    __syncthreads();
  }

  constexpr int num_iters = K / BLOCK_K;
  const int bid = blockIdx.x;

  auto scheduler = [&](int global_cluster_id) -> std::tuple<int, int, int, int> {
    int problem_id = 0;
    int local_cluster_id = global_cluster_id;

    if (global_cluster_id < GEMM_BLOCK_END0 / CTA_GROUP) {
      problem_id = 0;
      local_cluster_id = global_cluster_id;
    } else if (global_cluster_id < GEMM_BLOCK_END1 / CTA_GROUP) {
      problem_id = 1;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END0 / CTA_GROUP;
    } else if (global_cluster_id < GEMM_BLOCK_END2 / CTA_GROUP) {
      problem_id = 2;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END1 / CTA_GROUP;
    } else if (global_cluster_id < GEMM_BLOCK_END3 / CTA_GROUP) {
      problem_id = 3;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END2 / CTA_GROUP;
    } else if (global_cluster_id < GEMM_BLOCK_END4 / CTA_GROUP) {
      problem_id = 4;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END3 / CTA_GROUP;
    } else if (global_cluster_id < GEMM_BLOCK_END5 / CTA_GROUP) {
      problem_id = 5;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END4 / CTA_GROUP;
    } else if (global_cluster_id < GEMM_BLOCK_END6 / CTA_GROUP) {
      problem_id = 6;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END5 / CTA_GROUP;
    } else if (global_cluster_id < GEMM_BLOCK_END7 / CTA_GROUP) {
      problem_id = 7;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END6 / CTA_GROUP;
    }

    int N = Ns<PROBLEM_SIZE>[problem_id];

    const int grid_n = DIVUP(N, BLOCK_N) ;
    // The cluster Ids are arranged in a (grid_m / 2 , grid_n) matrix in a row major layout. 
    const int bid_m = local_cluster_id / grid_n * CTA_GROUP + cta_rank;
    const int bid_n = local_cluster_id % grid_n;

    const int off_m = bid_m * BLOCK_M;
    const int cluster_off_n = bid_n * BLOCK_N; // N offset for this cluster
    const int off_n = cluster_off_n + cta_rank * (BLOCK_N / CTA_GROUP); // N offset for this CTA
    return {off_m, off_n, cluster_off_n, problem_id};
  };

  // warp-specialization
  if (warp_id == NUM_WARPS - 2 && elect_sync()) {
    // TMA warp

    auto issue_tma = [&](int iter_k, int stage_id, int off_m, int off_n, int problem_id) {
      const int mbar_addr = (tma_mbar_addr + stage_id * 8) & 0xFEFFFFFF;  // CTA0's barrier
      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 problem = get_problem(problem_id);

      const CUtensorMap* A_tmap = std::get<0>(problem);
      const CUtensorMap* B_tmap = std::get<1>(problem);    

      const int M = std::get<5>(problem);
      const int N = std::get<6>(problem);

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

      // issue TMA
      const int off_k = iter_k * BLOCK_K;
      tma_3d_gmem2smem<CTA_GROUP>(A_smem, A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
      tma_3d_gmem2smem<CTA_GROUP>(B_smem, B_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);

      const int rest_m = off_m / 32 / 4; // M atom is (32, 4)
      const int rest_n = off_n / 32 / 4; // N atom is (32, 4)

      const int rest_k = off_k / 16 / 4; // The SF atom is ((32, 4), (16, 4))

      const CUtensorMap* SFA_tmap = std::get<2>(problem);
      const CUtensorMap* SFB_tmap = std::get<3>(problem);

      SFALayout sfa_layout = std::get<7>(problem);
      SFALayout sfb_layout = std::get<8>(problem);
  
      // Divide by 8 since underlying type is INT64
      int sfa_offset = crd2idx(make_coord(make_coord(_0{}, rest_m), make_coord(_0{}, rest_k)), sfa_layout) / 8;
      tma_1d_gmem2smem<CTA_GROUP>(SFA_smem, SFA_tmap, sfa_offset, mbar_addr, cache_A);

      int sfb_offset = crd2idx(make_coord(make_coord(_0{}, rest_n), make_coord(_0{}, rest_k)), sfb_layout) / 8;
      tma_1d_gmem2smem<CTA_GROUP>(SFB_smem, SFB_tmap, sfb_offset, mbar_addr, cache_B);

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

    int mma_phase = 1; // Start with phase 1 as there are no pending MMA operations
    int tma_pipeline_stage = 0;
    for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
      auto [off_m, off_n, dummy, problem_id] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id

      for (int iter_k = 0; iter_k < num_iters; iter_k++) {
        // wait MMA
        // NB when debugging, it will crash w/o this barrier. You cannot keep issuing TMAs w/o the
        // signal the mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase); has completed.
        mbarrier_wait(mma_mbar_addr + tma_pipeline_stage * 8, mma_phase);
        issue_tma(iter_k, tma_pipeline_stage, off_m, off_n, problem_id);

        tma_pipeline_stage = (tma_pipeline_stage + 1) % NUM_STAGES;

        if (tma_pipeline_stage == 0) {
          mma_phase ^= 1;
        }
      }
    }
  }
  else if (warp_id == NUM_WARPS - 1) {
    // allocate tmem
    const int addr = static_cast<int>(__cvta_generic_to_shared(tmem_addr));
    asm volatile("tcgen05.alloc.cta_group::%2.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(addr), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
    if (cta_rank == 0 && elect_sync()) {
      // MMA warp
      // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-instruction-descriptor
      // fp4 MMA doesn't support MMA_M=64. Hence, we will use MMA_M=128 and ignore the rest.
      constexpr int MMA_N = BLOCK_N;
      constexpr int MMA_M = 128 * CTA_GROUP;
      constexpr int16_t cta_mask = (1 << CTA_GROUP) - 1;


      constexpr uint32_t i_desc = (1U << 7U)   // atype=E2M1
                                | (1U << 10U)  // btype=E2M1
                                | ((uint32_t)MMA_N >> 3U << 17U)
                                | ((uint32_t)MMA_M >> 7U << 27U)
                                ;

      int tma_phase = 0;
      int mma_pipeline_stage = 0;
      int epilogue_phase = 1;
      int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result

      for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
        mbarrier_wait(epilogue_mbar_addr + 8 * mainloop_stage, epilogue_phase);

        auto [off_m, off_n, dummy, dummy1] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id

        // Given a k slice of the scale factors which is an atom, the scale factors are stored in groups of four columns in tmem
        // These will be aligned to rows 0, SF_TILE_ROWS, SF_TILE_ROWS * 2, ...
        // And the sf 0..31 will be in index_col=0, 32..63 index_col =1 etc
        // Given an arbitrary row_offset in M or N, in order to find the tmem column offset,
        // Find the residue at the SF tile row granularity, (row_offset % SF_TILE_ROWS) and
        // since each tmem col stores 32 SF, move over by the desired number of columns
        // sf_tmem_offset = (row_offset % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;

        const int M_sf_tmem_offset = (off_m % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
        const int N_sf_tmem_offset = (off_n % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
          // wait TMA
          mbarrier_wait(tma_mbar_addr + mma_pipeline_stage * 8, tma_phase);

          const int A_smem = smem + mma_pipeline_stage * 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;

          // set up shared memory descriptors for A and B
          // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-shared-memory-descriptor
          // 128-byte swizzling. LBO is implied to be 1.
          auto make_desc_AB = [](int addr) -> uint64_t {
            // GMMA::Layout_K_SW128_Atom<cutlass::float_e2m1_t>{} is (_8,_256):(_256,_1) so to get to the start of the next row is 256 / 2 = 128 bytes
            // so to skip 8 rows, we need 8 * 128 bytes.
            const int SBO = 8 * 128; 
            return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
          };
          // no swizzling
          auto make_desc_SF = [](int addr) -> uint64_t {
            // Since we are using no swizzling, the atom has size (8, 16 = 128 / 8). These atoms
            // are tiled (SF_ATOMS_PER_BLOCK_K = 4, 1). So the SBO is 8 * 16
            const int SBO = 8 * 16;
            return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
          };

          // tcgen05.cp -> tcgen05.mma should be pipelined correctly per PTX doc
          // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-memory-consistency-model-pipelined-instructions
          // cutlass issues all of smem->tmem BEFORE mma
          // https://github.com/NVIDIA/cutlass/blob/v4.3.2/include/cutlass/gemm/collective/sm100_blockscaled_mma_warpspecialized.hpp#L1013-L1016
          constexpr uint64_t SF_desc = make_desc_SF(0);
          // The below should really be SF_desc + matrix-descriptor-encode(x) , where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
          // but (x & 0x3FFFF) >> 4 simplifies to x >> 4.
          const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);
          const uint64_t SFB_desc = SF_desc + ((uint64_t)SFB_smem >> 4ULL);

          for (int k = 0; k < BLOCK_K / MMA_K; k++) {
            // The below should really be SFA_desc + (uint64_t)k * matrix-descriptor-encode(512ULL); where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
            // but (x & 0x3FFFF) >> 4 simplifies to x >> 4. So this is SFA_desc + (uint64_t)k * (512ULL >> 4ULL)
            // So we are copying 512 bytes chunks , so adjust the shmem desc accordingly.
            uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL);  // 4 columns, 512 bytes of 128x4 / 32x4x4
            uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);
            // tmem addresses
            // Tensor Memory addresses are 32-bit wide and specify two components.

            // Lane index
            // Column index
            // The layout is as follows:

            // 31 16
            // 15 0
            // Lane index
            // Column index

            // We are using tcgen05.cp.cta_group::1.32x128b.warpx4
            // See https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
            // so the slices ((32, 4), (16, 4)) slices of the scale factors are in the 32 rows (lanes) of tmem

            tcgen05_cp_nvfp4<CTA_GROUP>(SFA_tmem + k * 4, sfa_desc);
            tcgen05_cp_nvfp4<CTA_GROUP>(SFB_tmem + k * 4, sfb_desc);
          }

          // Compare to tiled_mma
          // k1 selects the (BLOCK_M, 256) tile.
          // k2 selects the (BLOCK_M, 64) tile, whose rows are swizzled.
          // NOTE: this doesn't work with BLOCK_N=32, since apparently tcgen05.mma requires SFB_tmem
          // to have 2-column (8-byte) alignment (looks like not documented).
          // HOWEVER: works with a cluster with BLOCK_N = 64 - maybe only the issuing thread has the alignment
          // requirement.
          for (int k1 = 0; k1 < BLOCK_K / 256; k1++) { // BLOCK_K == 256 so this loop only executes once
            for (int k2 = 0; k2 < 256 / MMA_K; k2++) {  // this is the inner loop needed since MMA_K is in general < 256 = BLOCK_K
              // The layout in shmem is the result of the TMA: {256, BLOCK_M, BLOCK_K / 256};
              // Notice that in TMA, it’s implied that the shared memory destination has the natural contiguous layout of the given shape
              // (we don’t specify shared memory stride anywhere).
              // since cuTensorMapEncodeTiled assumes the tensors are col major cute::LayoutLeft
              // The layout is then (256, BLOCK_M, BLOCK_K / 256) : (1, 256, 256 * BLOCK_M) elements 
              // https://gau-nernst.github.io/tcgen05/#decipher-tcgen05
              //
              // Thus given k1 and k2, the coords are (k2 * MMA_K, 0, k1)
              // so the offset in elements is k1 * 256 * BLOCK_M + k2 * MMA_K = k1 * 256 * BLOCK_M + k2 * 64  (MMA_K = 64)
              // so the offset in bytes = k1 * 128 * BLOCK_M + k2 * 32
              // which is the expression below.
      
              // uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
              // crd2idx(make_coord(k2 * MMA_K, 0, k1), asmem_layout) is in elements, divide by 2 to get bytes
              uint64_t a_desc = make_desc_AB(A_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), asmem_layout) / 2);
              uint64_t b_desc = make_desc_AB(B_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), bsmem_layout) / 2);

              // tmem addresses
              // Tensor Memory addresses are 32-bit wide and specify two components.

              // Lane index
              // Column index
              // The layout is as follows:

              // 31 16
              // 15 0
              // Lane index
              // Column index

              int k_sf = k1 * 4 + k2; // the scale factors are in blocks of 4. This is the block idx
              // k_sf is mulitplied by 4 to get us to the actual column.

              const int scale_A_tmem = SFA_tmem + k_sf * 4 + M_sf_tmem_offset;
              const int scale_B_tmem = SFB_tmem + k_sf * 4 + N_sf_tmem_offset;

              const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
              tcgen05_mma_nvfp4<CTA_GROUP>(/*dst tmem*/ BLOCK_N * mainloop_stage, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
            } // for k2
          } // for k1

          // signal MMA done
          // @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
          // tcgen05_commit_mcast does not work if CTA_GROUP = 1
          // tcgen05_commit_mcast<CTA_GROUP>(mma_mbar_addr + mma_pipeline_stage * 8, cta_mask);  // signal MMA done
          // @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
          asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
                      :: "r"(mma_mbar_addr + mma_pipeline_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");

          mma_pipeline_stage = (mma_pipeline_stage + 1) % NUM_STAGES;

          if (mma_pipeline_stage == 0) {
            tma_phase ^= 1;
          }
        }

        // signal mainloop done
        // tcgen05_commit_mcast<CTA_GROUP>(mainloop_mbar_addr + mainloop_stage * 8, cta_mask);  // signal mainloop done
        asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
                    :: "r"(mainloop_mbar_addr + mainloop_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
        mainloop_stage = (mainloop_stage + 1) % 2; 

        if (mainloop_stage == 0) {
          epilogue_phase ^= 1;
        }
      }  // for bid
    } // cta_rank == 0 && elect_sync()
  } // warp_id == NUM_WARPS - 1
  else if (tid < BLOCK_M) {
    // epilogue warps
    int mainloop_phase = 0;
    int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result
    
    for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
      auto result = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
      int off_m = std::get<0>(result);
      int cluster_off_n = std::get<2>(result);
      int problem_id = std::get<3>(result);

      auto problem = get_problem(problem_id);

      const int M = std::get<5>(problem);
      const int N = std::get<6>(problem);

      half* C_ptr = std::get<4>(problem);

      // wait mainloop
      mbarrier_wait(mainloop_mbar_addr + mainloop_stage * 8, mainloop_phase);
      asm volatile("tcgen05.fence::after_thread_sync;");

      int tmem_buffer_offset = BLOCK_N * mainloop_stage;

      auto epilogue_M_major = [&]() {
        // C is M-major
        constexpr int WIDTH = std::min(BLOCK_N, 64);  // using 128 might be slower

        for (int n = 0; n < (BLOCK_N + WIDTH - 1)/ WIDTH; n++) {
          float tmp[WIDTH];  // if WIDTH=128, we are using 128 registers here
          if constexpr (WIDTH == 128) tcgen05_ld_32x32bx128(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
          if constexpr (WIDTH == 64) tcgen05_ld_32x32bx64(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
          if constexpr (WIDTH == 32) tcgen05_ld_32x32bx32(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
          asm volatile("tcgen05.wait::ld.sync.aligned;");

          for (int i = 0; i < WIDTH; i++) {
            // NB row and col are transposed.
            const int row = cluster_off_n + n * WIDTH + i;
            const int col = off_m + tid;
            if (row < N && col < M) {
              C_ptr[row * M + col] = __float2half(tmp[i]);
            }
          }
        }
      };
      auto epilogue_N_major = [&]() {
        // C is N-major
        for (int m = 0; m < 32 / 16; m++) {
          float tmp[BLOCK_N / 2];
          if constexpr (BLOCK_N == 128) tcgen05_ld_16x256bx16(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
          if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
          if constexpr (BLOCK_N == 32) tcgen05_ld_16x256bx4(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
          asm volatile("tcgen05.wait::ld.sync.aligned;");

          for (int i = 0; i < (BLOCK_N + 7)/ 8; i++) {
            // TODO replace this with cute tensors.
            const int row = off_m + warp_id * 32 + m * 16 + lane_id / 4;
            const int col = cluster_off_n + i * 8 + (lane_id % 4) * 2;
            if (row < M) {
              reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 0], tmp[i * 4 + 1]});
            }
            if (row + 8 < M) {
              reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 2], tmp[i * 4 + 3]});
            }
          }
        }
      };

      if constexpr (SWAP_AB) {
        epilogue_M_major();
      } else {
        epilogue_N_major();
      }
      if (elect_sync()) {
        // // signal Epilogue done
        // constexpr int16_t cta_mask = 3;
        // tcgen05_commit_mcast<CTA_GROUP>(epilogue_mbar_addr + 8 * mainloop_stage, cta_mask);  // signal Epilogue done
        const int mbar_addr = (epilogue_mbar_addr + mainloop_stage * 8) & 0xFEFFFFFF;
        asm volatile("mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];" :: "r"(mbar_addr) : "memory");
      }
      
      mainloop_stage = (mainloop_stage + 1) % 2; 
      if (mainloop_stage == 0) {
        mainloop_phase ^= 1;
      }
    } // for bid
  } //epilogue warps

  // All warps.
  if constexpr (CTA_GROUP > 1) {
    // Dont exit w/o the peer CTA.
    asm volatile("barrier.cluster.arrive.release.aligned;");
    asm volatile("barrier.cluster.wait.acquire.aligned;");
  } else {
    __syncthreads();  // all threads finish reading data from tmem
  }
  // deallocate tmem. tmem address should be 0.
  if (warp_id == 0) {
    asm volatile("tcgen05.dealloc.cta_group::%2.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
  }
}

template <
  int K,
  int BLOCK_M,
  int BLOCK_N,
  int BLOCK_K,
  int NUM_STAGES,
  bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p1(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs) {

  constexpr int CTA_GROUP = 2;
  static_assert(BLOCK_K % 256 == 0);

  constexpr int rest_K = K / 16 / 4;
  constexpr int grid = NUM_SMS;
  constexpr int tb_size = BLOCK_M + 2 * WARP_SIZE;
  constexpr int AB_size = (BLOCK_M + BLOCK_N / CTA_GROUP) * (BLOCK_K / 2);
  constexpr int SFAB_size = 128 * (BLOCK_K / 16) * 2;
  constexpr int smem_size = (AB_size + SFAB_size) * NUM_STAGES;

  // This layout is used as TMA assumes col major, this is a colasced canonical layout.
  auto layout_smem_A = make_layout(make_shape(256, BLOCK_M, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
  // Each member of the cluster will get a cta_rank specific offset
  auto layout_smem_B = make_layout(make_shape(256, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_N / CTA_GROUP));

  constexpr int SF_TILE_ROWS = 128;
  constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
  constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns

  auto make_sf_layout = [&](int rest) {
    return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)), 
                              make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
  };
  
  // This is for the type only
  auto layout_sfa = make_sf_layout(16); 
  auto layout_sfb = make_sf_layout(16); 

  auto populate_kernel_config = [&ABC, &SFAB, &outputs](CUtensorMap* A_tmaps, CUtensorMap* B_tmaps, CUtensorMap* SFA_tmaps, CUtensorMap* SFB_tmaps,
    half** C_ptrs, int* Ms, int* Ns, auto problem_size_const) {
    constexpr int PROBLEM_SIZE = decltype(problem_size_const)::value;

    for (int workItem = 0; workItem < PROBLEM_SIZE; ++workItem) {
      std::tuple<at::Tensor, at::Tensor, at::Tensor> abc_tuple = ABC[workItem]; 
      auto A = std::get<0>(abc_tuple);
      auto B = std::get<1>(abc_tuple);

      std::tuple<at::Tensor, at::Tensor> sfab_tuple = SFAB[workItem]; 
      auto SFA = std::get<0>(sfab_tuple);
      auto SFB = std::get<1>(sfab_tuple);

      at::Tensor C = outputs[workItem];

      const int M = A.size(0);
      const int N = B.size(0);

      auto A_ptr   = reinterpret_cast<const char *>(A.data_ptr());
      auto B_ptr  = reinterpret_cast<const char *>(B.data_ptr());
      auto SFA_ptr = reinterpret_cast<const char *>(SFA.data_ptr());
      auto SFB_ptr = reinterpret_cast<const char *>(SFB.data_ptr());
      C_ptrs[workItem]   = reinterpret_cast<half *>(C.data_ptr());

      int new_M = M;
      int new_N = N;

      if constexpr (SWAP_AB) {
        std::swap(A_ptr, B_ptr);
        std::swap(SFA_ptr, SFB_ptr);
        std::swap(new_M, new_N);
      }

      Ms[workItem] = new_M;
      Ns[workItem] = new_N;

      init_AB_tmap(&A_tmaps[workItem], A_ptr, new_M, K, BLOCK_M, BLOCK_K);
      init_AB_tmap(&B_tmaps[workItem], B_ptr, new_N, K, BLOCK_N / CTA_GROUP, BLOCK_K);

      // CUtensorMap SFA_tmap0, SFB_tmap0;
      // SF Atom is ((32, 4), (16, 4))
      const int rest_M = DIVUP(new_M, SF_TILE_ROWS);
      const int rest_N = DIVUP(new_N, SF_TILE_ROWS);

      auto new_M_padded = SF_TILE_ROWS * rest_M; // 128 is the number of rows in the atom.
      auto new_N_padded = SF_TILE_ROWS * rest_N;

      init_SF_tmap(&SFA_tmaps[workItem], SFA_ptr, new_M_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
      init_SF_tmap(&SFB_tmaps[workItem], SFB_ptr, new_N_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
    }
  };

  constexpr int PROBLEM_SIZE = 8;
  CUtensorMap host_A_tmaps[PROBLEM_SIZE];
  CUtensorMap host_B_tmaps[PROBLEM_SIZE];
  CUtensorMap host_SFA_tmaps[PROBLEM_SIZE];
  CUtensorMap host_SFB_tmaps[PROBLEM_SIZE];
  half* host_C_ptrs[PROBLEM_SIZE];
  int host_Ms[PROBLEM_SIZE];
  int host_Ns[PROBLEM_SIZE];

  populate_kernel_config(host_A_tmaps, host_B_tmaps, host_SFA_tmaps, host_SFB_tmaps, host_C_ptrs, host_Ms, host_Ns, std::integral_constant<int, PROBLEM_SIZE>{});

  cudaMemcpyToSymbol(A_tmaps<PROBLEM_SIZE>, host_A_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(B_tmaps<PROBLEM_SIZE>, host_B_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(SFA_tmaps<PROBLEM_SIZE>, host_SFA_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(SFB_tmaps<PROBLEM_SIZE>, host_SFB_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(C_ptrs<PROBLEM_SIZE>, host_C_ptrs, PROBLEM_SIZE * sizeof(half*));
  cudaMemcpyToSymbol(Ms<PROBLEM_SIZE>, host_Ms, PROBLEM_SIZE * sizeof(int));
  cudaMemcpyToSymbol(Ns<PROBLEM_SIZE>, host_Ns, PROBLEM_SIZE * sizeof(int));

  auto this_kernel = multi_gemm_kernel_p1<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES,
                            decltype(layout_sfa), decltype(layout_sfb),
                            SWAP_AB, PROBLEM_SIZE, CTA_GROUP>;

  if (smem_size > 48'000)
    cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);

  this_kernel<<<grid, tb_size, smem_size>>>();
  return {outputs[0], outputs[1], outputs[2], outputs[3], outputs[4], outputs[5], outputs[6], outputs[7]};

}


template std::vector<at::Tensor> launch_gemm_cstr_p1<2048, 128, 64, 256, 9, true>(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs);
"""

CPP_SRC_1 = r"""
#include <vector>

#include <torch/script.h>

template <
  int K,
  int BLOCK_M,
  int BLOCK_N,
  int BLOCK_K,
  int NUM_STAGES,
  bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p1(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs);

std::vector<at::Tensor> gemm_cstr_p1(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs) {
    auto C = launch_gemm_cstr_p1<2048, 128, 64, 256, 9, true>(ABC, SFAB, outputs);

// #define LAUNCH(K_, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES) \
//   else if (K == K_) C = gemm_launch<K_, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>(A, B, SFA, SFB, C, buf);

//   if (false) {}
//   LAUNCH(16384, 128, 128, 256, 6)
//   LAUNCH( 7168, 128,  64, 256, 8)
//   LAUNCH( 2048, 128,  64, 256, 8)
//   // the rest
//   LAUNCH( 256, 128, 64, 256, 6)
//   LAUNCH( 512, 128, 64, 256, 6)
//   LAUNCH(1536, 128, 64, 256, 6)
//   LAUNCH(2304, 128, 64, 256, 6)

// #undef LAUNCH

  return C;
}

TORCH_LIBRARY(blockscaled_cstr_p1, m) {
  m.def("gemm_cstr_p1((Tensor, Tensor, Tensor)[] ABC, (Tensor, Tensor)[] SFAB, Tensor[] outputs) -> Tensor[]");
  m.impl("gemm_cstr_p1", &gemm_cstr_p1);
}
"""

CUDA_SRC_2  = r"""
using namespace cute;

constexpr int MAX_TMEM_COLS = 512;
constexpr int NUM_SMS = 148;

constexpr int NUM_BLOCKS0 = 72;
constexpr int NUM_BLOCKS1 = 120;


constexpr int GEMM_BLOCK_END0 = NUM_BLOCKS0;
constexpr int GEMM_BLOCK_END1 = GEMM_BLOCK_END0 + NUM_BLOCKS1;

constexpr int NUM_BLOCKS = GEMM_BLOCK_END1;

template<int PROBLEM_SIZE>
__constant__ CUtensorMap A_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ CUtensorMap B_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFA_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFB_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ half* C_ptrs[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ int Ms[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ int Ns[PROBLEM_SIZE];

template <
  int K,
  int BLOCK_M,
  int BLOCK_N,
  int BLOCK_K,
  int NUM_STAGES,
  class SFALayout,
  class SFBLayout,
  bool SWAP_AB,
  int PROBLEM_SIZE,
  int CTA_GROUP
>
__global__
__cluster_dims__(CTA_GROUP, 1, 1)
__launch_bounds__(BLOCK_M + 2 * WARP_SIZE)  // __launch_bounds__(MAX_THREADS_PER_BLOCK, MIN_BLOCKS_PER_MP)
void multi_gemm_kernel_p2() {
  // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-matrix-shape
  static_assert(BLOCK_M == 128); // There doesnt seem to be any code for tiling block_M

  // This layout is used as TMA assumes col major, this is a colasced canonical layout.
  auto asmem_layout = make_layout(make_shape(_256{}, BLOCK_M, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
  // Each member of the cluster will get a cta_rank specific offset
  auto bsmem_layout = make_layout(make_shape(_256{}, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_N / CTA_GROUP));

  // an SF atom is ((32, 4), (16, 4)), the shmem layout is ((128 * SF_ATOMS_PER_BLOCK_K, 4)) for TMA purposes
  // Since we are using CU_TENSOR_MAP_SWIZZLE_NONE, the TMA atom has shape  (8, 128bits) or (8, 16) bytes, so SBO below is 8 * 16
  constexpr int SF_TILE_ROWS = 128;
  constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
  constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns
  constexpr int SF_ROWS_PER_TMEM_COL = 32; // each tmem col has 32 rows - since we are using tcgen05.cp.cta_group::1.32x128b.warpx4

  using Problem = std::tuple<const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, half*, int, int, SFALayout, SFBLayout>;

  auto get_problem = [&](int problem_id) -> Problem {
  
    auto make_sf_layout = [](int rest) {
      constexpr int rest_K = K / 16 / 4;
      return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)), 
                                make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
    };

    const CUtensorMap* A_tmap = &A_tmaps<PROBLEM_SIZE>[problem_id];
    const CUtensorMap* B_tmap = &B_tmaps<PROBLEM_SIZE>[problem_id];
    const CUtensorMap* SFA_tmap = &SFA_tmaps<PROBLEM_SIZE>[problem_id];
    const CUtensorMap* SFB_tmap = &SFB_tmaps<PROBLEM_SIZE>[problem_id];

    half* C = C_ptrs<PROBLEM_SIZE>[problem_id];

    int M = Ms<PROBLEM_SIZE>[problem_id];
    int N = Ns<PROBLEM_SIZE>[problem_id];

    const int rest_M = DIVUP(M, SF_TILE_ROWS);
    const int rest_N = DIVUP(N, SF_TILE_ROWS);
    SFALayout sfa_layout = make_sf_layout(rest_M);
    SFBLayout sfb_layout = make_sf_layout(rest_N);

    return {A_tmap, B_tmap, SFA_tmap, SFB_tmap, C, M, N, sfa_layout, sfb_layout};
  };

  // CTA rank in a cluster
  int cta_rank;
  asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));

  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;

  // set up smem
  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; // BLOCK_M * BLOCK_K elements, each element is 4 bits so divide by 2 to get bytes.
  constexpr int B_size = BLOCK_N * BLOCK_K / 2 / CTA_GROUP; // each CTA only loads half of B
  constexpr int SFA_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K;  // always copy one atom - this is ATOM_SIZE_IN_BYTES * ATOMS_PER_BLOCK_K
  constexpr int SFB_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K;
  constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;

  // set up mbarriers and tmem
  // we have NUM_STAGES mbars for TMA
  //         NUM_STAGES mbars for MMA
  //                  1 mbar  for mainloop
  #pragma nv_diag_suppress static_var_with_dynamic_init
  __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
  __shared__ int tmem_addr[1];  // tmem address is 32-bit

  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;
  const int epilogue_mbar_addr = mainloop_mbar_addr + 2 * 8;

  // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
  // each MMA consumes:
  // - (128, 64) of A -> (128, 4) of SFA -> reshaped as (32, 4', 4) -> 4 tmem columns
  constexpr int SFA_tmem = BLOCK_N * 2; // Double buffer the mma output.
  constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K); // we need BLOCK_K / MMA_K many scale factors for the "tiledMMA"

  if (warp_id == 0 && elect_sync()) {
    for (int prefetch_idx = 0; prefetch_idx < 8; ++prefetch_idx) {
      asm volatile("prefetch.tensormap [%0];" :: "l"(&A_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
      asm volatile("prefetch.tensormap [%0];" :: "l"(&B_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
      asm volatile("prefetch.tensormap [%0];" :: "l"(&SFA_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
      asm volatile("prefetch.tensormap [%0];" :: "l"(&SFB_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
    }
  }  else if (warp_id == 1 && elect_sync()) {
    for (int i = 0; i < NUM_STAGES; i++) {
      mbarrier_init(tma_mbar_addr + i * 8, CTA_GROUP); // one thread in each cluster CTA reports.
      mbarrier_init(mma_mbar_addr + i * 8, 1);
    }
    mbarrier_init(mainloop_mbar_addr, 1); // barrier for first tmem mma buffer
    mbarrier_init(mainloop_mbar_addr + 8, 1); // barrier for second tmem mma buffer

    mbarrier_init(epilogue_mbar_addr, 4 * CTA_GROUP);
    mbarrier_init(epilogue_mbar_addr + 8, 4 * CTA_GROUP);

    asm volatile("fence.mbarrier_init.release.cluster;");  // visible to async proxy
  }
  
  if constexpr (CTA_GROUP > 1) {
    // visible to all threads in a cluster
    asm volatile("barrier.cluster.arrive.relaxed.aligned;");
    asm volatile("barrier.cluster.wait.acquire.aligned;");
  }
  else {
    // visible to all threads in a threadblock
    __syncthreads();
  }

  constexpr int num_iters = K / BLOCK_K;
  const int bid = blockIdx.x;

  auto scheduler = [&](int global_cluster_id) -> std::tuple<int, int, int, int> {
    int problem_id = 0;
    int local_cluster_id = global_cluster_id;

    if (global_cluster_id < GEMM_BLOCK_END0 / CTA_GROUP) {
      problem_id = 0;
      local_cluster_id = global_cluster_id;
    } else if (global_cluster_id < GEMM_BLOCK_END1 / CTA_GROUP) {
      problem_id = 1;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END0 / CTA_GROUP;
    }

    int N = Ns<PROBLEM_SIZE>[problem_id];

    const int grid_n = DIVUP(N, BLOCK_N) ;
    // The cluster Ids are arranged in a (grid_m / 2 , grid_n) matrix in a row major layout. 
    const int bid_m = local_cluster_id / grid_n * CTA_GROUP + cta_rank;
    const int bid_n = local_cluster_id % grid_n;

    const int off_m = bid_m * BLOCK_M;
    const int cluster_off_n = bid_n * BLOCK_N; // N offset for this cluster
    const int off_n = cluster_off_n + cta_rank * (BLOCK_N / CTA_GROUP); // N offset for this CTA
    return {off_m, off_n, cluster_off_n, problem_id};
  };

  // warp-specialization
  if (warp_id == NUM_WARPS - 2 && elect_sync()) {
    // TMA warp

    auto issue_tma = [&](int iter_k, int stage_id, int off_m, int off_n, int problem_id) {
      const int mbar_addr = (tma_mbar_addr + stage_id * 8) & 0xFEFFFFFF;  // CTA0's barrier
      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 problem = get_problem(problem_id);

      const CUtensorMap* A_tmap = std::get<0>(problem);
      const CUtensorMap* B_tmap = std::get<1>(problem);    

      const int M = std::get<5>(problem);
      const int N = std::get<6>(problem);

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

      // issue TMA
      const int off_k = iter_k * BLOCK_K;
      tma_3d_gmem2smem<CTA_GROUP>(A_smem, A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
      tma_3d_gmem2smem<CTA_GROUP>(B_smem, B_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);

      const int rest_m = off_m / 32 / 4; // M atom is (32, 4)
      const int rest_n = off_n / 32 / 4; // N atom is (32, 4)

      const int rest_k = off_k / 16 / 4; // The SF atom is ((32, 4), (16, 4))

      const CUtensorMap* SFA_tmap = std::get<2>(problem);
      const CUtensorMap* SFB_tmap = std::get<3>(problem);

      SFALayout sfa_layout = std::get<7>(problem);
      SFALayout sfb_layout = std::get<8>(problem);
  
      // Divide by 8 since underlying type is INT64
      int sfa_offset = crd2idx(make_coord(make_coord(_0{}, rest_m), make_coord(_0{}, rest_k)), sfa_layout) / 8;
      tma_1d_gmem2smem<CTA_GROUP>(SFA_smem, SFA_tmap, sfa_offset, mbar_addr, cache_A);

      int sfb_offset = crd2idx(make_coord(make_coord(_0{}, rest_n), make_coord(_0{}, rest_k)), sfb_layout) / 8;
      tma_1d_gmem2smem<CTA_GROUP>(SFB_smem, SFB_tmap, sfb_offset, mbar_addr, cache_B);

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

    int mma_phase = 1; // Start with phase 1 as there are no pending MMA operations
    int tma_pipeline_stage = 0;
    for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
      auto [off_m, off_n, dummy, problem_id] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id

      for (int iter_k = 0; iter_k < num_iters; iter_k++) {
        // wait MMA
        // NB when debugging, it will crash w/o this barrier. You cannot keep issuing TMAs w/o the
        // signal the mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase); has completed.
        mbarrier_wait(mma_mbar_addr + tma_pipeline_stage * 8, mma_phase);
        issue_tma(iter_k, tma_pipeline_stage, off_m, off_n, problem_id);

        tma_pipeline_stage = (tma_pipeline_stage + 1) % NUM_STAGES;

        if (tma_pipeline_stage == 0) {
          mma_phase ^= 1;
        }
      }
    }
  }
  else if (warp_id == NUM_WARPS - 1) {
    // allocate tmem
    const int addr = static_cast<int>(__cvta_generic_to_shared(tmem_addr));
    asm volatile("tcgen05.alloc.cta_group::%2.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(addr), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
    if (cta_rank == 0 && elect_sync()) {
      // MMA warp
      // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-instruction-descriptor
      // fp4 MMA doesn't support MMA_M=64. Hence, we will use MMA_M=128 and ignore the rest.
      constexpr int MMA_N = BLOCK_N;
      constexpr int MMA_M = 128 * CTA_GROUP;
      constexpr int16_t cta_mask = (1 << CTA_GROUP) - 1;


      constexpr uint32_t i_desc = (1U << 7U)   // atype=E2M1
                                | (1U << 10U)  // btype=E2M1
                                | ((uint32_t)MMA_N >> 3U << 17U)
                                | ((uint32_t)MMA_M >> 7U << 27U)
                                ;

      int tma_phase = 0;
      int mma_pipeline_stage = 0;
      int epilogue_phase = 1;
      int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result

      for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
        mbarrier_wait(epilogue_mbar_addr + 8 * mainloop_stage, epilogue_phase);

        auto [off_m, off_n, dummy, dummy1] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id

        // Given a k slice of the scale factors which is an atom, the scale factors are stored in groups of four columns in tmem
        // These will be aligned to rows 0, SF_TILE_ROWS, SF_TILE_ROWS * 2, ...
        // And the sf 0..31 will be in index_col=0, 32..63 index_col =1 etc
        // Given an arbitrary row_offset in M or N, in order to find the tmem column offset,
        // Find the residue at the SF tile row granularity, (row_offset % SF_TILE_ROWS) and
        // since each tmem col stores 32 SF, move over by the desired number of columns
        // sf_tmem_offset = (row_offset % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;

        const int M_sf_tmem_offset = (off_m % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
        const int N_sf_tmem_offset = (off_n % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
          // wait TMA
          mbarrier_wait(tma_mbar_addr + mma_pipeline_stage * 8, tma_phase);

          const int A_smem = smem + mma_pipeline_stage * 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;

          // set up shared memory descriptors for A and B
          // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-shared-memory-descriptor
          // 128-byte swizzling. LBO is implied to be 1.
          auto make_desc_AB = [](int addr) -> uint64_t {
            // GMMA::Layout_K_SW128_Atom<cutlass::float_e2m1_t>{} is (_8,_256):(_256,_1) so to get to the start of the next row is 256 / 2 = 128 bytes
            // so to skip 8 rows, we need 8 * 128 bytes.
            const int SBO = 8 * 128; 
            return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
          };
          // no swizzling
          auto make_desc_SF = [](int addr) -> uint64_t {
            // Since we are using no swizzling, the atom has size (8, 16 = 128 / 8). These atoms
            // are tiled (SF_ATOMS_PER_BLOCK_K = 4, 1). So the SBO is 8 * 16
            const int SBO = 8 * 16;
            return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
          };

          // tcgen05.cp -> tcgen05.mma should be pipelined correctly per PTX doc
          // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-memory-consistency-model-pipelined-instructions
          // cutlass issues all of smem->tmem BEFORE mma
          // https://github.com/NVIDIA/cutlass/blob/v4.3.2/include/cutlass/gemm/collective/sm100_blockscaled_mma_warpspecialized.hpp#L1013-L1016
          constexpr uint64_t SF_desc = make_desc_SF(0);
          // The below should really be SF_desc + matrix-descriptor-encode(x) , where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
          // but (x & 0x3FFFF) >> 4 simplifies to x >> 4.
          const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);
          const uint64_t SFB_desc = SF_desc + ((uint64_t)SFB_smem >> 4ULL);

          for (int k = 0; k < BLOCK_K / MMA_K; k++) {
            // The below should really be SFA_desc + (uint64_t)k * matrix-descriptor-encode(512ULL); where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
            // but (x & 0x3FFFF) >> 4 simplifies to x >> 4. So this is SFA_desc + (uint64_t)k * (512ULL >> 4ULL)
            // So we are copying 512 bytes chunks , so adjust the shmem desc accordingly.
            uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL);  // 4 columns, 512 bytes of 128x4 / 32x4x4
            uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);
            // tmem addresses
            // Tensor Memory addresses are 32-bit wide and specify two components.

            // Lane index
            // Column index
            // The layout is as follows:

            // 31 16
            // 15 0
            // Lane index
            // Column index

            // We are using tcgen05.cp.cta_group::1.32x128b.warpx4
            // See https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
            // so the slices ((32, 4), (16, 4)) slices of the scale factors are in the 32 rows (lanes) of tmem

            tcgen05_cp_nvfp4<CTA_GROUP>(SFA_tmem + k * 4, sfa_desc);
            tcgen05_cp_nvfp4<CTA_GROUP>(SFB_tmem + k * 4, sfb_desc);
          }

          // Compare to tiled_mma
          // k1 selects the (BLOCK_M, 256) tile.
          // k2 selects the (BLOCK_M, 64) tile, whose rows are swizzled.
          // NOTE: this doesn't work with BLOCK_N=32, since apparently tcgen05.mma requires SFB_tmem
          // to have 2-column (8-byte) alignment (looks like not documented).
          // HOWEVER: works with a cluster with BLOCK_N = 64 - maybe only the issuing thread has the alignment
          // requirement.
          for (int k1 = 0; k1 < BLOCK_K / 256; k1++) { // BLOCK_K == 256 so this loop only executes once
            for (int k2 = 0; k2 < 256 / MMA_K; k2++) {  // this is the inner loop needed since MMA_K is in general < 256 = BLOCK_K
              // The layout in shmem is the result of the TMA: {256, BLOCK_M, BLOCK_K / 256};
              // Notice that in TMA, it’s implied that the shared memory destination has the natural contiguous layout of the given shape
              // (we don’t specify shared memory stride anywhere).
              // since cuTensorMapEncodeTiled assumes the tensors are col major cute::LayoutLeft
              // The layout is then (256, BLOCK_M, BLOCK_K / 256) : (1, 256, 256 * BLOCK_M) elements 
              // https://gau-nernst.github.io/tcgen05/#decipher-tcgen05
              //
              // Thus given k1 and k2, the coords are (k2 * MMA_K, 0, k1)
              // so the offset in elements is k1 * 256 * BLOCK_M + k2 * MMA_K = k1 * 256 * BLOCK_M + k2 * 64  (MMA_K = 64)
              // so the offset in bytes = k1 * 128 * BLOCK_M + k2 * 32
              // which is the expression below.
      
              // uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
              // crd2idx(make_coord(k2 * MMA_K, 0, k1), asmem_layout) is in elements, divide by 2 to get bytes
              uint64_t a_desc = make_desc_AB(A_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), asmem_layout) / 2);
              uint64_t b_desc = make_desc_AB(B_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), bsmem_layout) / 2);

              // tmem addresses
              // Tensor Memory addresses are 32-bit wide and specify two components.

              // Lane index
              // Column index
              // The layout is as follows:

              // 31 16
              // 15 0
              // Lane index
              // Column index

              int k_sf = k1 * 4 + k2; // the scale factors are in blocks of 4. This is the block idx
              // k_sf is mulitplied by 4 to get us to the actual column.

              const int scale_A_tmem = SFA_tmem + k_sf * 4 + M_sf_tmem_offset;
              const int scale_B_tmem = SFB_tmem + k_sf * 4 + N_sf_tmem_offset;

              const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
              tcgen05_mma_nvfp4<CTA_GROUP>(/*dst tmem*/ BLOCK_N * mainloop_stage, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
            } // for k2
          } // for k1

          // signal MMA done
          // @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
          // tcgen05_commit_mcast does not work if CTA_GROUP = 1
          // tcgen05_commit_mcast<CTA_GROUP>(mma_mbar_addr + mma_pipeline_stage * 8, cta_mask);  // signal MMA done
          // @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
          asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
                      :: "r"(mma_mbar_addr + mma_pipeline_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");

          mma_pipeline_stage = (mma_pipeline_stage + 1) % NUM_STAGES;

          if (mma_pipeline_stage == 0) {
            tma_phase ^= 1;
          }
        }

        // signal mainloop done
        // tcgen05_commit_mcast<CTA_GROUP>(mainloop_mbar_addr + mainloop_stage * 8, cta_mask);  // signal mainloop done
        asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
                    :: "r"(mainloop_mbar_addr + mainloop_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
        mainloop_stage = (mainloop_stage + 1) % 2; 

        if (mainloop_stage == 0) {
          epilogue_phase ^= 1;
        }
      }  // for bid
    } // cta_rank == 0 && elect_sync()
  } // warp_id == NUM_WARPS - 1
  else if (tid < BLOCK_M) {
    // epilogue warps
    int mainloop_phase = 0;
    int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result
    
    for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
      auto result = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
      int off_m = std::get<0>(result);
      int cluster_off_n = std::get<2>(result);
      int problem_id = std::get<3>(result);

      auto problem = get_problem(problem_id);

      const int M = std::get<5>(problem);
      const int N = std::get<6>(problem);

      half* C_ptr = std::get<4>(problem);

      // wait mainloop
      mbarrier_wait(mainloop_mbar_addr + mainloop_stage * 8, mainloop_phase);
      asm volatile("tcgen05.fence::after_thread_sync;");

      int tmem_buffer_offset = BLOCK_N * mainloop_stage;

      auto epilogue_M_major = [&]() {
        // C is M-major
        constexpr int WIDTH = std::min(BLOCK_N, 64);  // using 128 might be slower

        for (int n = 0; n < (BLOCK_N + WIDTH - 1)/ WIDTH; n++) {
          float tmp[WIDTH];  // if WIDTH=128, we are using 128 registers here
          if constexpr (WIDTH == 128) tcgen05_ld_32x32bx128(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
          if constexpr (WIDTH == 64) tcgen05_ld_32x32bx64(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
          if constexpr (WIDTH == 32) tcgen05_ld_32x32bx32(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
          asm volatile("tcgen05.wait::ld.sync.aligned;");

          for (int i = 0; i < WIDTH; i++) {
            // NB row and col are transposed.
            const int row = cluster_off_n + n * WIDTH + i;
            const int col = off_m + tid;
            if (row < N && col < M) {
              C_ptr[row * M + col] = __float2half(tmp[i]);
            }
          }
        }
      };
      auto epilogue_N_major = [&]() {
        // C is N-major
        for (int m = 0; m < 32 / 16; m++) {
          float tmp[BLOCK_N / 2];
          if constexpr (BLOCK_N == 128) tcgen05_ld_16x256bx16(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
          if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
          if constexpr (BLOCK_N == 32) tcgen05_ld_16x256bx4(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
          asm volatile("tcgen05.wait::ld.sync.aligned;");

          for (int i = 0; i < (BLOCK_N + 7)/ 8; i++) {
            // TODO replace this with cute tensors.
            const int row = off_m + warp_id * 32 + m * 16 + lane_id / 4;
            const int col = cluster_off_n + i * 8 + (lane_id % 4) * 2;
            if (row < M) {
              reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 0], tmp[i * 4 + 1]});
            }
            if (row + 8 < M) {
              reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 2], tmp[i * 4 + 3]});
            }
          }
        }
      };

      if constexpr (SWAP_AB) {
        epilogue_M_major();
      } else {
        epilogue_N_major();
      }
      if (elect_sync()) {
        // // signal Epilogue done
        // constexpr int16_t cta_mask = 3;
        // tcgen05_commit_mcast<CTA_GROUP>(epilogue_mbar_addr + 8 * mainloop_stage, cta_mask);  // signal Epilogue done
        const int mbar_addr = (epilogue_mbar_addr + mainloop_stage * 8) & 0xFEFFFFFF;
        asm volatile("mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];" :: "r"(mbar_addr) : "memory");
      }
      
      mainloop_stage = (mainloop_stage + 1) % 2; 
      if (mainloop_stage == 0) {
        mainloop_phase ^= 1;
      }
    } // for bid
  } //epilogue warps

  // All warps.
  if constexpr (CTA_GROUP > 1) {
    // Dont exit w/o the peer CTA.
    asm volatile("barrier.cluster.arrive.release.aligned;");
    asm volatile("barrier.cluster.wait.acquire.aligned;");
  } else {
    __syncthreads();  // all threads finish reading data from tmem
  }
  // deallocate tmem. tmem address should be 0.
  if (warp_id == 0) {
    asm volatile("tcgen05.dealloc.cta_group::%2.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
  }
}

template <
  int K,
  int BLOCK_M,
  int BLOCK_N,
  int BLOCK_K,
  int NUM_STAGES,
  bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p2(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs) {

  constexpr int CTA_GROUP = 2;
  static_assert(BLOCK_K % 256 == 0);

  constexpr int rest_K = K / 16 / 4;
  constexpr int grid = NUM_SMS;
  constexpr int tb_size = BLOCK_M + 2 * WARP_SIZE;
  constexpr int AB_size = (BLOCK_M + BLOCK_N / CTA_GROUP) * (BLOCK_K / 2);
  constexpr int SFAB_size = 128 * (BLOCK_K / 16) * 2;
  constexpr int smem_size = (AB_size + SFAB_size) * NUM_STAGES;

  // This layout is used as TMA assumes col major, this is a colasced canonical layout.
  auto layout_smem_A = make_layout(make_shape(256, BLOCK_M, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
  // Each member of the cluster will get a cta_rank specific offset
  auto layout_smem_B = make_layout(make_shape(256, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_N / CTA_GROUP));

  constexpr int SF_TILE_ROWS = 128;
  constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
  constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns

  auto make_sf_layout = [&](int rest) {
    return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)), 
                              make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
  };
  
  // This is for the type only
  auto layout_sfa = make_sf_layout(16); 
  auto layout_sfb = make_sf_layout(16); 

  auto populate_kernel_config = [&ABC, &SFAB, &outputs](CUtensorMap* A_tmaps, CUtensorMap* B_tmaps, CUtensorMap* SFA_tmaps, CUtensorMap* SFB_tmaps,
    half** C_ptrs, int* Ms, int* Ns, auto problem_size_const) {
    constexpr int PROBLEM_SIZE = decltype(problem_size_const)::value;

    for (int workItem = 0; workItem < PROBLEM_SIZE; ++workItem) {
      std::tuple<at::Tensor, at::Tensor, at::Tensor> abc_tuple = ABC[workItem]; 
      auto A = std::get<0>(abc_tuple);
      auto B = std::get<1>(abc_tuple);

      std::tuple<at::Tensor, at::Tensor> sfab_tuple = SFAB[workItem]; 
      auto SFA = std::get<0>(sfab_tuple);
      auto SFB = std::get<1>(sfab_tuple);

      at::Tensor C = outputs[workItem];

      const int M = A.size(0);
      const int N = B.size(0);

      auto A_ptr   = reinterpret_cast<const char *>(A.data_ptr());
      auto B_ptr  = reinterpret_cast<const char *>(B.data_ptr());
      auto SFA_ptr = reinterpret_cast<const char *>(SFA.data_ptr());
      auto SFB_ptr = reinterpret_cast<const char *>(SFB.data_ptr());
      C_ptrs[workItem]   = reinterpret_cast<half *>(C.data_ptr());

      int new_M = M;
      int new_N = N;

      if constexpr (SWAP_AB) {
        std::swap(A_ptr, B_ptr);
        std::swap(SFA_ptr, SFB_ptr);
        std::swap(new_M, new_N);
      }

      Ms[workItem] = new_M;
      Ns[workItem] = new_N;

      init_AB_tmap(&A_tmaps[workItem], A_ptr, new_M, K, BLOCK_M, BLOCK_K);
      init_AB_tmap(&B_tmaps[workItem], B_ptr, new_N, K, BLOCK_N / CTA_GROUP, BLOCK_K);

      // CUtensorMap SFA_tmap0, SFB_tmap0;
      // SF Atom is ((32, 4), (16, 4))
      const int rest_M = DIVUP(new_M, SF_TILE_ROWS);
      const int rest_N = DIVUP(new_N, SF_TILE_ROWS);

      auto new_M_padded = SF_TILE_ROWS * rest_M; // 128 is the number of rows in the atom.
      auto new_N_padded = SF_TILE_ROWS * rest_N;

      init_SF_tmap(&SFA_tmaps[workItem], SFA_ptr, new_M_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
      init_SF_tmap(&SFB_tmaps[workItem], SFB_ptr, new_N_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
    }
  };

  constexpr int PROBLEM_SIZE = 2;
  CUtensorMap host_A_tmaps[PROBLEM_SIZE];
  CUtensorMap host_B_tmaps[PROBLEM_SIZE];
  CUtensorMap host_SFA_tmaps[PROBLEM_SIZE];
  CUtensorMap host_SFB_tmaps[PROBLEM_SIZE];
  half* host_C_ptrs[PROBLEM_SIZE];
  int host_Ms[PROBLEM_SIZE];
  int host_Ns[PROBLEM_SIZE];

  populate_kernel_config(host_A_tmaps, host_B_tmaps, host_SFA_tmaps, host_SFB_tmaps, host_C_ptrs, host_Ms, host_Ns, std::integral_constant<int, PROBLEM_SIZE>{});

  cudaMemcpyToSymbol(A_tmaps<PROBLEM_SIZE>, host_A_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(B_tmaps<PROBLEM_SIZE>, host_B_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(SFA_tmaps<PROBLEM_SIZE>, host_SFA_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(SFB_tmaps<PROBLEM_SIZE>, host_SFB_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(C_ptrs<PROBLEM_SIZE>, host_C_ptrs, PROBLEM_SIZE * sizeof(half*));
  cudaMemcpyToSymbol(Ms<PROBLEM_SIZE>, host_Ms, PROBLEM_SIZE * sizeof(int));
  cudaMemcpyToSymbol(Ns<PROBLEM_SIZE>, host_Ns, PROBLEM_SIZE * sizeof(int));

  auto this_kernel = multi_gemm_kernel_p2<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES,
                            decltype(layout_sfa), decltype(layout_sfb),
                            SWAP_AB, PROBLEM_SIZE, CTA_GROUP>;

  if (smem_size > 48'000)
    cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);

  this_kernel<<<grid, tb_size, smem_size>>>();
  return {outputs[0], outputs[1]};

}


template std::vector<at::Tensor> launch_gemm_cstr_p2<4096, 128, 64, 256, 9, true>(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs);
"""

CPP_SRC_2 = r"""
#include <vector>

#include <torch/script.h>

template <
  int K,
  int BLOCK_M,
  int BLOCK_N,
  int BLOCK_K,
  int NUM_STAGES,
  bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p2(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs);

std::vector<at::Tensor> gemm_cstr_p2(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs) {
    auto C = launch_gemm_cstr_p2<4096, 128, 64, 256, 9, true>(ABC, SFAB, outputs);

  return C;
}

TORCH_LIBRARY(blockscaled_cstr_p2, m) {
  m.def("gemm_cstr_p2((Tensor, Tensor, Tensor)[] ABC, (Tensor, Tensor)[] SFAB, Tensor[] outputs) -> Tensor[]");
  m.impl("gemm_cstr_p2", &gemm_cstr_p2);
}
"""

CUDA_SRC_3  = r"""
using namespace cute;

constexpr int MAX_TMEM_COLS = 512;
constexpr int NUM_SMS = 148;

constexpr int NUM_BLOCKS0 = 64;
constexpr int NUM_BLOCKS1 = 192;


constexpr int GEMM_BLOCK_END0 = NUM_BLOCKS0;
constexpr int GEMM_BLOCK_END1 = GEMM_BLOCK_END0 + NUM_BLOCKS1;

constexpr int NUM_BLOCKS = GEMM_BLOCK_END1;

template<int PROBLEM_SIZE>
__constant__ CUtensorMap A_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ CUtensorMap B_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFA_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFB_tmaps[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ half* C_ptrs[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ int Ms[PROBLEM_SIZE];

template<int PROBLEM_SIZE>
__constant__ int Ns[PROBLEM_SIZE];

template <
  int K,
  int BLOCK_M,
  int BLOCK_N,
  int BLOCK_K,
  int NUM_STAGES,
  class SFALayout,
  class SFBLayout,
  bool SWAP_AB,
  int PROBLEM_SIZE,
  int CTA_GROUP
>
__global__
__cluster_dims__(CTA_GROUP, 1, 1)
__launch_bounds__(BLOCK_M + 2 * WARP_SIZE)  // __launch_bounds__(MAX_THREADS_PER_BLOCK, MIN_BLOCKS_PER_MP)
void multi_gemm_kernel_p3() {
  // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-matrix-shape
  static_assert(BLOCK_M == 128); // There doesnt seem to be any code for tiling block_M

  // This layout is used as TMA assumes col major, this is a colasced canonical layout.
  auto asmem_layout = make_layout(make_shape(_256{}, BLOCK_M, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
  // Each member of the cluster will get a cta_rank specific offset
  auto bsmem_layout = make_layout(make_shape(_256{}, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_N / CTA_GROUP));

  // an SF atom is ((32, 4), (16, 4)), the shmem layout is ((128 * SF_ATOMS_PER_BLOCK_K, 4)) for TMA purposes
  // Since we are using CU_TENSOR_MAP_SWIZZLE_NONE, the TMA atom has shape  (8, 128bits) or (8, 16) bytes, so SBO below is 8 * 16
  constexpr int SF_TILE_ROWS = 128;
  constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
  constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns
  constexpr int SF_ROWS_PER_TMEM_COL = 32; // each tmem col has 32 rows - since we are using tcgen05.cp.cta_group::1.32x128b.warpx4

  using Problem = std::tuple<const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, half*, int, int, SFALayout, SFBLayout>;

  auto get_problem = [&](int problem_id) -> Problem {
  
    auto make_sf_layout = [](int rest) {
      constexpr int rest_K = K / 16 / 4;
      return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)), 
                                make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
    };

    const CUtensorMap* A_tmap = &A_tmaps<PROBLEM_SIZE>[problem_id];
    const CUtensorMap* B_tmap = &B_tmaps<PROBLEM_SIZE>[problem_id];
    const CUtensorMap* SFA_tmap = &SFA_tmaps<PROBLEM_SIZE>[problem_id];
    const CUtensorMap* SFB_tmap = &SFB_tmaps<PROBLEM_SIZE>[problem_id];

    half* C = C_ptrs<PROBLEM_SIZE>[problem_id];

    int M = Ms<PROBLEM_SIZE>[problem_id];
    int N = Ns<PROBLEM_SIZE>[problem_id];

    const int rest_M = DIVUP(M, SF_TILE_ROWS);
    const int rest_N = DIVUP(N, SF_TILE_ROWS);
    SFALayout sfa_layout = make_sf_layout(rest_M);
    SFBLayout sfb_layout = make_sf_layout(rest_N);

    return {A_tmap, B_tmap, SFA_tmap, SFB_tmap, C, M, N, sfa_layout, sfb_layout};
  };

  // CTA rank in a cluster
  int cta_rank;
  asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));

  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;

  // set up smem
  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; // BLOCK_M * BLOCK_K elements, each element is 4 bits so divide by 2 to get bytes.
  constexpr int B_size = BLOCK_N * BLOCK_K / 2 / CTA_GROUP; // each CTA only loads half of B
  constexpr int SFA_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K;  // always copy one atom - this is ATOM_SIZE_IN_BYTES * ATOMS_PER_BLOCK_K
  constexpr int SFB_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K;
  constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;

  // set up mbarriers and tmem
  // we have NUM_STAGES mbars for TMA
  //         NUM_STAGES mbars for MMA
  //                  1 mbar  for mainloop
  #pragma nv_diag_suppress static_var_with_dynamic_init
  __shared__ int64_t mbars[NUM_STAGES * 2 + 1];
  __shared__ int tmem_addr[1];  // tmem address is 32-bit

  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;
  const int epilogue_mbar_addr = mainloop_mbar_addr + 2 * 8;

  // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
  // each MMA consumes:
  // - (128, 64) of A -> (128, 4) of SFA -> reshaped as (32, 4', 4) -> 4 tmem columns
  constexpr int SFA_tmem = BLOCK_N * 2; // Double buffer the mma output.
  constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K); // we need BLOCK_K / MMA_K many scale factors for the "tiledMMA"

  if (warp_id == 0 && elect_sync()) {
    for (int prefetch_idx = 0; prefetch_idx < 8; ++prefetch_idx) {
      asm volatile("prefetch.tensormap [%0];" :: "l"(&A_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
      asm volatile("prefetch.tensormap [%0];" :: "l"(&B_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
      asm volatile("prefetch.tensormap [%0];" :: "l"(&SFA_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
      asm volatile("prefetch.tensormap [%0];" :: "l"(&SFB_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
    }
  }  else if (warp_id == 1 && elect_sync()) {
    for (int i = 0; i < NUM_STAGES; i++) {
      mbarrier_init(tma_mbar_addr + i * 8, CTA_GROUP); // one thread in each cluster CTA reports.
      mbarrier_init(mma_mbar_addr + i * 8, 1);
    }
    mbarrier_init(mainloop_mbar_addr, 1); // barrier for first tmem mma buffer
    mbarrier_init(mainloop_mbar_addr + 8, 1); // barrier for second tmem mma buffer

    mbarrier_init(epilogue_mbar_addr, 4 * CTA_GROUP);
    mbarrier_init(epilogue_mbar_addr + 8, 4 * CTA_GROUP);

    asm volatile("fence.mbarrier_init.release.cluster;");  // visible to async proxy
  }
  
  if constexpr (CTA_GROUP > 1) {
    // visible to all threads in a cluster
    asm volatile("barrier.cluster.arrive.relaxed.aligned;");
    asm volatile("barrier.cluster.wait.acquire.aligned;");
  }
  else {
    // visible to all threads in a threadblock
    __syncthreads();
  }

  constexpr int num_iters = K / BLOCK_K;
  const int bid = blockIdx.x;

  auto scheduler = [&](int global_cluster_id) -> std::tuple<int, int, int, int> {
    int problem_id = 0;
    int local_cluster_id = global_cluster_id;

    if (global_cluster_id < GEMM_BLOCK_END0 / CTA_GROUP) {
      problem_id = 0;
      local_cluster_id = global_cluster_id;
    } else if (global_cluster_id < GEMM_BLOCK_END1 / CTA_GROUP) {
      problem_id = 1;
      local_cluster_id = global_cluster_id - GEMM_BLOCK_END0 / CTA_GROUP;
    }

    int N = Ns<PROBLEM_SIZE>[problem_id];

    const int grid_n = DIVUP(N, BLOCK_N) ;
    // The cluster Ids are arranged in a (grid_m / 2 , grid_n) matrix in a row major layout. 
    const int bid_m = local_cluster_id / grid_n * CTA_GROUP + cta_rank;
    const int bid_n = local_cluster_id % grid_n;

    const int off_m = bid_m * BLOCK_M;
    const int cluster_off_n = bid_n * BLOCK_N; // N offset for this cluster
    const int off_n = cluster_off_n + cta_rank * (BLOCK_N / CTA_GROUP); // N offset for this CTA
    return {off_m, off_n, cluster_off_n, problem_id};
  };

  // warp-specialization
  if (warp_id == NUM_WARPS - 2 && elect_sync()) {
    // TMA warp

    auto issue_tma = [&](int iter_k, int stage_id, int off_m, int off_n, int problem_id) {
      const int mbar_addr = (tma_mbar_addr + stage_id * 8) & 0xFEFFFFFF;  // CTA0's barrier
      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 problem = get_problem(problem_id);

      const CUtensorMap* A_tmap = std::get<0>(problem);
      const CUtensorMap* B_tmap = std::get<1>(problem);    

      const int M = std::get<5>(problem);
      const int N = std::get<6>(problem);

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

      // issue TMA
      const int off_k = iter_k * BLOCK_K;
      tma_3d_gmem2smem<CTA_GROUP>(A_smem, A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
      tma_3d_gmem2smem<CTA_GROUP>(B_smem, B_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);

      const int rest_m = off_m / 32 / 4; // M atom is (32, 4)
      const int rest_n = off_n / 32 / 4; // N atom is (32, 4)

      const int rest_k = off_k / 16 / 4; // The SF atom is ((32, 4), (16, 4))

      const CUtensorMap* SFA_tmap = std::get<2>(problem);
      const CUtensorMap* SFB_tmap = std::get<3>(problem);

      SFALayout sfa_layout = std::get<7>(problem);
      SFALayout sfb_layout = std::get<8>(problem);
  
      // Divide by 8 since underlying type is INT64
      int sfa_offset = crd2idx(make_coord(make_coord(_0{}, rest_m), make_coord(_0{}, rest_k)), sfa_layout) / 8;
      tma_1d_gmem2smem<CTA_GROUP>(SFA_smem, SFA_tmap, sfa_offset, mbar_addr, cache_A);

      int sfb_offset = crd2idx(make_coord(make_coord(_0{}, rest_n), make_coord(_0{}, rest_k)), sfb_layout) / 8;
      tma_1d_gmem2smem<CTA_GROUP>(SFB_smem, SFB_tmap, sfb_offset, mbar_addr, cache_B);

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

    int mma_phase = 1; // Start with phase 1 as there are no pending MMA operations
    int tma_pipeline_stage = 0;
    for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
      auto [off_m, off_n, dummy, problem_id] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id

      for (int iter_k = 0; iter_k < num_iters; iter_k++) {
        // wait MMA
        // NB when debugging, it will crash w/o this barrier. You cannot keep issuing TMAs w/o the
        // signal the mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase); has completed.
        mbarrier_wait(mma_mbar_addr + tma_pipeline_stage * 8, mma_phase);
        issue_tma(iter_k, tma_pipeline_stage, off_m, off_n, problem_id);

        tma_pipeline_stage = (tma_pipeline_stage + 1) % NUM_STAGES;

        if (tma_pipeline_stage == 0) {
          mma_phase ^= 1;
        }
      }
    }
  }
  else if (warp_id == NUM_WARPS - 1) {
    // allocate tmem
    const int addr = static_cast<int>(__cvta_generic_to_shared(tmem_addr));
    asm volatile("tcgen05.alloc.cta_group::%2.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(addr), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
    if (cta_rank == 0 && elect_sync()) {
      // MMA warp
      // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-instruction-descriptor
      // fp4 MMA doesn't support MMA_M=64. Hence, we will use MMA_M=128 and ignore the rest.
      constexpr int MMA_N = BLOCK_N;
      constexpr int MMA_M = 128 * CTA_GROUP;
      constexpr int16_t cta_mask = (1 << CTA_GROUP) - 1;


      constexpr uint32_t i_desc = (1U << 7U)   // atype=E2M1
                                | (1U << 10U)  // btype=E2M1
                                | ((uint32_t)MMA_N >> 3U << 17U)
                                | ((uint32_t)MMA_M >> 7U << 27U)
                                ;

      int tma_phase = 0;
      int mma_pipeline_stage = 0;
      int epilogue_phase = 1;
      int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result

      for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
        mbarrier_wait(epilogue_mbar_addr + 8 * mainloop_stage, epilogue_phase);

        auto [off_m, off_n, dummy, dummy1] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id


        const int M_sf_tmem_offset = (off_m % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
        const int N_sf_tmem_offset = (off_n % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;

        for (int iter_k = 0; iter_k < num_iters; iter_k++) {
          // wait TMA
          mbarrier_wait(tma_mbar_addr + mma_pipeline_stage * 8, tma_phase);

          const int A_smem = smem + mma_pipeline_stage * 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;

          // set up shared memory descriptors for A and B
          // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-shared-memory-descriptor
          // 128-byte swizzling. LBO is implied to be 1.
          auto make_desc_AB = [](int addr) -> uint64_t {
            // GMMA::Layout_K_SW128_Atom<cutlass::float_e2m1_t>{} is (_8,_256):(_256,_1) so to get to the start of the next row is 256 / 2 = 128 bytes
            // so to skip 8 rows, we need 8 * 128 bytes.
            const int SBO = 8 * 128; 
            return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
          };
          // no swizzling
          auto make_desc_SF = [](int addr) -> uint64_t {
            // Since we are using no swizzling, the atom has size (8, 16 = 128 / 8). These atoms
            // are tiled (SF_ATOMS_PER_BLOCK_K = 4, 1). So the SBO is 8 * 16
            const int SBO = 8 * 16;
            return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
          };

          // tcgen05.cp -> tcgen05.mma should be pipelined correctly per PTX doc
          // https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-memory-consistency-model-pipelined-instructions
          // cutlass issues all of smem->tmem BEFORE mma
          // https://github.com/NVIDIA/cutlass/blob/v4.3.2/include/cutlass/gemm/collective/sm100_blockscaled_mma_warpspecialized.hpp#L1013-L1016
          constexpr uint64_t SF_desc = make_desc_SF(0);
          // The below should really be SF_desc + matrix-descriptor-encode(x) , where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
          // but (x & 0x3FFFF) >> 4 simplifies to x >> 4.
          const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);
          const uint64_t SFB_desc = SF_desc + ((uint64_t)SFB_smem >> 4ULL);

          for (int k = 0; k < BLOCK_K / MMA_K; k++) {
            // The below should really be SFA_desc + (uint64_t)k * matrix-descriptor-encode(512ULL); where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
            // but (x & 0x3FFFF) >> 4 simplifies to x >> 4. So this is SFA_desc + (uint64_t)k * (512ULL >> 4ULL)
            // So we are copying 512 bytes chunks , so adjust the shmem desc accordingly.
            uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL);  // 4 columns, 512 bytes of 128x4 / 32x4x4
            uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);
            // tmem addresses
            // Tensor Memory addresses are 32-bit wide and specify two components.

            // Lane index
            // Column index
            // The layout is as follows:

            // 31 16
            // 15 0
            // Lane index
            // Column index

            // We are using tcgen05.cp.cta_group::1.32x128b.warpx4
            // See https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
            // so the slices ((32, 4), (16, 4)) slices of the scale factors are in the 32 rows (lanes) of tmem

            tcgen05_cp_nvfp4<CTA_GROUP>(SFA_tmem + k * 4, sfa_desc);
            tcgen05_cp_nvfp4<CTA_GROUP>(SFB_tmem + k * 4, sfb_desc);
          }

          // Compare to tiled_mma
          // k1 selects the (BLOCK_M, 256) tile.
          // k2 selects the (BLOCK_M, 64) tile, whose rows are swizzled.
          // NOTE: this doesn't work with BLOCK_N=32, since apparently tcgen05.mma requires SFB_tmem
          // to have 2-column (8-byte) alignment (looks like not documented).
          // HOWEVER: works with a cluster with BLOCK_N = 64 - maybe only the issuing thread has the alignment
          // requirement.
          for (int k1 = 0; k1 < BLOCK_K / 256; k1++) { // BLOCK_K == 256 so this loop only executes once
            for (int k2 = 0; k2 < 256 / MMA_K; k2++) {  // this is the inner loop needed since MMA_K is in general < 256 = BLOCK_K
      
              // uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
              // crd2idx(make_coord(k2 * MMA_K, 0, k1), asmem_layout) is in elements, divide by 2 to get bytes
              uint64_t a_desc = make_desc_AB(A_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), asmem_layout) / 2);
              uint64_t b_desc = make_desc_AB(B_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), bsmem_layout) / 2);


              int k_sf = k1 * 4 + k2; // the scale factors are in blocks of 4. This is the block idx
              // k_sf is mulitplied by 4 to get us to the actual column.

              const int scale_A_tmem = SFA_tmem + k_sf * 4 + M_sf_tmem_offset;
              const int scale_B_tmem = SFB_tmem + k_sf * 4 + N_sf_tmem_offset;

              const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
              tcgen05_mma_nvfp4<CTA_GROUP>(/*dst tmem*/ BLOCK_N * mainloop_stage, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
            } // for k2
          } // for k1

          // signal MMA done
          // @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
          // tcgen05_commit_mcast does not work if CTA_GROUP = 1
          // tcgen05_commit_mcast<CTA_GROUP>(mma_mbar_addr + mma_pipeline_stage * 8, cta_mask);  // signal MMA done
          // @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
          asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
                      :: "r"(mma_mbar_addr + mma_pipeline_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");

          mma_pipeline_stage = (mma_pipeline_stage + 1) % NUM_STAGES;

          if (mma_pipeline_stage == 0) {
            tma_phase ^= 1;
          }
        }

        // signal mainloop done
        // tcgen05_commit_mcast<CTA_GROUP>(mainloop_mbar_addr + mainloop_stage * 8, cta_mask);  // signal mainloop done
        asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
                    :: "r"(mainloop_mbar_addr + mainloop_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
        mainloop_stage = (mainloop_stage + 1) % 2; 

        if (mainloop_stage == 0) {
          epilogue_phase ^= 1;
        }
      }  // for bid
    } // cta_rank == 0 && elect_sync()
  } // warp_id == NUM_WARPS - 1
  else if (tid < BLOCK_M) {
    // epilogue warps
    int mainloop_phase = 0;
    int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result
    
    for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
      auto result = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
      int off_m = std::get<0>(result);
      int cluster_off_n = std::get<2>(result);
      int problem_id = std::get<3>(result);

      auto problem = get_problem(problem_id);

      const int M = std::get<5>(problem);
      const int N = std::get<6>(problem);

      half* C_ptr = std::get<4>(problem);

      // wait mainloop
      mbarrier_wait(mainloop_mbar_addr + mainloop_stage * 8, mainloop_phase);
      asm volatile("tcgen05.fence::after_thread_sync;");

      int tmem_buffer_offset = BLOCK_N * mainloop_stage;

      auto epilogue_M_major = [&]() {
        // C is M-major
        constexpr int WIDTH = std::min(BLOCK_N, 64);  // using 128 might be slower

        for (int n = 0; n < (BLOCK_N + WIDTH - 1)/ WIDTH; n++) {
          float tmp[WIDTH];  // if WIDTH=128, we are using 128 registers here
          if constexpr (WIDTH == 128) tcgen05_ld_32x32bx128(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
          if constexpr (WIDTH == 64) tcgen05_ld_32x32bx64(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
          if constexpr (WIDTH == 32) tcgen05_ld_32x32bx32(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
          asm volatile("tcgen05.wait::ld.sync.aligned;");

          for (int i = 0; i < WIDTH; i++) {
            // NB row and col are transposed.
            const int row = cluster_off_n + n * WIDTH + i;
            const int col = off_m + tid;
            if (row < N && col < M) {
              C_ptr[row * M + col] = __float2half(tmp[i]);
            }
          }
        }
      };
      auto epilogue_N_major = [&]() {
        // C is N-major
        for (int m = 0; m < 32 / 16; m++) {
          float tmp[BLOCK_N / 2];
          if constexpr (BLOCK_N == 128) tcgen05_ld_16x256bx16(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
          if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
          if constexpr (BLOCK_N == 32) tcgen05_ld_16x256bx4(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
          asm volatile("tcgen05.wait::ld.sync.aligned;");

          for (int i = 0; i < (BLOCK_N + 7)/ 8; i++) {
            // TODO replace this with cute tensors.
            const int row = off_m + warp_id * 32 + m * 16 + lane_id / 4;
            const int col = cluster_off_n + i * 8 + (lane_id % 4) * 2;
            if (row < M) {
              reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 0], tmp[i * 4 + 1]});
            }
            if (row + 8 < M) {
              reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 2], tmp[i * 4 + 3]});
            }
          }
        }
      };

      if constexpr (SWAP_AB) {
        epilogue_M_major();
      } else {
        epilogue_N_major();
      }
      if (elect_sync()) {
        // // signal Epilogue done
        // constexpr int16_t cta_mask = 3;
        // tcgen05_commit_mcast<CTA_GROUP>(epilogue_mbar_addr + 8 * mainloop_stage, cta_mask);  // signal Epilogue done
        const int mbar_addr = (epilogue_mbar_addr + mainloop_stage * 8) & 0xFEFFFFFF;
        asm volatile("mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];" :: "r"(mbar_addr) : "memory");
      }
      
      mainloop_stage = (mainloop_stage + 1) % 2; 
      if (mainloop_stage == 0) {
        mainloop_phase ^= 1;
      }
    } // for bid
  } //epilogue warps

  // All warps.
  if constexpr (CTA_GROUP > 1) {
    // Dont exit w/o the peer CTA.
    asm volatile("barrier.cluster.arrive.release.aligned;");
    asm volatile("barrier.cluster.wait.acquire.aligned;");
  } else {
    __syncthreads();  // all threads finish reading data from tmem
  }
  // deallocate tmem. tmem address should be 0.
  if (warp_id == 0) {
    asm volatile("tcgen05.dealloc.cta_group::%2.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
  }
}

template <
  int K,
  int BLOCK_M,
  int BLOCK_N,
  int BLOCK_K,
  int NUM_STAGES,
  bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p3(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs) {

  constexpr int CTA_GROUP = 2;
  static_assert(BLOCK_K % 256 == 0);

  constexpr int rest_K = K / 16 / 4;
  constexpr int grid = NUM_SMS;
  constexpr int tb_size = BLOCK_M + 2 * WARP_SIZE;
  constexpr int AB_size = (BLOCK_M + BLOCK_N / CTA_GROUP) * (BLOCK_K / 2);
  constexpr int SFAB_size = 128 * (BLOCK_K / 16) * 2;
  constexpr int smem_size = (AB_size + SFAB_size) * NUM_STAGES;

  // This layout is used as TMA assumes col major, this is a colasced canonical layout.
  auto layout_smem_A = make_layout(make_shape(256, BLOCK_M, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
  // Each member of the cluster will get a cta_rank specific offset
  auto layout_smem_B = make_layout(make_shape(256, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_N / CTA_GROUP));

  constexpr int SF_TILE_ROWS = 128;
  constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
  constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns

  auto make_sf_layout = [&](int rest) {
    return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)), 
                              make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
  };
  
  // This is for the type only
  auto layout_sfa = make_sf_layout(16); 
  auto layout_sfb = make_sf_layout(16); 

  auto populate_kernel_config = [&ABC, &SFAB, &outputs](CUtensorMap* A_tmaps, CUtensorMap* B_tmaps, CUtensorMap* SFA_tmaps, CUtensorMap* SFB_tmaps,
    half** C_ptrs, int* Ms, int* Ns, auto problem_size_const) {
    constexpr int PROBLEM_SIZE = decltype(problem_size_const)::value;

    for (int workItem = 0; workItem < PROBLEM_SIZE; ++workItem) {
      std::tuple<at::Tensor, at::Tensor, at::Tensor> abc_tuple = ABC[workItem]; 
      auto A = std::get<0>(abc_tuple);
      auto B = std::get<1>(abc_tuple);

      std::tuple<at::Tensor, at::Tensor> sfab_tuple = SFAB[workItem]; 
      auto SFA = std::get<0>(sfab_tuple);
      auto SFB = std::get<1>(sfab_tuple);

      at::Tensor C = outputs[workItem];

      const int M = A.size(0);
      const int N = B.size(0);

      auto A_ptr   = reinterpret_cast<const char *>(A.data_ptr());
      auto B_ptr  = reinterpret_cast<const char *>(B.data_ptr());
      auto SFA_ptr = reinterpret_cast<const char *>(SFA.data_ptr());
      auto SFB_ptr = reinterpret_cast<const char *>(SFB.data_ptr());
      C_ptrs[workItem]   = reinterpret_cast<half *>(C.data_ptr());

      int new_M = M;
      int new_N = N;

      if constexpr (SWAP_AB) {
        std::swap(A_ptr, B_ptr);
        std::swap(SFA_ptr, SFB_ptr);
        std::swap(new_M, new_N);
      }

      Ms[workItem] = new_M;
      Ns[workItem] = new_N;

      init_AB_tmap(&A_tmaps[workItem], A_ptr, new_M, K, BLOCK_M, BLOCK_K);
      init_AB_tmap(&B_tmaps[workItem], B_ptr, new_N, K, BLOCK_N / CTA_GROUP, BLOCK_K);

      // CUtensorMap SFA_tmap0, SFB_tmap0;
      // SF Atom is ((32, 4), (16, 4))
      const int rest_M = DIVUP(new_M, SF_TILE_ROWS);
      const int rest_N = DIVUP(new_N, SF_TILE_ROWS);

      auto new_M_padded = SF_TILE_ROWS * rest_M; // 128 is the number of rows in the atom.
      auto new_N_padded = SF_TILE_ROWS * rest_N;

      init_SF_tmap(&SFA_tmaps[workItem], SFA_ptr, new_M_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
      init_SF_tmap(&SFB_tmaps[workItem], SFB_ptr, new_N_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
    }
  };

  constexpr int PROBLEM_SIZE = 2;
  CUtensorMap host_A_tmaps[PROBLEM_SIZE];
  CUtensorMap host_B_tmaps[PROBLEM_SIZE];
  CUtensorMap host_SFA_tmaps[PROBLEM_SIZE];
  CUtensorMap host_SFB_tmaps[PROBLEM_SIZE];
  half* host_C_ptrs[PROBLEM_SIZE];
  int host_Ms[PROBLEM_SIZE];
  int host_Ns[PROBLEM_SIZE];

  populate_kernel_config(host_A_tmaps, host_B_tmaps, host_SFA_tmaps, host_SFB_tmaps, host_C_ptrs, host_Ms, host_Ns, std::integral_constant<int, PROBLEM_SIZE>{});

  cudaMemcpyToSymbol(A_tmaps<PROBLEM_SIZE>, host_A_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(B_tmaps<PROBLEM_SIZE>, host_B_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(SFA_tmaps<PROBLEM_SIZE>, host_SFA_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(SFB_tmaps<PROBLEM_SIZE>, host_SFB_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
  cudaMemcpyToSymbol(C_ptrs<PROBLEM_SIZE>, host_C_ptrs, PROBLEM_SIZE * sizeof(half*));
  cudaMemcpyToSymbol(Ms<PROBLEM_SIZE>, host_Ms, PROBLEM_SIZE * sizeof(int));
  cudaMemcpyToSymbol(Ns<PROBLEM_SIZE>, host_Ns, PROBLEM_SIZE * sizeof(int));

  auto this_kernel = multi_gemm_kernel_p3<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES,
                            decltype(layout_sfa), decltype(layout_sfb),
                            SWAP_AB, PROBLEM_SIZE, CTA_GROUP>;

  if (smem_size > 48'000)
    cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);

  this_kernel<<<grid, tb_size, smem_size>>>();
  return {outputs[0], outputs[1]};

}


template std::vector<at::Tensor> launch_gemm_cstr_p3<1536, 128, 64, 256, 9, true>(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs);
"""

CPP_SRC_3 = r"""
#include <vector>

#include <torch/script.h>

template <
  int K,
  int BLOCK_M,
  int BLOCK_N,
  int BLOCK_K,
  int NUM_STAGES,
  bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p3(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs);

std::vector<at::Tensor> gemm_cstr_p3(
  c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
  c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
  c10::List<at::Tensor> outputs) {
    auto C = launch_gemm_cstr_p3<1536, 128, 64, 256, 9, true>(ABC, SFAB, outputs);

  return C;
}

TORCH_LIBRARY(blockscaled_cstr_p3, m) {
  m.def("gemm_cstr_p3((Tensor, Tensor, Tensor)[] ABC, (Tensor, Tensor)[] SFAB, Tensor[] outputs) -> Tensor[]");
  m.impl("gemm_cstr_p3", &gemm_cstr_p3);
}
"""
ON_VERDA = False
optional_args = {"extra_include_paths": ['/root/cutlass/include/', '/root/cutlass/tools/util/include/',]} if ON_VERDA else {}

module = load_inline(
    "some_module_name_0",
    cpp_sources=CPP_SRC_0,
    cuda_sources=COMMON_CU + CUDA_SRC_0,
    verbose=True,
    is_python_module=False,
    no_implicit_headers=True,
    extra_cuda_cflags=[
        "-O3",
        "-gencode=arch=compute_100a,code=sm_100a",
        "--use_fast_math",
        "--expt-relaxed-constexpr",
        "--relocatable-device-code=false",
        "-lineinfo",
        "-Xptxas=-v",
    ],
    extra_ldflags=["-lcuda"],
    **optional_args
)

module = load_inline(
    "some_module_name_1",
    cpp_sources=CPP_SRC_1,
    cuda_sources=COMMON_CU + CUDA_SRC_1,
    verbose=True,
    is_python_module=False,
    no_implicit_headers=True,
    extra_cuda_cflags=[
        "-O3",
        "-gencode=arch=compute_100a,code=sm_100a",
        "--use_fast_math",
        "--expt-relaxed-constexpr",
        "--relocatable-device-code=false",
        "-lineinfo",
        "-Xptxas=-v",
    ],
    extra_ldflags=["-lcuda"],
    **optional_args
)

module = load_inline(
    "some_module_name_2",
    cpp_sources=CPP_SRC_2,
    cuda_sources=COMMON_CU + CUDA_SRC_2,
    verbose=True,
    is_python_module=False,
    no_implicit_headers=True,
    extra_cuda_cflags=[
        "-O3",
        "-gencode=arch=compute_100a,code=sm_100a",
        "--use_fast_math",
        "--expt-relaxed-constexpr",
        "--relocatable-device-code=false",
        "-lineinfo",
        "-Xptxas=-v",
    ],
    extra_ldflags=["-lcuda"],
    **optional_args
)

module = load_inline(
    "some_module_name_3",
    cpp_sources=CPP_SRC_3,
    cuda_sources=COMMON_CU + CUDA_SRC_3,
    verbose=True,
    is_python_module=False,
    no_implicit_headers=True,
    extra_cuda_cflags=[
        "-O3",
        "-gencode=arch=compute_100a,code=sm_100a",
        "--use_fast_math",
        "--expt-relaxed-constexpr",
        "--relocatable-device-code=false",
        "-lineinfo",
        "-Xptxas=-v",
    ],
    extra_ldflags=["-lcuda"],
    **optional_args
)
################################################################################
########################## Test custom kernel ##################################
################################################################################

# Scaling factor vector size
sf_vec_size = 16

# Helper function for ceiling division
def ceil_div(a, b):
    return (a + b - 1) // b

# Helper function to convert scale factor tensor to blocked format
# Helper function to convert scale factor tensor to blocked format
def to_blocked(input_matrix):
    rows, cols = input_matrix.shape

    # Please ensure rows and cols are multiples of 128 and 4 respectively
    n_row_blocks = ceil_div(rows, 128)
    n_col_blocks = ceil_div(cols, 4)
    padded_rows = n_row_blocks * 128
    padded_cols = n_col_blocks * 4

    # Pad the input matrix if necessary
    if padded_rows != rows or padded_cols != cols:
        padded = torch.nn.functional.pad(
            input_matrix,
            (0, padded_cols - cols, 0, padded_rows - rows),
            mode="constant",
            value=0,
        )
    else:
        padded = input_matrix
    blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
    rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)

    return rearranged.flatten()


def ref_kernel(
    data: input_t,
) -> output_t:
    """
    PyTorch reference implementation of NVFP4 block-scaled group GEMM.
    """
    abc_tensors, sfasfb_tensors, _, problem_sizes = data
    
    result_tensors = []
    for i, (
        (a_ref, b_ref, c_ref),
        (sfa_ref, sfb_ref),
        (m, n, k, l),
    ) in enumerate(
        zip(
            abc_tensors,
            sfasfb_tensors,
            problem_sizes,
        )
    ):
        for l_idx in range(l):
            # Convert the scale factor tensor to blocked format
            scale_a = to_blocked(sfa_ref[:, :, l_idx])
            scale_b = to_blocked(sfb_ref[:, :, l_idx])
            # (m, k) @ (n, k).T -> (m, n)
            res = torch._scaled_mm(
                a_ref[:, :, l_idx].view(torch.float4_e2m1fn_x2),
                b_ref[:, :, l_idx].transpose(0, 1).view(torch.float4_e2m1fn_x2),
                scale_a.cuda(),
                scale_b.cuda(),
                bias=None,
                out_dtype=torch.float16,
            )
            c_ref[:, :, l_idx] = res
        result_tensors.append((c_ref))
    return result_tensors


# Helper function to prepare the scale factor tensors for both reference
# kernel and customize kernel. The customized data layout can be found in:
# https://docs.nvidia.com/cuda/cublas/index.html?highlight=fp4#d-block-scaling-factors-layout
def create_reordered_scale_factor_tensor(l, mn, k, ref_f8_tensor):
    sf_k = ceil_div(k, sf_vec_size)
    atom_m = (32, 4)
    atom_k = 4
    mma_shape = (
        l,  # batch size
        ceil_div(mn, atom_m[0] * atom_m[1]),
        ceil_div(sf_k, atom_k),
        atom_m[0],
        atom_m[1],
        atom_k,
    )
    # Create the reordered scale factor tensor (32, 4, rest_m, 4, rest_k, l) on GPU.
    mma_permute_order = (3, 4, 1, 5, 2, 0)
    # Generate a random int8 tensor, then convert to float8_e4m3fn
    rand_int_tensor = torch.randint(1, 3, mma_shape, dtype=torch.int8, device='cuda')
    reordered_f8_tensor = rand_int_tensor.to(dtype=torch.float8_e4m3fn)
    # Permute according to mma_permute_order
    reordered_f8_tensor = reordered_f8_tensor.permute(*mma_permute_order)

    # Move ref_f8_tensor to GPU if not already there
    if ref_f8_tensor.device.type == 'cpu':
        ref_f8_tensor = ref_f8_tensor.cuda()

    # GPU-side vectorized reordering (replaces slow CPU nested loops)
    # Create index grids for all dimensions
    i_idx = torch.arange(mn, device='cuda')
    j_idx = torch.arange(sf_k, device='cuda')
    b_idx = torch.arange(l, device='cuda')
    
    # Create meshgrid for all combinations of (i, j, b)
    i_grid, j_grid, b_grid = torch.meshgrid(i_idx, j_idx, b_idx, indexing='ij')
    
    # Calculate target indices in vectorized manner
    mm = i_grid // (atom_m[0] * atom_m[1])
    mm32 = i_grid % atom_m[0]
    mm4 = (i_grid % 128) // atom_m[0]
    kk = j_grid // atom_k
    kk4 = j_grid % atom_k
    
    # Perform the reordering with advanced indexing (all on GPU)
    reordered_f8_tensor[mm32, mm4, mm, kk4, kk, b_grid] = ref_f8_tensor[i_grid, j_grid, b_grid]
    
    return reordered_f8_tensor


def _create_fp4_tensors(l, mn, k):
    # generate uint8 tensor, then convert to float4e2m1fn_x2 data type
    # generate all bit patterns
    ref_i8 = torch.randint(255, size=(l, mn, k // 2), dtype=torch.uint8, device="cuda") # * 0 + 34 # Remove comment to make inputs one 34 = b'0010'0010

    # for each nibble, only keep the sign bit and 2 LSBs
    # the possible values are [-1.5, -1, -0.5, 0, +0.5, +1, +1.5]
    ref_i8 = ref_i8 & 0b1011_1011
    return ref_i8.permute(1, 2, 0).view(torch.float4_e2m1fn_x2)


def generate_input(
    m: tuple,
    n: tuple,
    k: tuple,
    g: int,
    seed: int,
):
    """
    Generate input tensors for NVFP4 block-scaled group GEMM. 
    Each group can have different m, n, k, l.
    
    Args:
        problem_sizes: List of tuples (m, n, k, l) for each problem
        m: Number of rows in matrix A
        n: Number of columns in matrix B
        k: Number of columns in A and rows of B
        l: Batch size, always is 1
        groups: Number of groups
        seed: Random seed for reproducibility
    
    Returns:
        Tuple of (list(tuple(a, b, c)), list(tuple(sfa, sfb)), list(tuple(sfa_reordered, sfb_reordered)), list(tuple(m, n, k, l))) where each group has its own a, b, c, sfa, sfb.
            a: [m, k, l] - Input matrix in torch.float4e2m1fn_x2 data type
            b: [n, k, l] - Input matrix in torch.float4e2m1fn_x2 data type
            sfa: [m, k // 16, l] - Input scale factors in torch.float8e4m3fn data type
            sfb: [n, k // 16, l] - Input scale factors in torch.float8e4m3fn data type
            sfa_reordered: [32, 4, rest_m, 4, rest_k, l] - Input scale factors in torch.float8e4m3fn data type
            sfb_reordered: [32, 4, rest_n, 4, rest_k, l] - Input scale factors in torch.float8e4m3fn data type
            c: [m, n, l] - Output matrix in torch.float16 data type
    """
    torch.manual_seed(seed)
    
    abc_tensors = []
    sfasfb_tensors = []
    sfasfb_reordered_tensors = []
    problem_sizes = []
    l = 1
    # Generate a, b, c, sfa, sfb tensors for all groups
    for group_idx in range(g):
        mi = m[group_idx]
        ni = n[group_idx]
        ki = k[group_idx]
        a_ref = _create_fp4_tensors(l, mi, ki)
        b_ref = _create_fp4_tensors(l, ni, ki)

        c_ref = torch.randn((l, mi, ni), dtype=torch.float16, device="cuda").permute(
            1, 2, 0
        )

        sf_k = ceil_div(ki, sf_vec_size)
         
        sfa_ref_cpu_random = torch.randint(
            1, 3, (l, mi, sf_k), dtype=torch.int8
        ) # * 0 + 1 # Remove comment to make sfs one
        sfa_ref_cpu = sfa_ref_cpu_random.to(dtype=torch.float8_e4m3fn).permute(1, 2, 0)
        sfb_ref_cpu_random = torch.randint(
            1, 3, (l, ni, sf_k), dtype=torch.int8
        ) # * 0 + 1 # Remove comment to make sfs one
        sfb_ref_cpu = sfb_ref_cpu_random.to(dtype=torch.float8_e4m3fn).permute(1, 2, 0)

        sfa_reordered = create_reordered_scale_factor_tensor(l, mi, ki, sfa_ref_cpu)
        sfb_reordered = create_reordered_scale_factor_tensor(l, ni, ki, sfb_ref_cpu)

        abc_tensors.append((a_ref, b_ref, c_ref))
        sfasfb_tensors.append((sfa_ref_cpu, sfb_ref_cpu))
        sfasfb_reordered_tensors.append((sfa_reordered, sfb_reordered))
        problem_sizes.append((mi, ni, ki, l))
    return (abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes)

################################################################################
########################## End Test custom kernel ##############################
################################################################################

start = 0
BIG_BUFFER = torch.zeros(int(2e10), dtype=torch.float16, device="cuda")

def allocate(c: torch.Tensor):
    if not USE_ALLOCATE:
      return c
    global start
    end = start + c.numel()
    buf = BIG_BUFFER[start:end].as_strided(c.shape, c.stride())
    start = end
    return buf

USE_ALLOCATE = False

result_tensor_func = allocate if USE_ALLOCATE else lambda x: x 


def custom_kernel(data: input_t) -> output_t:
    """
    Reference implementation of block-scale fp4 group gemm
    Args:
        data: list of tuples (abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes) where:
            abc_tensors: list of tuples (a, b, c) where 
                a is torch.Tensor[float4e2m1fn_x2] of shape [m, k // 2, l]
                b is torch.Tensor[float4e2m1fn_x2] of shape [n, k // 2, l]
                c is torch.Tensor[float16] of shape [m, n, l]
            sfasfb_tensors: list of tuples (sfa, sfb) where 
                sfa is torch.Tensor[float8_e4m3fnuz] of shape [m, k // 16, l]
                sfb is torch.Tensor[float8_e4m3fnuz] of shape [n, k // 16, l]
            sfasfb_reordered_tensors: list of tuples (sfa_reordered, sfb_reordered) where 
                sfa_reordered is torch.Tensor[float8_e4m3fnuz] of shape [32, 4, rest_m, 4, rest_k, l]
                sfb_reordered is torch.Tensor[float8_e4m3fnuz] of shape [32, 4, rest_n, 4, rest_k, l]
            problem_sizes: list of tuples (m, n, k, l)
        each group has its own a, b, c, sfa, sfb with different m, n, k, l problem sizes
        l should always be 1 for each group.
    Returns:
        list of tuples (c) where c is torch.Tensor[float16] of shape [m, n, l]
    """
    abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes = data

    output_tensors = []
    for item in abc_tensors:
        _, _, c_ref = item
        output_tensors.append(result_tensor_func(c_ref))

    if abc_tensors[0][0].shape[1] == 7168 // 2:
      return torch.ops.blockscaled_cstr_p0.gemm_cstr_p0(abc_tensors, sfasfb_reordered_tensors, output_tensors)
    elif abc_tensors[0][0].shape[1] == 2048 // 2:
      return torch.ops.blockscaled_cstr_p1.gemm_cstr_p1(abc_tensors, sfasfb_reordered_tensors, output_tensors)
    elif abc_tensors[0][0].shape[1] == 4096 // 2:
      return torch.ops.blockscaled_cstr_p2.gemm_cstr_p2(abc_tensors, sfasfb_reordered_tensors, output_tensors)
    elif abc_tensors[0][0].shape[1] == 1536 // 2:
      return torch.ops.blockscaled_cstr_p3.gemm_cstr_p3(abc_tensors, sfasfb_reordered_tensors, output_tensors)

    return ref_kernel(data)

if __name__ == '__main__':
  benchmarks = [[8, [80, 176, 128, 72, 64, 248, 96, 160], [4096, 4096, 4096, 4096, 4096, 4096, 4096, 4096], [7168, 7168, 7168, 7168, 7168, 7168, 7168, 7168]],\
                [8, [40, 76, 168, 72, 164, 148, 196, 160], [7168, 7168, 7168, 7168, 7168, 7168, 7168, 7168], [2048, 2048, 2048, 2048, 2048, 2048, 2048, 2048]],\
                [2, [192, 320], [3072, 3072], [4096, 4096]],\
                [2, [128, 384], [4096, 4096], [1536, 1536]]]

  for benchmark in benchmarks:
    G, M, N, K = benchmark
    print(f"{M=}, {N=}, {K=}")
    input_data = generate_input(M, N, K, G, 1234)

    raw_results = custom_kernel(input_data)
    ref_results = ref_kernel(input_data)

    print(raw_results[0] - ref_results[0])
scrolls · 3661 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON