Skip to content
KernelIndex
Search⌘K

submission 383109

hekailove · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-modal-nvfp4-dual-gemm-383109?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 dual GEMMsuite of 4 cases
NVIDIA B200
15.1µs
#66 of 161
2026-01-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6ca9dab93a0f90aa3098b5739ac4238643459a1b2710b8cefad62165906a092b
license declaredunknown
license concludedunknown
authorshekailove
imported2026-08-15

Techniques

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

fused-epiloguetemplate <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CACHE_MODE, int PAD, int SWZ, int EPILOGUE_MODE, int MBAR_MODE>
mbarrier__device__ __forceinline__ void mbarrier_init(int mbar_addr, int count) {
shared-memory__device__ __forceinline__ int out_smem_col(int col_h2, int row) {
tcgen05asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
tma"cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "
vector-width = half2__device__ __forceinline__ uint32_t half2_as_u32(half2 x) {

Kernel source

submission.py1782 lines



import os

os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
os.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":16:8")
os.environ.setdefault("CUDA_MODULE_LOADING", "LAZY")

import torch
from torch.utils.cpp_extension import load_inline


_CUDA_SRC = r"""
#include <cuda.h>
#include <cudaTypedefs.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>

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

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

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

#ifndef ENABLE_CLK64_DIAG
#define ENABLE_CLK64_DIAG 0
#endif

#ifndef ENABLE_V4_EPILOGUE
#define ENABLE_V4_EPILOGUE 0
#endif

#if ENABLE_CLK64_DIAG
__device__ __align__(16) unsigned long long g_clk64_diag[65536];
#endif

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

template <int SWZ>
__device__ __forceinline__ int out_smem_col(int col_h2, int row) {
  if constexpr (SWZ) return col_h2 ^ ((row & 7) << 2);
  return col_h2;
}

// 精度策略:使用 `ex2.approx` + `rcp.approx` 近似 sigmoid;目标是使最终 fp16 输出满足 rtol/atol=1e-3 门槛。
__device__ __forceinline__ float fast_sigmoid(float x) {
  float ex2;
  asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex2) : "f"((-x) * 1.4426950408889634f));
  float denom = 1.0f + ex2;
  float rcp;
  asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(rcp) : "f"(denom));
  return rcp;
}

__device__ __forceinline__ uint32_t elect_sync() {
  uint32_t pred = 0;
  asm volatile(
    "{\n\t"
    ".reg .pred %%px;\n\t"
    "elect.sync _|%%px, %1;\n\t"
    "@%%px mov.s32 %0, 1;\n\t"
    "}\n\t"
    : "+r"(pred)
    : "r"(0xFFFFFFFF)
  );
  return pred;
}

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

__device__ __forceinline__ void mbarrier_wait(int mbar_addr, int phase) {
  uint32_t ticks = 0x989680;
  asm volatile(
    "{\n\t"
    ".reg .pred P1;\n\t"
    ".reg .u32 B;\n\t"
    "mov.u32 B, 1;\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"
    "nanosleep.u32 B;\n\t"
    "shl.b32 B, B, 1;\n\t"
    "min.u32 B, B, 32;\n\t"
    "bra.uni LAB_WAIT;\n\t"
    "DONE:\n\t"
    "}\n\t"
    :: "r"(mbar_addr), "r"(phase), "r"(ticks)
  );
}

__device__ __forceinline__ void tma_gmem2smem(int dst, const void *src, int size, int mbar_addr, uint64_t cache_policy) {
  asm volatile(
    "cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "
    "[%0], [%1], %2, [%3], %4;"
    :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy)
  );
}

__device__ __forceinline__ void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy) {
  asm volatile(
    "cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "
    "[%0], [%1, {%2, %3, %4}], [%5], %6;"
    :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "l"(cache_policy)
    : "memory"
  );
}

__device__ __forceinline__ void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
  asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
}

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

__device__ __forceinline__ uint32_t half2_as_u32(half2 x) {
  union {
    half2 h;
    uint32_t u;
  } v;
  v.h = x;
  return v.u;
}

__device__ __forceinline__ void stg_v2_b32(const void *p, uint32_t a, uint32_t b) {
  asm volatile("st.global.v2.b32 [%0], {%1, %2};" :: "l"(p), "r"(a), "r"(b) : "memory");
}

__device__ __forceinline__ void stg_v4_b32(const void *p, uint32_t a, uint32_t b, uint32_t c, uint32_t d) {
  asm volatile("st.global.v4.b32 [%0], {%1, %2, %3, %4};" :: "l"(p), "r"(a), "r"(b), "r"(c), "r"(d) : "memory");
}

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

template <const char *SHAPE_, const char *NUM_>
__device__ __forceinline__ 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_)
  );
}

__device__ __forceinline__ void tcgen05_ld_16x256bx8(float *tmp, int row, int col) {
  tcgen05_ld_32regs<SHAPE::_16x256b, NUM::x8>(tmp, row, col);
}

template <const char *SHAPE_, const char *NUM_>
__device__ __forceinline__ 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_)
  );
}

__device__ __forceinline__ void tcgen05_ld_16x256bx16(float *tmp, int row, int col) {
  tcgen05_ld_64regs<SHAPE::_16x256b, NUM::x16>(tmp, row, col);
}

static inline void ck_cu(CUresult err) {
  if (err == CUDA_SUCCESS) return;
  const char *msg = nullptr;
  if (cuGetErrorString(err, &msg) != CUDA_SUCCESS) msg = "cu err";
  TORCH_CHECK(false, msg);
}

static inline void init_AB_tmap(
  CUtensorMap *tmap,
  const char *ptr,
  uint64_t global_h, uint64_t global_w,
  uint32_t shared_h, uint32_t shared_w,
  CUtensorMapL2promotion l2_promo
) {
  // Host-side 固定开销优化:缓存 cuTensorMapEncodeTiled 的结果,key 包含指针与维度/tiling。
  // 注意:缓存命中仅在 key 完全一致时复用;否则必须重新 encode 以保证正确性。
  constexpr int TM_CACHE_SIZE = 8;
  struct TMCacheEntry {
    const void *ptr;
    uint64_t global_h;
    uint64_t global_w;
    uint32_t shared_h;
    uint32_t shared_w;
    uint32_t l2_promo;
    CUtensorMap tmap;
    uint32_t valid;
  };
  static TMCacheEntry cache[TM_CACHE_SIZE];
  #pragma unroll
  for (int i = 0; i < TM_CACHE_SIZE; i++) {
    const TMCacheEntry &e = cache[i];
    if (e.valid &&
        e.ptr == (const void *)ptr &&
        e.global_h == global_h &&
        e.global_w == global_w &&
        e.shared_h == shared_h &&
        e.shared_w == shared_w &&
        e.l2_promo == (uint32_t)l2_promo) {
      *tmap = e.tmap;
      return;
    }
  }

  constexpr uint32_t rank = 3;
  uint64_t globalDim[rank]       = {256, global_h, global_w / 256};
  uint64_t globalStrides[rank-1] = {global_w / 2, 128};
  uint32_t boxDim[rank]          = {256, shared_h, shared_w / 256};
  uint32_t elementStrides[rank]  = {1, 1, 1};

  CUtensorMap tmp;
  auto err = cuTensorMapEncodeTiled(
    &tmp,
    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,
    l2_promo,
    CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
  );
  ck_cu(err);
  *tmap = tmp;

  static uint32_t victim = 0;
  const uint32_t slot = victim++ & (TM_CACHE_SIZE - 1);
  cache[slot].ptr = (const void *)ptr;
  cache[slot].global_h = global_h;
  cache[slot].global_w = global_w;
  cache[slot].shared_h = shared_h;
  cache[slot].shared_w = shared_w;
  cache[slot].l2_promo = (uint32_t)l2_promo;
  cache[slot].tmap = tmp;
  cache[slot].valid = 1;
}

template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CACHE_MODE>
__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void gemm_f32_kernel(
  const __grid_constant__ CUtensorMap A_tmap,
  const __grid_constant__ CUtensorMap B_tmap,
  const char *SFA_ptr,
  const char *SFB_ptr,
  float *C_ptr,
  int M, int N
) {
  const int tid = threadIdx.x;
  const int bid = blockIdx.y;

  const int lane_id = tid & 31;
  const int warp_id = tid >> 5;

  const int grid_m = M / BLOCK_M;
  const int grid_n = N / BLOCK_N;
  const int bid_m = bid / grid_n;
  const int bid_n = bid - bid_m * grid_n;

	  const int off_m = bid_m * BLOCK_M;
	  const int off_n = bid_n * BLOCK_N;

#if ENABLE_CLK64_DIAG
	  const int diag_base = bid * 8;
#endif

	  constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

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

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

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

  if (warp_id == 0 && elect_sync()) {
    for (int i = 0; i < NUM_STAGES * 2 + 1; i++) mbarrier_init(tma_mbar_addr + i * 8, 1);
    asm volatile("fence.mbarrier_init.release.cluster;");
  } else if (warp_id == 1) {
    asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(BLOCK_N * 2));
  }
  __syncthreads();

  constexpr int num_iters = K / BLOCK_K;

  if (warp_id == NUM_WARPS - 2 && elect_sync()) {
    constexpr uint64_t cache_A = (CACHE_MODE == 0) ? EVICT_FIRST : EVICT_LAST;
    constexpr uint64_t cache_B = (CACHE_MODE == 0) ? EVICT_LAST : EVICT_FIRST;

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

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

      const int rest_k = K / 16 / 4;
      const char *SFA_src = SFA_ptr + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
      const char *SFB_src = SFB_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
      tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_A);
      tma_gmem2smem(SFB_smem, SFB_src, SFB_size, mbar_addr, cache_B);

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

    constexpr int PRELOAD = (num_iters < NUM_STAGES) ? num_iters : NUM_STAGES;
    for (int iter_k = 0; iter_k < PRELOAD; iter_k++) issue_tma(iter_k, iter_k);
    for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
      const int stage_id = iter_k % NUM_STAGES;
      const int mma_phase = (iter_k / NUM_STAGES - 1) & 1;
      mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
      issue_tma(iter_k, stage_id);
    }
	  } else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
#if ENABLE_CLK64_DIAG
	    const unsigned long long t0 = clock64();
#endif
	    constexpr int MMA_N = BLOCK_N;
	    constexpr int MMA_M = 128;
	    constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)MMA_N >> 3U << 17U) | ((uint32_t)MMA_M >> 7U << 27U);

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

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

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

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

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

      for (int k1 = 0; k1 < BLOCK_K / 256; k1++)
        for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
          uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
          uint64_t b_desc = make_desc_AB(B_smem + k1 * BLOCK_N * 128 + k2 * 32);

          const int k_sf = k1 * 4 + k2;
          const int scale_A_tmem = SFA_tmem + k_sf * 4;
          int scale_B_tmem;
          if constexpr (BLOCK_N == 128) {
            scale_B_tmem = SFB_tmem + k_sf * 4;
          } else {
            scale_B_tmem = SFB_tmem + k_sf * 4 + (bid_n & 1) * (BLOCK_N / 32);
          }

          const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
          tcgen05_mma_nvfp4(0, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
        }

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

	    asm volatile(
	      "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
	      :: "r"(mainloop_mbar_addr)
	      : "memory"
	    );
#if ENABLE_CLK64_DIAG
	    g_clk64_diag[diag_base + 1] = clock64() - t0;
#endif
	  } else if (tid < BLOCK_M) {
    if (elect_sync()) mbarrier_wait(mainloop_mbar_addr, 0);
    __syncwarp();
    asm volatile("tcgen05.fence::after_thread_sync;");

    for (int mm = 0; mm < 2; mm++) {
      float tmp[BLOCK_N / 2];
      if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + mm * 16, 0);
      else tcgen05_ld_16x256bx16(tmp, warp_id * 32 + mm * 16, 0);
      asm volatile("tcgen05.wait::ld.sync.aligned;");

      #pragma unroll
      for (int i = 0; i < BLOCK_N / 8; i++) {
        const int row = off_m + warp_id * 32 + mm * 16 + lane_id / 4;
        const int col = off_n + i * 8 + (lane_id & 3) * 2;
        reinterpret_cast<float2 *>(C_ptr + (row + 0) * N + col)[0] = float2{tmp[i * 4 + 0], tmp[i * 4 + 1]};
        reinterpret_cast<float2 *>(C_ptr + (row + 8) * N + col)[0] = float2{tmp[i * 4 + 2], tmp[i * 4 + 3]};
      }
    }

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

template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CACHE_MODE, int PAD, int SWZ>
__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void gemm_silu_mul_kernel(
  const __grid_constant__ CUtensorMap A_tmap,
  const __grid_constant__ CUtensorMap B_tmap,
  const char *SFA_ptr,
  const char *SFB_ptr,
  const float *G1_ptr,
  half *Out_ptr,
  int M, int N
) {
  const int tid = threadIdx.x;
  const int bid = blockIdx.y;

  const int lane_id = tid & 31;
  const int warp_id = tid >> 5;

  const int grid_m = M / BLOCK_M;
  const int grid_n = N / BLOCK_N;
  const int bid_m = bid / grid_n;
  const int bid_n = bid - bid_m * grid_n;

  const int off_m = bid_m * BLOCK_M;
  const int off_n = bid_n * BLOCK_N;

#if ENABLE_CLK64_DIAG
  const int diag_base = bid * 8;
#endif

  constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

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

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

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

  if (warp_id == 0 && elect_sync()) {
    for (int i = 0; i < NUM_STAGES * 2 + 1; i++) mbarrier_init(tma_mbar_addr + i * 8, 1);
    asm volatile("fence.mbarrier_init.release.cluster;");
  } else if (warp_id == 1) {
    asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(BLOCK_N * 2));
  }
  __syncthreads();

	  constexpr int num_iters = K / BLOCK_K;
	
	  if (warp_id == NUM_WARPS - 2 && elect_sync()) {
#if ENABLE_CLK64_DIAG
	    const unsigned long long t0 = clock64();
#endif
	    constexpr uint64_t cache_A = (CACHE_MODE == 0) ? EVICT_LAST : EVICT_FIRST;
	    constexpr uint64_t cache_B = (CACHE_MODE == 0) ? EVICT_FIRST : EVICT_LAST;

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

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

      const int rest_k = K / 16 / 4;
      const char *SFA_src = SFA_ptr + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
      const char *SFB_src = SFB_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
      tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_A);
      tma_gmem2smem(SFB_smem, SFB_src, SFB_size, mbar_addr, cache_B);

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

	    constexpr int PRELOAD = (num_iters < NUM_STAGES) ? num_iters : NUM_STAGES;
	    for (int iter_k = 0; iter_k < PRELOAD; iter_k++) issue_tma(iter_k, iter_k);
	    for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
	      const int stage_id = iter_k % NUM_STAGES;
	      const int mma_phase = (iter_k / NUM_STAGES - 1) & 1;
	      mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
	      issue_tma(iter_k, stage_id);
	    }
#if ENABLE_CLK64_DIAG
	    g_clk64_diag[diag_base + 0] = clock64() - t0;
#endif
	  } else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
    constexpr int MMA_N = BLOCK_N;
    constexpr int MMA_M = 128;
    constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)MMA_N >> 3U << 17U) | ((uint32_t)MMA_M >> 7U << 27U);

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

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

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

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

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

      for (int k1 = 0; k1 < BLOCK_K / 256; k1++)
        for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
          uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
          uint64_t b_desc = make_desc_AB(B_smem + k1 * BLOCK_N * 128 + k2 * 32);

          const int k_sf = k1 * 4 + k2;
          const int scale_A_tmem = SFA_tmem + k_sf * 4;
          int scale_B_tmem;
          if constexpr (BLOCK_N == 128) {
            scale_B_tmem = SFB_tmem + k_sf * 4;
          } else {
            scale_B_tmem = SFB_tmem + k_sf * 4 + (bid_n & 1) * (BLOCK_N / 32);
          }

          const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
          tcgen05_mma_nvfp4(0, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
        }

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

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

    half2 *out_smem = reinterpret_cast<half2 *>(smem_ptr);
    constexpr int OUT_H2_STRIDE = BLOCK_N / 2 + PAD;

    for (int mm = 0; mm < 2; mm++) {
      float tmp[BLOCK_N / 2];
      if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + mm * 16, 0);
      else tcgen05_ld_16x256bx16(tmp, warp_id * 32 + mm * 16, 0);
      asm volatile("tcgen05.wait::ld.sync.aligned;");

      #pragma unroll
      for (int i = 0; i < BLOCK_N / 8; i++) {
        const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
        const int lc = i * 8 + (lane_id & 3) * 2;
        const int row = off_m + lr;
        const int col = off_n + lc;
        const float2 x0 = reinterpret_cast<const float2 *>(G1_ptr + (row + 0) * N + col)[0];
        const float2 x8 = reinterpret_cast<const float2 *>(G1_ptr + (row + 8) * N + col)[0];

        float2 o0;
        float2 o8;
        float s;
        s = fast_sigmoid(x0.x);
        o0.x = (x0.x * s) * tmp[i * 4 + 0];
        s = fast_sigmoid(x0.y);
        o0.y = (x0.y * s) * tmp[i * 4 + 1];
        s = fast_sigmoid(x8.x);
        o8.x = (x8.x * s) * tmp[i * 4 + 2];
        s = fast_sigmoid(x8.y);
        o8.y = (x8.y * s) * tmp[i * 4 + 3];

        const int hc = lc >> 1;
        out_smem[(lr + 0) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 0)] = __float22half2_rn(o0);
        out_smem[(lr + 8) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 8)] = __float22half2_rn(o8);
      }
    }

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

    constexpr int WRITER_WARPS = BLOCK_M / WARP_SIZE;
    constexpr int ROWS_PER_WARP = BLOCK_M / WRITER_WARPS;
    const int r_base = warp_id * ROWS_PER_WARP;
    for (int rr = 0; rr < ROWS_PER_WARP; rr++) {
      const int r = r_base + rr;
      const int row = off_m + r;
      half2 *dst = reinterpret_cast<half2 *>(Out_ptr + row * N + off_n);
      const half2 *src = out_smem + r * OUT_H2_STRIDE;
      if constexpr (BLOCK_N == 64) {
        dst[lane_id] = src[out_smem_col<SWZ>(lane_id, r)];
      } else {
        dst[lane_id] = src[out_smem_col<SWZ>(lane_id, r)];
        dst[lane_id + 32] = src[out_smem_col<SWZ>(lane_id + 32, r)];
      }
    }
  }
}

template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CACHE_MODE, int PAD, int SWZ, int EPILOGUE_MODE, int MBAR_MODE>
__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void dual_gemm_silu_kernel(
  const __grid_constant__ CUtensorMap A_tmap,
  const __grid_constant__ CUtensorMap B1_tmap,
  const __grid_constant__ CUtensorMap B2_tmap,
  const char *SFA_ptr,
  const char *SFB1_ptr,
  const char *SFB2_ptr,
  half *Out_ptr,
  int M, int N
) {
  const int tid = threadIdx.x;

  const int lane_id = tid & 31;
  const int warp_id = tid >> 5;

  const int bid_m = (int)blockIdx.y;
  const int bid_n = (int)blockIdx.x;

  const int off_m = bid_m * BLOCK_M;
  const int off_n = bid_n * BLOCK_N;

#if ENABLE_CLK64_DIAG
  const int diag_base = (bid_m * (N / BLOCK_N) + bid_n) * 8;
#endif

  constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;

  extern __shared__ __align__(1024) char smem_ptr[];
  const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));

  constexpr int A_size = BLOCK_M * BLOCK_K / 2;
  constexpr int B_size = BLOCK_N * BLOCK_K / 2;
  constexpr int SFA_size = 128 * BLOCK_K / 16;
  constexpr int SFB_size = 128 * BLOCK_K / 16;
  constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;

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

  constexpr int OUT2_TMEM = 2 * BLOCK_N;
  constexpr int SFA_tmem = BLOCK_N;
  constexpr int SFB1_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K);
  constexpr int SFB2_tmem = SFB1_tmem + 4 * (BLOCK_K / MMA_K);

  if (warp_id == 0 && elect_sync()) {
    for (int i = 0; i < NUM_STAGES * 2 + 1; i++) mbarrier_init(tma_mbar_addr + i * 8, 1);
    asm volatile("fence.mbarrier_init.release.cluster;");
  } else if (warp_id == 1) {
    asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(BLOCK_N * 4));
  }
  __syncthreads();

  constexpr int num_iters = K / BLOCK_K;

  if (warp_id == NUM_WARPS - 2 && elect_sync()) {
#if ENABLE_CLK64_DIAG
    const unsigned long long t_tma = clock64();
#endif
    constexpr uint64_t cache_A = (CACHE_MODE == 0) ? EVICT_LAST : EVICT_FIRST;
    constexpr uint64_t cache_B = (CACHE_MODE == 0) ? EVICT_FIRST : EVICT_LAST;

    auto issue_tma = [&](int iter_k, int stage_id) {
      const int mbar_addr = tma_mbar_addr + stage_id * 8;

      const int A_smem = smem + stage_id * STAGE_SIZE;
      const int B1_smem = A_smem + A_size;
      const int B2_smem = B1_smem + B_size;
      const int SFA_smem = B2_smem + B_size;
      const int SFB1_smem = SFA_smem + SFA_size;
      const int SFB2_smem = SFB1_smem + SFB_size;

      const int off_k = iter_k * BLOCK_K;
      tma_3d_gmem2smem(A_smem, &A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
      tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
      tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);

      const int rest_k = K / 16 / 4;
      const char *SFA_src = SFA_ptr + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
      const char *SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
      const char *SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
      tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_A);
      tma_gmem2smem(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cache_B);
      tma_gmem2smem(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cache_B);

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

    constexpr int PRELOAD = (num_iters < NUM_STAGES) ? num_iters : NUM_STAGES;
    for (int iter_k = 0; iter_k < PRELOAD; iter_k++) issue_tma(iter_k, iter_k);
    for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
      const int stage_id = iter_k % NUM_STAGES;
      const int mma_phase = (iter_k / NUM_STAGES - 1) & 1;
      mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
      issue_tma(iter_k, stage_id);
    }
#if ENABLE_CLK64_DIAG
    g_clk64_diag[diag_base + 0] = clock64() - t_tma;
#endif
  } else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
    constexpr int MMA_N = BLOCK_N;
    constexpr int MMA_M = 128;
    constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)MMA_N >> 3U << 17U) | ((uint32_t)MMA_M >> 7U << 27U);

#if ENABLE_CLK64_DIAG
    const unsigned long long t_mma = clock64();
#endif
    for (int iter_k = 0; iter_k < num_iters; iter_k++) {
      const int stage_id = iter_k % NUM_STAGES;
      const int tma_phase = (iter_k / NUM_STAGES) & 1;
      mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);

      const int A_smem = smem + stage_id * STAGE_SIZE;
      const int B1_smem = A_smem + A_size;
      const int B2_smem = B1_smem + B_size;
      const int SFA_smem = B2_smem + B_size;
      const int SFB1_smem = SFA_smem + SFA_size;
      const int SFB2_smem = SFB1_smem + SFB_size;

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

      constexpr uint64_t SF_desc = make_desc_SF(0);
      const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);
      const uint64_t SFB1_desc = SF_desc + ((uint64_t)SFB1_smem >> 4ULL);
      const uint64_t SFB2_desc = SF_desc + ((uint64_t)SFB2_smem >> 4ULL);

      for (int k = 0; k < BLOCK_K / MMA_K; k++) {
        uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL);
        uint64_t sfb1_desc = SFB1_desc + (uint64_t)k * (512ULL >> 4ULL);
        uint64_t sfb2_desc = SFB2_desc + (uint64_t)k * (512ULL >> 4ULL);
        tcgen05_cp_nvfp4(SFA_tmem + k * 4, sfa_desc);
        tcgen05_cp_nvfp4(SFB1_tmem + k * 4, sfb1_desc);
        tcgen05_cp_nvfp4(SFB2_tmem + k * 4, sfb2_desc);
      }

      for (int k1 = 0; k1 < BLOCK_K / 256; k1++)
        for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
          uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
          uint64_t b1_desc = make_desc_AB(B1_smem + k1 * BLOCK_N * 128 + k2 * 32);
          uint64_t b2_desc = make_desc_AB(B2_smem + k1 * BLOCK_N * 128 + k2 * 32);

          const int k_sf = k1 * 4 + k2;
          const int scale_A_tmem = SFA_tmem + k_sf * 4;

          int scale_B1_tmem;
          int scale_B2_tmem;
          if constexpr (BLOCK_N == 128) {
            scale_B1_tmem = SFB1_tmem + k_sf * 4;
            scale_B2_tmem = SFB2_tmem + k_sf * 4;
          } else {
            const int off = (bid_n & 1) * (BLOCK_N / 32);
            scale_B1_tmem = SFB1_tmem + k_sf * 4 + off;
            scale_B2_tmem = SFB2_tmem + k_sf * 4 + off;
          }

          const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
          tcgen05_mma_nvfp4(0, a_desc, b1_desc, i_desc, scale_A_tmem, scale_B1_tmem, enable_input_d);
          tcgen05_mma_nvfp4(OUT2_TMEM, a_desc, b2_desc, i_desc, scale_A_tmem, scale_B2_tmem, enable_input_d);
        }

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

    asm volatile(
      "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
      :: "r"(mainloop_mbar_addr)
      : "memory"
    );
#if ENABLE_CLK64_DIAG
    g_clk64_diag[diag_base + 1] = clock64() - t_mma;
#endif
  } else if (tid < BLOCK_M) {
    if constexpr (MBAR_MODE == 0) {
      if (warp_id == 0 && lane_id == 0) mbarrier_wait(mainloop_mbar_addr, 0);
      asm volatile("bar.sync 3, %0;" :: "r"(BLOCK_M) : "memory");
    } else {
      mbarrier_wait(mainloop_mbar_addr, 0);
    }
	    asm volatile("tcgen05.fence::after_thread_sync;");
#if ENABLE_CLK64_DIAG
	    const unsigned long long t_epi = clock64();
#endif
	
	    half2 *out_smem;
	    constexpr int OUT_H2_STRIDE = BLOCK_N / 2 + PAD;

    if constexpr (EPILOGUE_MODE == 1 && BLOCK_N == 64) {
      out_smem = reinterpret_cast<half2 *>(smem_ptr);

      for (int mm = 0; mm < 2; mm++) {
        float g1_reg[BLOCK_N / 2];
        tcgen05_ld_16x256bx8(g1_reg, warp_id * 32 + mm * 16, 0);
        asm volatile("tcgen05.wait::ld.sync.aligned;");

        float tmp[BLOCK_N / 2];
        tcgen05_ld_16x256bx8(tmp, warp_id * 32 + mm * 16, OUT2_TMEM);
        asm volatile("tcgen05.wait::ld.sync.aligned;");

        #pragma unroll
        for (int i = 0; i < BLOCK_N / 8; i++) {
          const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
          const int lc = i * 8 + (lane_id & 3) * 2;
          const float2 x0 = float2{g1_reg[i * 4 + 0], g1_reg[i * 4 + 1]};
          const float2 x8 = float2{g1_reg[i * 4 + 2], g1_reg[i * 4 + 3]};

          float2 o0;
          float2 o8;
          float ex0, ex1, ex2, ex3;
          asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex0) : "f"((-x0.x) * 1.4426950408889634f));
          asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex1) : "f"((-x0.y) * 1.4426950408889634f));
          asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex2) : "f"((-x8.x) * 1.4426950408889634f));
          asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex3) : "f"((-x8.y) * 1.4426950408889634f));

          const float d0 = 1.0f + ex0;
          const float d1 = 1.0f + ex1;
          const float d2 = 1.0f + ex2;
          const float d3 = 1.0f + ex3;

          float s0, s1, s2, s3;
          asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s0) : "f"(d0));
          asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s1) : "f"(d1));
          asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s2) : "f"(d2));
          asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s3) : "f"(d3));

          o0.x = (x0.x * s0) * tmp[i * 4 + 0];
          o0.y = (x0.y * s1) * tmp[i * 4 + 1];
          o8.x = (x8.x * s2) * tmp[i * 4 + 2];
          o8.y = (x8.y * s3) * tmp[i * 4 + 3];

          const int hc = lc >> 1;
          out_smem[(lr + 0) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 0)] = __float22half2_rn(o0);
          out_smem[(lr + 8) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 8)] = __float22half2_rn(o8);
        }
      }
    } else if constexpr (EPILOGUE_MODE == 2 && BLOCK_N == 128) {
      out_smem = reinterpret_cast<half2 *>(smem_ptr);

      #pragma unroll
      for (int mm = 0; mm < 2; mm++) {
        const int tmem_row = warp_id * 32 + mm * 16;

        #pragma unroll
        for (int seg = 0; seg < 2; seg++) {
          constexpr int SEG_N = 64;
          constexpr int SEG_H2 = SEG_N / 2;

          float g1_reg[SEG_H2];
          tcgen05_ld_16x256bx8(g1_reg, tmem_row, seg * SEG_N);
          asm volatile("tcgen05.wait::ld.sync.aligned;");

          float tmp[SEG_H2];
          tcgen05_ld_16x256bx8(tmp, tmem_row, OUT2_TMEM + seg * SEG_N);
          asm volatile("tcgen05.wait::ld.sync.aligned;");

          #pragma unroll
          for (int i = 0; i < SEG_N / 8; i++) {
            const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
            const int lc = seg * SEG_N + i * 8 + (lane_id & 3) * 2;
            const float2 x0 = float2{g1_reg[i * 4 + 0], g1_reg[i * 4 + 1]};
            const float2 x8 = float2{g1_reg[i * 4 + 2], g1_reg[i * 4 + 3]};

            float2 o0;
            float2 o8;
            float s;
            s = fast_sigmoid(x0.x);
            o0.x = (x0.x * s) * tmp[i * 4 + 0];
            s = fast_sigmoid(x0.y);
            o0.y = (x0.y * s) * tmp[i * 4 + 1];
            s = fast_sigmoid(x8.x);
            o8.x = (x8.x * s) * tmp[i * 4 + 2];
            s = fast_sigmoid(x8.y);
            o8.y = (x8.y * s) * tmp[i * 4 + 3];

            const int hc = lc >> 1;
            out_smem[(lr + 0) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 0)] = __float22half2_rn(o0);
            out_smem[(lr + 8) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 8)] = __float22half2_rn(o8);
          }
        }
      }
		    } else if constexpr ((EPILOGUE_MODE == 4 || EPILOGUE_MODE == 6 || EPILOGUE_MODE == 7 || EPILOGUE_MODE == 8) && BLOCK_N == 64) {
		      // 无 out_smem:直写 global(v2/v4),依赖 bar.sync 确保 tcgen05.dealloc 的安全时序。
	      #pragma unroll
	      for (int mm = 0; mm < 2; mm++) {
	        const int tmem_row = warp_id * 32 + mm * 16;
	        constexpr int SEG_N = 64;
	        constexpr int SEG_H2 = SEG_N / 2;

	        float g1_reg[SEG_H2];
	        float tmp[SEG_H2];
	        tcgen05_ld_16x256bx8(g1_reg, tmem_row, 0);
	        tcgen05_ld_16x256bx8(tmp, tmem_row, OUT2_TMEM);
	        asm volatile("tcgen05.wait::ld.sync.aligned;");
	
	        const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
	        const int row0 = off_m + lr;
	        const int row8 = row0 + 8;
	        half2 *out0_h2 = reinterpret_cast<half2 *>(Out_ptr + row0 * N + off_n);
	        half2 *out8_h2 = reinterpret_cast<half2 *>(Out_ptr + row8 * N + off_n);
	
	        #pragma unroll
	        for (int j = 0; j < SEG_N / 16; j++) {
	          const int idx0 = j * 8;
	          const int idx1 = idx0 + 4;
	
	          float2 o0_0;
	          float2 o8_0;
	          if constexpr (EPILOGUE_MODE == 7) {
	            o0_0.x = tmp[idx0 + 0];
	            o0_0.y = tmp[idx0 + 1];
	            o8_0.x = tmp[idx0 + 2];
	            o8_0.y = tmp[idx0 + 3];
	          } else if constexpr (EPILOGUE_MODE == 8) {
	            o0_0.x = g1_reg[idx0 + 0];
	            o0_0.y = g1_reg[idx0 + 1];
	            o8_0.x = g1_reg[idx0 + 2];
	            o8_0.y = g1_reg[idx0 + 3];
	          } else {
	            const float2 x0 = float2{g1_reg[idx0 + 0], g1_reg[idx0 + 1]};
	            const float2 x8 = float2{g1_reg[idx0 + 2], g1_reg[idx0 + 3]};
	
	            float ex0, ex1, ex2, ex3;
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex0) : "f"((-x0.x) * 1.4426950408889634f));
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex1) : "f"((-x0.y) * 1.4426950408889634f));
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex2) : "f"((-x8.x) * 1.4426950408889634f));
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex3) : "f"((-x8.y) * 1.4426950408889634f));
	
	            const float d0 = 1.0f + ex0;
	            const float d1 = 1.0f + ex1;
	            const float d2 = 1.0f + ex2;
	            const float d3 = 1.0f + ex3;
	
	            float s0, s1, s2, s3;
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s0) : "f"(d0));
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s1) : "f"(d1));
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s2) : "f"(d2));
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s3) : "f"(d3));
	
	            const float p0 = x0.x * tmp[idx0 + 0];
	            const float p1 = x0.y * tmp[idx0 + 1];
	            const float p2 = x8.x * tmp[idx0 + 2];
	            const float p3 = x8.y * tmp[idx0 + 3];
	
	            o0_0.x = p0 * s0;
	            o0_0.y = p1 * s1;
	            o8_0.x = p2 * s2;
	            o8_0.y = p3 * s3;
	          }
	
	          float2 o0_1;
	          float2 o8_1;
	          if constexpr (EPILOGUE_MODE == 7) {
	            o0_1.x = tmp[idx1 + 0];
	            o0_1.y = tmp[idx1 + 1];
	            o8_1.x = tmp[idx1 + 2];
	            o8_1.y = tmp[idx1 + 3];
	          } else if constexpr (EPILOGUE_MODE == 8) {
	            o0_1.x = g1_reg[idx1 + 0];
	            o0_1.y = g1_reg[idx1 + 1];
	            o8_1.x = g1_reg[idx1 + 2];
	            o8_1.y = g1_reg[idx1 + 3];
	          } else {
	            const float2 x0 = float2{g1_reg[idx1 + 0], g1_reg[idx1 + 1]};
	            const float2 x8 = float2{g1_reg[idx1 + 2], g1_reg[idx1 + 3]};
	
	            float ex0, ex1, ex2, ex3;
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex0) : "f"((-x0.x) * 1.4426950408889634f));
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex1) : "f"((-x0.y) * 1.4426950408889634f));
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex2) : "f"((-x8.x) * 1.4426950408889634f));
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex3) : "f"((-x8.y) * 1.4426950408889634f));
	
	            const float d0 = 1.0f + ex0;
	            const float d1 = 1.0f + ex1;
	            const float d2 = 1.0f + ex2;
	            const float d3 = 1.0f + ex3;
	
	            float s0, s1, s2, s3;
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s0) : "f"(d0));
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s1) : "f"(d1));
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s2) : "f"(d2));
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s3) : "f"(d3));
	
	            const float p0 = x0.x * tmp[idx1 + 0];
	            const float p1 = x0.y * tmp[idx1 + 1];
	            const float p2 = x8.x * tmp[idx1 + 2];
	            const float p3 = x8.y * tmp[idx1 + 3];
	
	            o0_1.x = p0 * s0;
	            o0_1.y = p1 * s1;
	            o8_1.x = p2 * s2;
	            o8_1.y = p3 * s3;
	          }
	
	          const uint32_t u0_0 = half2_as_u32(__float22half2_rn(o0_0));
	          const uint32_t u1_0 = half2_as_u32(__float22half2_rn(o0_1));
	          const uint32_t u0_8 = half2_as_u32(__float22half2_rn(o8_0));
	          const uint32_t u1_8 = half2_as_u32(__float22half2_rn(o8_1));
	
	          const uint64_t both0 = ((uint64_t)u1_0 << 32) | (uint64_t)u0_0;
	          const uint64_t both8 = ((uint64_t)u1_8 << 32) | (uint64_t)u0_8;
	          const uint64_t peer0 = __shfl_xor_sync(0xFFFFFFFF, both0, 1);
	          const uint64_t peer8 = __shfl_xor_sync(0xFFFFFFFF, both8, 1);
	
	          const uint32_t peer_u0_0 = (uint32_t)peer0;
	          const uint32_t peer_u1_0 = (uint32_t)(peer0 >> 32);
	          const uint32_t peer_u0_8 = (uint32_t)peer8;
	          const uint32_t peer_u1_8 = (uint32_t)(peer8 >> 32);
	
		          const int t = lane_id & 3;
		          const int t_odd = t & 1;
		          const uint32_t lo0 = t_odd ? peer_u1_0 : u0_0;
		          const uint32_t hi0 = t_odd ? u1_0 : peer_u0_0;
		          const uint32_t lo8 = t_odd ? peer_u1_8 : u0_8;
		          const uint32_t hi8 = t_odd ? u1_8 : peer_u0_8;
		
			          if constexpr ((EPILOGUE_MODE == 4 || EPILOGUE_MODE == 7 || EPILOGUE_MODE == 8) || (EPILOGUE_MODE == 6 && !ENABLE_V4_EPILOGUE)) {
		            const int store_off = (t & 2) + ((t & 1) << 2);
		            const half2 *out0_p = out0_h2 + j * 8 + store_off;
		            const half2 *out8_p = out8_h2 + j * 8 + store_off;
		            stg_v2_b32(out0_p, lo0, hi0);
		            stg_v2_b32(out8_p, lo8, hi8);
		          } else {
		            const int base_lane = lane_id & ~3;
		            if (t == 0) {
		              const uint32_t lo0_p = __shfl_sync(0xFFFFFFFF, lo0, base_lane + 2);
		              const uint32_t hi0_p = __shfl_sync(0xFFFFFFFF, hi0, base_lane + 2);
		              const uint32_t lo8_p = __shfl_sync(0xFFFFFFFF, lo8, base_lane + 2);
		              const uint32_t hi8_p = __shfl_sync(0xFFFFFFFF, hi8, base_lane + 2);
		              const half2 *out0_p = out0_h2 + j * 8;
		              const half2 *out8_p = out8_h2 + j * 8;
		              stg_v4_b32(out0_p, lo0, hi0, lo0_p, hi0_p);
		              stg_v4_b32(out8_p, lo8, hi8, lo8_p, hi8_p);
		            } else if (t == 1) {
		              const uint32_t lo0_p = __shfl_sync(0xFFFFFFFF, lo0, base_lane + 3);
		              const uint32_t hi0_p = __shfl_sync(0xFFFFFFFF, hi0, base_lane + 3);
		              const uint32_t lo8_p = __shfl_sync(0xFFFFFFFF, lo8, base_lane + 3);
		              const uint32_t hi8_p = __shfl_sync(0xFFFFFFFF, hi8, base_lane + 3);
		              const half2 *out0_p = out0_h2 + j * 8 + 4;
		              const half2 *out8_p = out8_h2 + j * 8 + 4;
		              stg_v4_b32(out0_p, lo0, hi0, lo0_p, hi0_p);
		              stg_v4_b32(out8_p, lo8, hi8, lo8_p, hi8_p);
		            }
		          }
		        }
		      }
	
	      asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
	      if (warp_id == 0) asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 4));
#if ENABLE_CLK64_DIAG
	      if (lane_id == 0) g_clk64_diag[diag_base + 2 + warp_id] = clock64() - t_epi;
#endif
	      return;
	    } else if constexpr ((EPILOGUE_MODE == 3 || EPILOGUE_MODE == 5 || EPILOGUE_MODE == 7 || EPILOGUE_MODE == 8) && BLOCK_N == 128) {
	      // 无 out_smem:直接写回 global,依赖 bar.sync 确保 tcgen05.dealloc 的安全时序。
	      #pragma unroll
	      for (int mm = 0; mm < 2; mm++) {
	        const int tmem_row = warp_id * 32 + mm * 16;

        #pragma unroll
        for (int seg = 0; seg < 2; seg++) {
          constexpr int SEG_N = 64;
          constexpr int SEG_H2 = SEG_N / 2;

          float g1_reg[SEG_H2];
          float tmp[SEG_H2];
          tcgen05_ld_16x256bx8(g1_reg, tmem_row, seg * SEG_N);
          tcgen05_ld_16x256bx8(tmp, tmem_row, OUT2_TMEM + seg * SEG_N);
          asm volatile("tcgen05.wait::ld.sync.aligned;");
	
	          const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
	          const int row0 = off_m + lr;
	          const int row8 = row0 + 8;
	          half2 *out0_h2 = reinterpret_cast<half2 *>(Out_ptr + row0 * N + off_n) + seg * 32;
	          half2 *out8_h2 = reinterpret_cast<half2 *>(Out_ptr + row8 * N + off_n) + seg * 32;
	
	          #pragma unroll
	          for (int j = 0; j < SEG_N / 16; j++) {
	            const int idx0 = j * 8;
	            const int idx1 = idx0 + 4;
	
	            const float2 x0_0 = float2{g1_reg[idx0 + 0], g1_reg[idx0 + 1]};
	            const float2 x8_0 = float2{g1_reg[idx0 + 2], g1_reg[idx0 + 3]};
	            const float2 x0_1 = float2{g1_reg[idx1 + 0], g1_reg[idx1 + 1]};
	            const float2 x8_1 = float2{g1_reg[idx1 + 2], g1_reg[idx1 + 3]};
	
	            float ex0_0, ex0_1, ex1_0, ex1_1, ex2_0, ex2_1, ex3_0, ex3_1;
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex0_0) : "f"((-x0_0.x) * 1.4426950408889634f));
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex0_1) : "f"((-x0_1.x) * 1.4426950408889634f));
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex1_0) : "f"((-x0_0.y) * 1.4426950408889634f));
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex1_1) : "f"((-x0_1.y) * 1.4426950408889634f));
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex2_0) : "f"((-x8_0.x) * 1.4426950408889634f));
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex2_1) : "f"((-x8_1.x) * 1.4426950408889634f));
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex3_0) : "f"((-x8_0.y) * 1.4426950408889634f));
	            asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex3_1) : "f"((-x8_1.y) * 1.4426950408889634f));
	
	            const float d0_0 = 1.0f + ex0_0;
	            const float d0_1 = 1.0f + ex0_1;
	            const float d1_0 = 1.0f + ex1_0;
	            const float d1_1 = 1.0f + ex1_1;
	            const float d2_0 = 1.0f + ex2_0;
	            const float d2_1 = 1.0f + ex2_1;
	            const float d3_0 = 1.0f + ex3_0;
	            const float d3_1 = 1.0f + ex3_1;
	
	            float s0_0, s0_1, s1_0, s1_1, s2_0, s2_1, s3_0, s3_1;
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s0_0) : "f"(d0_0));
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s0_1) : "f"(d0_1));
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s1_0) : "f"(d1_0));
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s1_1) : "f"(d1_1));
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s2_0) : "f"(d2_0));
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s2_1) : "f"(d2_1));
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s3_0) : "f"(d3_0));
	            asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s3_1) : "f"(d3_1));
	
	            float2 o0_0;
	            float2 o8_0;
	            float2 o0_1;
	            float2 o8_1;
	
	            if constexpr (EPILOGUE_MODE == 7) {
	              o0_0.x = tmp[idx0 + 0];
	              o0_0.y = tmp[idx0 + 1];
	              o8_0.x = tmp[idx0 + 2];
	              o8_0.y = tmp[idx0 + 3];
	              o0_1.x = tmp[idx1 + 0];
	              o0_1.y = tmp[idx1 + 1];
	              o8_1.x = tmp[idx1 + 2];
	              o8_1.y = tmp[idx1 + 3];
	            } else if constexpr (EPILOGUE_MODE == 8) {
	              o0_0.x = x0_0.x;
	              o0_0.y = x0_0.y;
	              o8_0.x = x8_0.x;
	              o8_0.y = x8_0.y;
	              o0_1.x = x0_1.x;
	              o0_1.y = x0_1.y;
	              o8_1.x = x8_1.x;
	              o8_1.y = x8_1.y;
	            } else {
	              o0_0.x = (x0_0.x * tmp[idx0 + 0]) * s0_0;
	              o0_0.y = (x0_0.y * tmp[idx0 + 1]) * s1_0;
	              o8_0.x = (x8_0.x * tmp[idx0 + 2]) * s2_0;
	              o8_0.y = (x8_0.y * tmp[idx0 + 3]) * s3_0;
	              o0_1.x = (x0_1.x * tmp[idx1 + 0]) * s0_1;
	              o0_1.y = (x0_1.y * tmp[idx1 + 1]) * s1_1;
	              o8_1.x = (x8_1.x * tmp[idx1 + 2]) * s2_1;
	              o8_1.y = (x8_1.y * tmp[idx1 + 3]) * s3_1;
	            }
	
	            const uint32_t u0_0 = half2_as_u32(__float22half2_rn(o0_0));
	            const uint32_t u1_0 = half2_as_u32(__float22half2_rn(o0_1));
	            const uint32_t u0_8 = half2_as_u32(__float22half2_rn(o8_0));
	            const uint32_t u1_8 = half2_as_u32(__float22half2_rn(o8_1));
	
	            const uint64_t both0 = ((uint64_t)u1_0 << 32) | (uint64_t)u0_0;
	            const uint64_t both8 = ((uint64_t)u1_8 << 32) | (uint64_t)u0_8;
	            const uint64_t peer0 = __shfl_xor_sync(0xFFFFFFFF, both0, 1);
	            const uint64_t peer8 = __shfl_xor_sync(0xFFFFFFFF, both8, 1);
	
	            const uint32_t peer_u0_0 = (uint32_t)peer0;
	            const uint32_t peer_u1_0 = (uint32_t)(peer0 >> 32);
	            const uint32_t peer_u0_8 = (uint32_t)peer8;
	            const uint32_t peer_u1_8 = (uint32_t)(peer8 >> 32);
	
	            const int t = lane_id & 3;
	            const int t_odd = t & 1;
	            const uint32_t lo0 = t_odd ? peer_u1_0 : u0_0;
	            const uint32_t hi0 = t_odd ? u1_0 : peer_u0_0;
	            const uint32_t lo8 = t_odd ? peer_u1_8 : u0_8;
	            const uint32_t hi8 = t_odd ? u1_8 : peer_u0_8;
	
		            if constexpr ((EPILOGUE_MODE == 3 || EPILOGUE_MODE == 7 || EPILOGUE_MODE == 8) || (EPILOGUE_MODE == 5 && !ENABLE_V4_EPILOGUE)) {
	              const int store_off = (t & 2) + ((t & 1) << 2);
	              const half2 *out0_p = out0_h2 + j * 8 + store_off;
	              const half2 *out8_p = out8_h2 + j * 8 + store_off;
	              stg_v2_b32(out0_p, lo0, hi0);
	              stg_v2_b32(out8_p, lo8, hi8);
	            } else {
	              const int base_lane = lane_id & ~3;
	              if (t == 0) {
	                const uint32_t lo0_p = __shfl_sync(0xFFFFFFFF, lo0, base_lane + 2);
	                const uint32_t hi0_p = __shfl_sync(0xFFFFFFFF, hi0, base_lane + 2);
	                const uint32_t lo8_p = __shfl_sync(0xFFFFFFFF, lo8, base_lane + 2);
	                const uint32_t hi8_p = __shfl_sync(0xFFFFFFFF, hi8, base_lane + 2);
	                const half2 *out0_p = out0_h2 + j * 8;
	                const half2 *out8_p = out8_h2 + j * 8;
	                stg_v4_b32(out0_p, lo0, hi0, lo0_p, hi0_p);
	                stg_v4_b32(out8_p, lo8, hi8, lo8_p, hi8_p);
	              } else if (t == 1) {
	                const uint32_t lo0_p = __shfl_sync(0xFFFFFFFF, lo0, base_lane + 3);
	                const uint32_t hi0_p = __shfl_sync(0xFFFFFFFF, hi0, base_lane + 3);
	                const uint32_t lo8_p = __shfl_sync(0xFFFFFFFF, lo8, base_lane + 3);
	                const uint32_t hi8_p = __shfl_sync(0xFFFFFFFF, hi8, base_lane + 3);
	                const half2 *out0_p = out0_h2 + j * 8 + 4;
	                const half2 *out8_p = out8_h2 + j * 8 + 4;
	                stg_v4_b32(out0_p, lo0, hi0, lo0_p, hi0_p);
	                stg_v4_b32(out8_p, lo8, hi8, lo8_p, hi8_p);
	              }
	            }
	          }
	        }
	      }
	
	      asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
	      if (warp_id == 0) asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 4));
#if ENABLE_CLK64_DIAG
	      if (lane_id == 0) g_clk64_diag[diag_base + 2 + warp_id] = clock64() - t_epi;
#endif
	      return;
	    } else {
      constexpr int G1_STRIDE = BLOCK_N + 8;
      float *g1_smem = reinterpret_cast<float *>(smem_ptr);
      out_smem = reinterpret_cast<half2 *>(g1_smem + BLOCK_M * G1_STRIDE);

      for (int mm = 0; mm < 2; mm++) {
        float tmp[BLOCK_N / 2];
        if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + mm * 16, 0);
        else tcgen05_ld_16x256bx16(tmp, warp_id * 32 + mm * 16, 0);
        asm volatile("tcgen05.wait::ld.sync.aligned;");

        #pragma unroll
        for (int i = 0; i < BLOCK_N / 8; i++) {
          const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
          const int lc = i * 8 + (lane_id & 3) * 2;
          reinterpret_cast<float2 *>(g1_smem + (lr + 0) * G1_STRIDE + lc)[0] = float2{tmp[i * 4 + 0], tmp[i * 4 + 1]};
          reinterpret_cast<float2 *>(g1_smem + (lr + 8) * G1_STRIDE + lc)[0] = float2{tmp[i * 4 + 2], tmp[i * 4 + 3]};
        }
      }

      for (int mm = 0; mm < 2; mm++) {
        float tmp[BLOCK_N / 2];
        if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + mm * 16, OUT2_TMEM);
        else tcgen05_ld_16x256bx16(tmp, warp_id * 32 + mm * 16, OUT2_TMEM);
        asm volatile("tcgen05.wait::ld.sync.aligned;");

        #pragma unroll
        for (int i = 0; i < BLOCK_N / 8; i++) {
          const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
          const int lc = i * 8 + (lane_id & 3) * 2;
          const float2 x0 = reinterpret_cast<const float2 *>(g1_smem + (lr + 0) * G1_STRIDE + lc)[0];
          const float2 x8 = reinterpret_cast<const float2 *>(g1_smem + (lr + 8) * G1_STRIDE + lc)[0];

          float2 o0;
          float2 o8;
          float s;
          s = fast_sigmoid(x0.x);
          o0.x = (x0.x * s) * tmp[i * 4 + 0];
          s = fast_sigmoid(x0.y);
          o0.y = (x0.y * s) * tmp[i * 4 + 1];
          s = fast_sigmoid(x8.x);
          o8.x = (x8.x * s) * tmp[i * 4 + 2];
          s = fast_sigmoid(x8.y);
          o8.y = (x8.y * s) * tmp[i * 4 + 3];

          const int hc = lc >> 1;
          out_smem[(lr + 0) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 0)] = __float22half2_rn(o0);
          out_smem[(lr + 8) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 8)] = __float22half2_rn(o8);
        }
      }
    }

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

    constexpr int WRITER_WARPS = BLOCK_M / WARP_SIZE;
    constexpr int ROWS_PER_WARP = BLOCK_M / WRITER_WARPS;
    const int r_base = warp_id * ROWS_PER_WARP;
    for (int rr = 0; rr < ROWS_PER_WARP; rr++) {
      const int r = r_base + rr;
      const int row = off_m + r;
      half2 *dst = reinterpret_cast<half2 *>(Out_ptr + row * N + off_n);
      const half2 *src = out_smem + r * OUT_H2_STRIDE;
      if constexpr (BLOCK_N == 64) {
        dst[lane_id] = src[out_smem_col<SWZ>(lane_id, r)];
      } else {
        dst[lane_id] = src[out_smem_col<SWZ>(lane_id, r)];
        dst[lane_id + 32] = src[out_smem_col<SWZ>(lane_id + 32, r)];
      }
    }
  }
}

template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
static inline void launch_gemm_f32(
  const at::Tensor& A,
  const at::Tensor& B,
  const at::Tensor& SFA,
  const at::Tensor& SFB,
  at::Tensor& C
) {
  const int M = (int)A.size(0);
  const int N = (int)B.size(0);
  const int grid_m = M / BLOCK_M;
  const int grid_n = N / BLOCK_N;
  const bool prefer_cache_A = grid_n >= grid_m;

  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());
  auto C_ptr = reinterpret_cast<float *>(C.data_ptr());

  CUtensorMap A_tmap, B_tmap;
  init_AB_tmap(&A_tmap, A_ptr, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K, CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE);
  init_AB_tmap(&B_tmap, B_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K, CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE);

  dim3 grid(1, (unsigned)(grid_m * grid_n));
  const int tb_size = BLOCK_M + 2 * WARP_SIZE;
  constexpr int AB_size = (BLOCK_M + BLOCK_N) * (BLOCK_K / 2);
  constexpr int SFAB_size = 128 * (BLOCK_K / 16) * 2;
  constexpr int smem_size = (AB_size + SFAB_size) * NUM_STAGES;

  auto kptr0 = gemm_f32_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, 0>;
  auto kptr1 = gemm_f32_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, 1>;
  auto kptr = prefer_cache_A ? kptr1 : kptr0;
  if constexpr (smem_size > 48'000) {
    static int attr0 = 0;
    static int attr1 = 0;
    if (prefer_cache_A) {
      if (__builtin_expect(attr1 == 0, 0)) {
        cudaFuncSetAttribute(kptr1, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
        attr1 = 1;
      }
    } else {
      if (__builtin_expect(attr0 == 0, 0)) {
        cudaFuncSetAttribute(kptr0, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
        attr0 = 1;
      }
    }
  }
  kptr<<<grid, tb_size, smem_size>>>(A_tmap, B_tmap, SFA_ptr, SFB_ptr, C_ptr, M, N);
}

template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int PAD, int SWZ>
static inline void launch_gemm_silu_mul(
  const at::Tensor& A,
  const at::Tensor& B,
  const at::Tensor& SFA,
  const at::Tensor& SFB,
  const at::Tensor& g1,
  at::Tensor& out
) {
  const int M = (int)A.size(0);
  const int N = (int)B.size(0);
  const int grid_m = M / BLOCK_M;
  const int grid_n = N / BLOCK_N;
  const bool prefer_cache_A = grid_n >= grid_m;

  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());
  auto G1_ptr = reinterpret_cast<const float *>(g1.data_ptr());
  auto Out_ptr = reinterpret_cast<half *>(out.data_ptr());

  CUtensorMap A_tmap, B_tmap;
  init_AB_tmap(&A_tmap, A_ptr, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K, CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE);
  init_AB_tmap(&B_tmap, B_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K, CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE);

  dim3 grid(1, (unsigned)(grid_m * grid_n));
  const int tb_size = BLOCK_M + 2 * WARP_SIZE;
  constexpr int AB_size = (BLOCK_M + BLOCK_N) * (BLOCK_K / 2);
  constexpr int SFAB_size = 128 * (BLOCK_K / 16) * 2;
  constexpr int smem_stage = (AB_size + SFAB_size) * NUM_STAGES;
  constexpr int OUT_H2_STRIDE = BLOCK_N / 2 + PAD;
  constexpr int smem_scratch = BLOCK_M * OUT_H2_STRIDE * (int)sizeof(half2);
  constexpr int smem_size = (smem_stage > smem_scratch) ? smem_stage : smem_scratch;

  auto kptr0 = gemm_silu_mul_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, 0, PAD, SWZ>;
  auto kptr1 = gemm_silu_mul_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, 1, PAD, SWZ>;
  auto kptr = prefer_cache_A ? kptr0 : kptr1;
  if constexpr (smem_size > 48'000) {
    static int attr0 = 0;
    static int attr1 = 0;
    if (prefer_cache_A) {
      if (__builtin_expect(attr0 == 0, 0)) {
        cudaFuncSetAttribute(kptr0, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
        attr0 = 1;
      }
    } else {
      if (__builtin_expect(attr1 == 0, 0)) {
        cudaFuncSetAttribute(kptr1, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
        attr1 = 1;
      }
    }
  }
  kptr<<<grid, tb_size, smem_size>>>(A_tmap, B_tmap, SFA_ptr, SFB_ptr, G1_ptr, Out_ptr, M, N);
}

template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CACHE_MODE, int PAD, int SWZ, int EPILOGUE_MODE, int MBAR_MODE>
static inline void launch_dual_gemm_silu(
  const at::Tensor& A,
  const at::Tensor& B1,
  const at::Tensor& B2,
  const at::Tensor& SFA,
  const at::Tensor& SFB1,
  const at::Tensor& SFB2,
  at::Tensor& out
) {
  const int M = (int)A.size(0);
  const int N = (int)B1.size(0);

  auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());
  auto B1_ptr = reinterpret_cast<const char *>(B1.data_ptr());
  auto B2_ptr = reinterpret_cast<const char *>(B2.data_ptr());
  auto SFA_ptr = reinterpret_cast<const char *>(SFA.data_ptr());
  auto SFB1_ptr = reinterpret_cast<const char *>(SFB1.data_ptr());
  auto SFB2_ptr = reinterpret_cast<const char *>(SFB2.data_ptr());
  auto Out_ptr = reinterpret_cast<half *>(out.data_ptr());

	  CUtensorMap A_tmap, B1_tmap, B2_tmap;
	  const CUtensorMapL2promotion A_l2 =
	    (CACHE_MODE == 0) ? CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_L2_256B
	                      : CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE;
	  const CUtensorMapL2promotion B_l2 =
	    (CACHE_MODE == 0) ? CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE
	                      : CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_L2_256B;
	  init_AB_tmap(&A_tmap, A_ptr, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K, A_l2);
	  init_AB_tmap(&B1_tmap, B1_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K, B_l2);
	  init_AB_tmap(&B2_tmap, B2_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K, B_l2);

  const int grid_m = M / BLOCK_M;
  const int grid_n = N / BLOCK_N;
  dim3 grid((unsigned)grid_n, (unsigned)grid_m);
  const int tb_size = BLOCK_M + 2 * WARP_SIZE;
  constexpr int AB_size = (BLOCK_M + 2 * BLOCK_N) * (BLOCK_K / 2);
  constexpr int SF_size = 128 * (BLOCK_K / 16) * 3;
  constexpr int smem_stage = (AB_size + SF_size) * NUM_STAGES;
  constexpr int G1_STRIDE = BLOCK_N + 8;
  constexpr int OUT_H2_STRIDE = BLOCK_N / 2 + PAD;
  constexpr bool REG_FUSED = (EPILOGUE_MODE == 1 && BLOCK_N == 64) || (EPILOGUE_MODE == 2 && BLOCK_N == 128);
  constexpr bool DIRECT_STORE =
    ((EPILOGUE_MODE == 3 || EPILOGUE_MODE == 5 || EPILOGUE_MODE == 7 || EPILOGUE_MODE == 8) && BLOCK_N == 128) ||
    ((EPILOGUE_MODE == 4 || EPILOGUE_MODE == 6 || EPILOGUE_MODE == 7 || EPILOGUE_MODE == 8) && BLOCK_N == 64);
  constexpr int smem_scratch = DIRECT_STORE ? 0 : (REG_FUSED
    ? (BLOCK_M * OUT_H2_STRIDE * (int)sizeof(half2))
    : (BLOCK_M * G1_STRIDE * (int)sizeof(float) + BLOCK_M * OUT_H2_STRIDE * (int)sizeof(half2)));
  constexpr int smem_size = (smem_stage > smem_scratch) ? smem_stage : smem_scratch;

  auto kptr = dual_gemm_silu_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, CACHE_MODE, PAD, SWZ, EPILOGUE_MODE, MBAR_MODE>;
  if constexpr (smem_size > 48'000) {
    static int attr = 0;
    if (__builtin_expect(attr == 0, 0)) {
      cudaFuncSetAttribute(kptr, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
      attr = 1;
    }
  }
  kptr<<<grid, tb_size, smem_size>>>(A_tmap, B1_tmap, B2_tmap, SFA_ptr, SFB1_ptr, SFB2_ptr, Out_ptr, M, N);
}

__global__ void silu_mul_f32_vec2(const float* __restrict__ x, const float* __restrict__ y, half* __restrict__ out, int64_t n2) {
  const int64_t idx = int64_t(blockIdx.x) * blockDim.x + threadIdx.x;
  if (idx >= n2) return;
  const float2 fx = reinterpret_cast<const float2*>(x)[idx];
  const float2 fy = reinterpret_cast<const float2*>(y)[idx];
  float2 o;
  const float sx0 = fast_sigmoid(fx.x);
  const float sx1 = fast_sigmoid(fx.y);
  o.x = (fx.x * sx0) * fy.x;
  o.y = (fx.y * sx1) * fy.y;
  reinterpret_cast<half2*>(out)[idx] = __float22half2_rn(o);
}

static inline void launch_silu_mul_f32(const at::Tensor& g1, const at::Tensor& g2, at::Tensor& out) {
  const int64_t n = out.numel();
  TORCH_CHECK((n & 1) == 0, "n");
  const int64_t n2 = n >> 1;
  const int threads = 256;
  const int blocks = (int)((n2 + threads - 1) / threads);
  silu_mul_f32_vec2<<<blocks, threads>>>(
    reinterpret_cast<const float*>(g1.data_ptr()),
    reinterpret_cast<const float*>(g2.data_ptr()),
    reinterpret_cast<half*>(out.data_ptr()),
    n2
  );
}

at::Tensor fused(
  const at::Tensor& A,
  const at::Tensor& B1,
  const at::Tensor& B2,
  const at::Tensor& SFA,
  const at::Tensor& SFB1,
  const at::Tensor& SFB2,
  at::Tensor& out,
  at::Tensor& g1,
  at::Tensor& g2
) {
  TORCH_CHECK(A.is_cuda() && B1.is_cuda() && B2.is_cuda(), "cuda");
  TORCH_CHECK(SFA.is_cuda() && SFB1.is_cuda() && SFB2.is_cuda(), "cuda");
  TORCH_CHECK(out.is_cuda() && g1.is_cuda() && g2.is_cuda(), "cuda");
  TORCH_CHECK(A.dim() == 3 && B1.dim() == 3 && B2.dim() == 3, "dim");
  TORCH_CHECK(out.dim() == 3 && g1.dim() == 3 && g2.dim() == 3, "dim");

  const int64_t M = A.size(0);
  const int64_t Kp = A.size(1);
  const int64_t L = A.size(2);
  const int64_t N = B1.size(0);
  TORCH_CHECK(L == 1, "l");
  TORCH_CHECK(B1.size(1) == Kp && B1.size(2) == L, "b1");
  TORCH_CHECK(B2.size(1) == Kp && B2.size(2) == L, "b2");
  TORCH_CHECK(out.size(0) == M && out.size(1) == N && out.size(2) == L, "out");
  TORCH_CHECK(g1.size(0) == M && g1.size(1) == N && g1.size(2) == L, "g1");
  TORCH_CHECK(g2.size(0) == M && g2.size(1) == N && g2.size(2) == L, "g2");

  TORCH_CHECK((M % 128) == 0, "m");
  TORCH_CHECK((N % 64) == 0, "n");

  const int K = (int)(Kp * 2);
		  if (K == 7168) {
		    if (M == 512 && N == 4096) {
		      constexpr int CACHE_MODE = 0;
		      constexpr int PAD = 2;
		      constexpr int SWZ = 1;
		      constexpr int MBAR_MODE = 0;
			      launch_dual_gemm_silu<7168, 128, 128, 256, 4, CACHE_MODE, PAD, SWZ, 5, MBAR_MODE>(A, B1, B2, SFA, SFB1, SFB2, out);
		    } else if (M == 512 && N == 3072) {
		      constexpr int CACHE_MODE = 0;
		      constexpr int PAD = 2;
		      constexpr int SWZ = 1;
		      constexpr int MBAR_MODE = 0;
			      launch_dual_gemm_silu<7168, 128, 128, 256, 4, CACHE_MODE, PAD, SWZ, 5, MBAR_MODE>(A, B1, B2, SFA, SFB1, SFB2, out);
		    } else if (M == 256 && N == 4096) {
		      constexpr int CACHE_MODE = 1;
		      constexpr int PAD = 2;
		      constexpr int SWZ = 1;
		      constexpr int MBAR_MODE = 0;
			      launch_dual_gemm_silu<7168, 128,  64, 256, 5, CACHE_MODE, PAD, SWZ, 6, MBAR_MODE>(A, B1, B2, SFA, SFB1, SFB2, out);
		    } else {
      launch_gemm_f32<7168, 128, 64, 256, 8>(A, B1, SFA, SFB1, g1);
      launch_gemm_f32<7168, 128, 64, 256, 8>(A, B2, SFA, SFB2, g2);
      launch_silu_mul_f32(g1, g2, out);
    }
  } else if (K == 4096) {
    if (M == 256 && N == 3072) {
      constexpr int CACHE_MODE = 0;
      constexpr int PAD = 1;
      constexpr int SWZ = 1;
      constexpr int MBAR_MODE = 0;
      launch_dual_gemm_silu<4096, 128,  64, 256, 5, CACHE_MODE, PAD, SWZ, 6, MBAR_MODE>(A, B1, B2, SFA, SFB1, SFB2, out);
    } else {
      launch_gemm_f32<4096, 128, 64, 256, 8>(A, B1, SFA, SFB1, g1);
      launch_gemm_f32<4096, 128, 64, 256, 8>(A, B2, SFA, SFB2, g2);
      launch_silu_mul_f32(g1, g2, out);
    }
  } else if (K == 2304) {
    launch_gemm_f32<2304, 128, 64, 256, 8>(A, B1, SFA, SFB1, g1);
    launch_gemm_f32<2304, 128, 64, 256, 8>(A, B2, SFA, SFB2, g2);
    launch_silu_mul_f32(g1, g2, out);
  } else if (K == 2048) {
    launch_gemm_f32<2048, 128, 64, 256, 8>(A, B1, SFA, SFB1, g1);
    launch_gemm_f32<2048, 128, 64, 256, 8>(A, B2, SFA, SFB2, g2);
    launch_silu_mul_f32(g1, g2, out);
  } else if (K == 1536) {
    launch_gemm_f32<1536, 128, 64, 256, 8>(A, B1, SFA, SFB1, g1);
    launch_gemm_f32<1536, 128, 64, 256, 8>(A, B2, SFA, SFB2, g2);
    launch_silu_mul_f32(g1, g2, out);
  } else if (K == 512) {
    launch_gemm_f32<512, 128, 64, 256, 8>(A, B1, SFA, SFB1, g1);
    launch_gemm_f32<512, 128, 64, 256, 8>(A, B2, SFA, SFB2, g2);
    launch_silu_mul_f32(g1, g2, out);
  } else if (K == 256) {
    launch_gemm_f32<256, 128, 64, 256, 8>(A, B1, SFA, SFB1, g1);
    launch_gemm_f32<256, 128, 64, 256, 8>(A, B2, SFA, SFB2, g2);
    launch_silu_mul_f32(g1, g2, out);
  } else {
    TORCH_CHECK(false, "k ", K);
  }

  return out;
}

TORCH_LIBRARY(nvfp4_dual_lib, m) {
  m.def("fused(Tensor A, Tensor B1, Tensor B2, Tensor SFA, Tensor SFB1, Tensor SFB2, Tensor(a!) out, Tensor(b!) g1, Tensor(c!) g2) -> Tensor");
  m.impl("fused", &fused);
}
"""


_loaded = False


def _load():
    global _loaded
    if _loaded:
        return
    load_inline(
        name="nvfp4_dual_ext_tc_v3",
        cpp_sources="",
        cuda_sources=_CUDA_SRC,
        functions=None,
        with_cuda=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"],
        verbose=False,
        is_python_module=False,
        no_implicit_headers=True,
    )
    _loaded = True
    torch.cuda.empty_cache()
    torch.cuda.ipc_collect()


_buf_cache = {}
_ranked_cache_cleared = False


def _is_ranked(m: int, n: int, k: int, l: int) -> bool:
    if l != 1:
        return False
    if k == 7168:
        return (m == 512 and (n == 4096 or n == 3072)) or (m == 256 and n == 4096)
    if k == 4096:
        return (m == 256) and (n == 3072)
    return False


def _get_buf(tag, shape, device):
    m, n, l = int(shape[0]), int(shape[1]), int(shape[2])
    need = m * n * l
    key = (device, int(tag))
    t = _buf_cache.get(key)
    if t is None or (t.device != device) or (t.numel() < need):
        t = torch.empty((need,), device=device, dtype=torch.float32)
        _buf_cache[key] = t
    return t[:need].view(m, n, l)


def custom_kernel(data):
    _load()
    a, b1, b2, _sfa, _sfb1, _sfb2, sfa_p, sfb1_p, sfb2_p, c = data
    m = int(a.shape[0])
    n = int(b1.shape[0])
    k = int(a.shape[1]) * 2
    l = int(a.shape[2])
    if _is_ranked(m, n, k, l):
        global _ranked_cache_cleared
        if not _ranked_cache_cleared:
            torch.cuda.empty_cache()
            torch.cuda.ipc_collect()
            _ranked_cache_cleared = True
        g1 = _get_buf(1, c.shape, a.device)
        g2 = _get_buf(2, c.shape, a.device)
        return torch.ops.nvfp4_dual_lib.fused(a, b1, b2, sfa_p, sfb1_p, sfb2_p, c, g1, g2)

    torch.cuda.empty_cache()
    g1 = torch.empty(c.shape, device=a.device, dtype=torch.float32)
    g2 = torch.empty(c.shape, device=a.device, dtype=torch.float32)
    out = torch.ops.nvfp4_dual_lib.fused(a, b1, b2, sfa_p, sfb1_p, sfb2_p, c, g1, g2)
    del g1, g2
    torch.cuda.empty_cache()
    torch.cuda.ipc_collect()
    return out


__all__ = ["custom_kernel"]
scrolls · 1782 lines total

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

Changes from previous submission

Against this author's previous submission submission 374289.

+ import os
+
+ os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
+ os.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":16:8")
+ os.environ.setdefault("CUDA_MODULE_LOADING", "LAZY")
+
import torch
from torch.utils.cpp_extension import load_inline
-
-
-
-
-
-
_CUDA_SRC = r"""
#include <cuda.h>
#include <cudaTypedefs.h>
⋯ 9 unchanged lines
constexpr uint64_t EVICT_FIRST = 0x12F0000000000000ULL;
constexpr uint64_t EVICT_LAST = 0x14F0000000000000ULL;
+ #ifndef ENABLE_CLK64_DIAG
+ #define ENABLE_CLK64_DIAG 0
+ #endif
+
+ #ifndef ENABLE_V4_EPILOGUE
+ #define ENABLE_V4_EPILOGUE 0
+ #endif
+
+ #if ENABLE_CLK64_DIAG
+ __device__ __align__(16) unsigned long long g_clk64_diag[65536];
+ #endif
+
__device__ __forceinline__ constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; }
+ template <int SWZ>
+ __device__ __forceinline__ int out_smem_col(int col_h2, int row) {
+ if constexpr (SWZ) return col_h2 ^ ((row & 7) << 2);
+ return col_h2;
+ }
+
+ // 精度策略:使用 `ex2.approx` + `rcp.approx` 近似 sigmoid;目标是使最终 fp16 输出满足 rtol/atol=1e-3 门槛。
+ __device__ __forceinline__ float fast_sigmoid(float x) {
+ float ex2;
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex2) : "f"((-x) * 1.4426950408889634f));
+ float denom = 1.0f + ex2;
+ float rcp;
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(rcp) : "f"(denom));
+ return rcp;
+ }
+
__device__ __forceinline__ uint32_t elect_sync() {
uint32_t pred = 0;
asm volatile(
⋯ 17 unchanged lines
asm volatile(
"{\n\t"
".reg .pred P1;\n\t"
+ ".reg .u32 B;\n\t"
+ "mov.u32 B, 1;\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"
+ "nanosleep.u32 B;\n\t"
+ "shl.b32 B, B, 1;\n\t"
+ "min.u32 B, B, 32;\n\t"
"bra.uni LAB_WAIT;\n\t"
"DONE:\n\t"
"}\n\t"
⋯ 42 unchanged lines
);
}
+ __device__ __forceinline__ uint32_t half2_as_u32(half2 x) {
+ union {
+ half2 h;
+ uint32_t u;
+ } v;
+ v.h = x;
+ return v.u;
+ }
+
+ __device__ __forceinline__ void stg_v2_b32(const void *p, uint32_t a, uint32_t b) {
+ asm volatile("st.global.v2.b32 [%0], {%1, %2};" :: "l"(p), "r"(a), "r"(b) : "memory");
+ }
+
+ __device__ __forceinline__ void stg_v4_b32(const void *p, uint32_t a, uint32_t b, uint32_t c, uint32_t d) {
+ asm volatile("st.global.v4.b32 [%0], {%1, %2, %3, %4};" :: "l"(p), "r"(a), "r"(b), "r"(c), "r"(d) : "memory");
+ }
+
struct SHAPE { static constexpr char _16x256b[] = ".16x256b"; };
struct NUM { static constexpr char x8[] = ".x8"; static constexpr char x16[] = ".x16"; };
⋯ 56 unchanged lines
CUtensorMap *tmap,
const char *ptr,
uint64_t global_h, uint64_t global_w,
- uint32_t shared_h, uint32_t shared_w
+ uint32_t shared_h, uint32_t shared_w,
+ CUtensorMapL2promotion l2_promo
) {
+ // Host-side 固定开销优化:缓存 cuTensorMapEncodeTiled 的结果,key 包含指针与维度/tiling。
+ // 注意:缓存命中仅在 key 完全一致时复用;否则必须重新 encode 以保证正确性。
+ constexpr int TM_CACHE_SIZE = 8;
+ struct TMCacheEntry {
+ const void *ptr;
+ uint64_t global_h;
+ uint64_t global_w;
+ uint32_t shared_h;
+ uint32_t shared_w;
+ uint32_t l2_promo;
+ CUtensorMap tmap;
+ uint32_t valid;
+ };
+ static TMCacheEntry cache[TM_CACHE_SIZE];
+ #pragma unroll
+ for (int i = 0; i < TM_CACHE_SIZE; i++) {
+ const TMCacheEntry &e = cache[i];
+ if (e.valid &&
+ e.ptr == (const void *)ptr &&
+ e.global_h == global_h &&
+ e.global_w == global_w &&
+ e.shared_h == shared_h &&
+ e.shared_w == shared_w &&
+ e.l2_promo == (uint32_t)l2_promo) {
+ *tmap = e.tmap;
+ return;
+ }
+ }
+
constexpr uint32_t rank = 3;
uint64_t globalDim[rank] = {256, global_h, global_w / 256};
uint64_t globalStrides[rank-1] = {global_w / 2, 128};
uint32_t boxDim[rank] = {256, shared_h, shared_w / 256};
uint32_t elementStrides[rank] = {1, 1, 1};
+ CUtensorMap tmp;
auto err = cuTensorMapEncodeTiled(
- tmap,
+ &tmp,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
rank,
(void *)ptr,
⋯ 3 unchanged lines
elementStrides,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B,
- CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
+ l2_promo,
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
);
ck_cu(err);
+ *tmap = tmp;
+
+ static uint32_t victim = 0;
+ const uint32_t slot = victim++ & (TM_CACHE_SIZE - 1);
+ cache[slot].ptr = (const void *)ptr;
+ cache[slot].global_h = global_h;
+ cache[slot].global_w = global_w;
+ cache[slot].shared_h = shared_h;
+ cache[slot].shared_w = shared_w;
+ cache[slot].l2_promo = (uint32_t)l2_promo;
+ cache[slot].tmap = tmp;
+ cache[slot].valid = 1;
}
- template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
+ template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CACHE_MODE>
__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void gemm_f32_kernel(
const __grid_constant__ CUtensorMap A_tmap,
⋯ 14 unchanged lines
const int bid_m = bid / grid_n;
const int bid_n = bid - bid_m * grid_n;
- const int off_m = bid_m * BLOCK_M;
- const int off_n = bid_n * BLOCK_N;
+ const int off_m = bid_m * BLOCK_M;
+ const int off_n = bid_n * BLOCK_N;
- constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;
+ #if ENABLE_CLK64_DIAG
+ const int diag_base = bid * 8;
+ #endif
+ constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;
+
extern __shared__ __align__(1024) char smem_ptr[];
const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
constexpr int A_size = BLOCK_M * BLOCK_K / 2;
⋯ 22 unchanged lines
constexpr int num_iters = K / BLOCK_K;
if (warp_id == NUM_WARPS - 2 && elect_sync()) {
- const uint64_t cache_A = EVICT_FIRST;
- const uint64_t cache_B = EVICT_LAST;
+ constexpr uint64_t cache_A = (CACHE_MODE == 0) ? EVICT_FIRST : EVICT_LAST;
+ constexpr uint64_t cache_B = (CACHE_MODE == 0) ? EVICT_LAST : EVICT_FIRST;
auto issue_tma = [&](int iter_k, int stage_id) {
const int mbar_addr = tma_mbar_addr + stage_id * 8;
⋯ 27 unchanged lines
mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
issue_tma(iter_k, stage_id);
}
- } else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
- constexpr int MMA_N = BLOCK_N;
- constexpr int MMA_M = 128;
- constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)MMA_N >> 3U << 17U) | ((uint32_t)MMA_M >> 7U << 27U);
+ } else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
+ #if ENABLE_CLK64_DIAG
+ const unsigned long long t0 = clock64();
+ #endif
+ constexpr int MMA_N = BLOCK_N;
+ constexpr int MMA_M = 128;
+ constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)MMA_N >> 3U << 17U) | ((uint32_t)MMA_M >> 7U << 27U);
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
const int stage_id = iter_k % NUM_STAGES;
⋯ 50 unchanged lines
);
}
- asm volatile(
- "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
- :: "r"(mainloop_mbar_addr)
- : "memory"
- );
- } else if (tid < BLOCK_M) {
- mbarrier_wait(mainloop_mbar_addr, 0);
+ asm volatile(
+ "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
+ :: "r"(mainloop_mbar_addr)
+ : "memory"
+ );
+ #if ENABLE_CLK64_DIAG
+ g_clk64_diag[diag_base + 1] = clock64() - t0;
+ #endif
+ } else if (tid < BLOCK_M) {
+ if (elect_sync()) mbarrier_wait(mainloop_mbar_addr, 0);
+ __syncwarp();
asm volatile("tcgen05.fence::after_thread_sync;");
for (int mm = 0; mm < 2; mm++) {
⋯ 16 unchanged lines
}
}
- template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
+ template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CACHE_MODE, int PAD, int SWZ>
__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void gemm_silu_mul_kernel(
const __grid_constant__ CUtensorMap A_tmap,
⋯ 18 unchanged lines
const int off_m = bid_m * BLOCK_M;
const int off_n = bid_n * BLOCK_N;
+ #if ENABLE_CLK64_DIAG
+ const int diag_base = bid * 8;
+ #endif
+
constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;
extern __shared__ __align__(1024) char smem_ptr[];
⋯ 21 unchanged lines
}
__syncthreads();
- constexpr int num_iters = K / BLOCK_K;
+ constexpr int num_iters = K / BLOCK_K;
+
+ if (warp_id == NUM_WARPS - 2 && elect_sync()) {
+ #if ENABLE_CLK64_DIAG
+ const unsigned long long t0 = clock64();
+ #endif
+ constexpr uint64_t cache_A = (CACHE_MODE == 0) ? EVICT_LAST : EVICT_FIRST;
+ constexpr uint64_t cache_B = (CACHE_MODE == 0) ? EVICT_FIRST : EVICT_LAST;
- if (warp_id == NUM_WARPS - 2 && elect_sync()) {
- const uint64_t cache_A = EVICT_LAST;
- const uint64_t cache_B = EVICT_FIRST;
- const uint64_t cache_SFA = EVICT_LAST;
- const uint64_t cache_SFB = EVICT_LAST;
-
auto issue_tma = [&](int iter_k, int stage_id) {
const int mbar_addr = tma_mbar_addr + stage_id * 8;
const int A_smem = smem + stage_id * STAGE_SIZE;
⋯ 18 unchanged lines
);
};
- constexpr int PRELOAD = (num_iters < NUM_STAGES) ? num_iters : NUM_STAGES;
- for (int iter_k = 0; iter_k < PRELOAD; iter_k++) issue_tma(iter_k, iter_k);
- for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
- const int stage_id = iter_k % NUM_STAGES;
- const int mma_phase = (iter_k / NUM_STAGES - 1) & 1;
- mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
- issue_tma(iter_k, stage_id);
- }
- } else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
+ constexpr int PRELOAD = (num_iters < NUM_STAGES) ? num_iters : NUM_STAGES;
+ for (int iter_k = 0; iter_k < PRELOAD; iter_k++) issue_tma(iter_k, iter_k);
+ for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
+ const int stage_id = iter_k % NUM_STAGES;
+ const int mma_phase = (iter_k / NUM_STAGES - 1) & 1;
+ mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
+ issue_tma(iter_k, stage_id);
+ }
+ #if ENABLE_CLK64_DIAG
+ g_clk64_diag[diag_base + 0] = clock64() - t0;
+ #endif
+ } else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
constexpr int MMA_N = BLOCK_N;
constexpr int MMA_M = 128;
constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)MMA_N >> 3U << 17U) | ((uint32_t)MMA_M >> 7U << 27U);
⋯ 59 unchanged lines
: "memory"
);
} else if (tid < BLOCK_M) {
- mbarrier_wait(mainloop_mbar_addr, 0);
+ if (elect_sync()) mbarrier_wait(mainloop_mbar_addr, 0);
+ __syncwarp();
asm volatile("tcgen05.fence::after_thread_sync;");
+ half2 *out_smem = reinterpret_cast<half2 *>(smem_ptr);
+ constexpr int OUT_H2_STRIDE = BLOCK_N / 2 + PAD;
+
for (int mm = 0; mm < 2; mm++) {
float tmp[BLOCK_N / 2];
if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + mm * 16, 0);
⋯ 2 unchanged lines
#pragma unroll
for (int i = 0; i < BLOCK_N / 8; i++) {
- const int row = off_m + warp_id * 32 + mm * 16 + lane_id / 4;
- const int col = off_n + i * 8 + (lane_id & 3) * 2;
+ const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
+ const int lc = i * 8 + (lane_id & 3) * 2;
+ const int row = off_m + lr;
+ const int col = off_n + lc;
const float2 x0 = reinterpret_cast<const float2 *>(G1_ptr + (row + 0) * N + col)[0];
const float2 x8 = reinterpret_cast<const float2 *>(G1_ptr + (row + 8) * N + col)[0];
float2 o0;
float2 o8;
float s;
- s = 1.0f / (1.0f + __expf(-x0.x));
+ s = fast_sigmoid(x0.x);
o0.x = (x0.x * s) * tmp[i * 4 + 0];
- s = 1.0f / (1.0f + __expf(-x0.y));
+ s = fast_sigmoid(x0.y);
o0.y = (x0.y * s) * tmp[i * 4 + 1];
- s = 1.0f / (1.0f + __expf(-x8.x));
+ s = fast_sigmoid(x8.x);
o8.x = (x8.x * s) * tmp[i * 4 + 2];
- s = 1.0f / (1.0f + __expf(-x8.y));
+ s = fast_sigmoid(x8.y);
o8.y = (x8.y * s) * tmp[i * 4 + 3];
- reinterpret_cast<half2 *>(Out_ptr + (row + 0) * N + col)[0] = __float22half2_rn(o0);
- reinterpret_cast<half2 *>(Out_ptr + (row + 8) * N + col)[0] = __float22half2_rn(o8);
+ const int hc = lc >> 1;
+ out_smem[(lr + 0) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 0)] = __float22half2_rn(o0);
+ out_smem[(lr + 8) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 8)] = __float22half2_rn(o8);
}
}
asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
if (warp_id == 0) asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 2));
+
+ constexpr int WRITER_WARPS = BLOCK_M / WARP_SIZE;
+ constexpr int ROWS_PER_WARP = BLOCK_M / WRITER_WARPS;
+ const int r_base = warp_id * ROWS_PER_WARP;
+ for (int rr = 0; rr < ROWS_PER_WARP; rr++) {
+ const int r = r_base + rr;
+ const int row = off_m + r;
+ half2 *dst = reinterpret_cast<half2 *>(Out_ptr + row * N + off_n);
+ const half2 *src = out_smem + r * OUT_H2_STRIDE;
+ if constexpr (BLOCK_N == 64) {
+ dst[lane_id] = src[out_smem_col<SWZ>(lane_id, r)];
+ } else {
+ dst[lane_id] = src[out_smem_col<SWZ>(lane_id, r)];
+ dst[lane_id + 32] = src[out_smem_col<SWZ>(lane_id + 32, r)];
+ }
+ }
}
}
- template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
+ template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int CACHE_MODE, int PAD, int SWZ, int EPILOGUE_MODE, int MBAR_MODE>
__global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void dual_gemm_silu_kernel(
const __grid_constant__ CUtensorMap A_tmap,
⋯ 6 unchanged lines
int M, int N
) {
const int tid = threadIdx.x;
- const int bid = blockIdx.y;
const int lane_id = tid & 31;
const int warp_id = tid >> 5;
- const int grid_n = N / BLOCK_N;
- const int bid_m = bid / grid_n;
- const int bid_n = bid - bid_m * grid_n;
+ const int bid_m = (int)blockIdx.y;
+ const int bid_n = (int)blockIdx.x;
const int off_m = bid_m * BLOCK_M;
const int off_n = bid_n * BLOCK_N;
+ #if ENABLE_CLK64_DIAG
+ const int diag_base = (bid_m * (N / BLOCK_N) + bid_n) * 8;
+ #endif
+
constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;
extern __shared__ __align__(1024) char smem_ptr[];
⋯ 27 unchanged lines
constexpr int num_iters = K / BLOCK_K;
if (warp_id == NUM_WARPS - 2 && elect_sync()) {
- const uint64_t cache_A = EVICT_LAST;
- const uint64_t cache_B = EVICT_FIRST;
- const uint64_t cache_SFA = EVICT_LAST;
- const uint64_t cache_SFB = EVICT_LAST;
+ #if ENABLE_CLK64_DIAG
+ const unsigned long long t_tma = clock64();
+ #endif
+ constexpr uint64_t cache_A = (CACHE_MODE == 0) ? EVICT_LAST : EVICT_FIRST;
+ constexpr uint64_t cache_B = (CACHE_MODE == 0) ? EVICT_FIRST : EVICT_LAST;
auto issue_tma = [&](int iter_k, int stage_id) {
const int mbar_addr = tma_mbar_addr + stage_id * 8;
⋯ 14 unchanged lines
const char *SFA_src = SFA_ptr + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
const char *SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
const char *SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;
- tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_SFA);
- tma_gmem2smem(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cache_SFB);
- tma_gmem2smem(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cache_SFB);
+ tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_A);
+ tma_gmem2smem(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cache_B);
+ tma_gmem2smem(SFB2_smem, SFB2_src, SFB_size, mbar_addr, cache_B);
asm volatile(
"mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
⋯ 10 unchanged lines
mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
issue_tma(iter_k, stage_id);
}
+ #if ENABLE_CLK64_DIAG
+ g_clk64_diag[diag_base + 0] = clock64() - t_tma;
+ #endif
} else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
constexpr int MMA_N = BLOCK_N;
constexpr int MMA_M = 128;
constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)MMA_N >> 3U << 17U) | ((uint32_t)MMA_M >> 7U << 27U);
+ #if ENABLE_CLK64_DIAG
+ const unsigned long long t_mma = clock64();
+ #endif
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
const int stage_id = iter_k % NUM_STAGES;
const int tma_phase = (iter_k / NUM_STAGES) & 1;
⋯ 66 unchanged lines
:: "r"(mainloop_mbar_addr)
: "memory"
);
- } else if (tid < BLOCK_M) {
- mbarrier_wait(mainloop_mbar_addr, 0);
+ #if ENABLE_CLK64_DIAG
+ g_clk64_diag[diag_base + 1] = clock64() - t_mma;
+ #endif
+ } else if (tid < BLOCK_M) {
+ if constexpr (MBAR_MODE == 0) {
+ if (warp_id == 0 && lane_id == 0) mbarrier_wait(mainloop_mbar_addr, 0);
+ asm volatile("bar.sync 3, %0;" :: "r"(BLOCK_M) : "memory");
+ } else {
+ mbarrier_wait(mainloop_mbar_addr, 0);
+ }
asm volatile("tcgen05.fence::after_thread_sync;");
+ #if ENABLE_CLK64_DIAG
+ const unsigned long long t_epi = clock64();
+ #endif
+
+ half2 *out_smem;
+ constexpr int OUT_H2_STRIDE = BLOCK_N / 2 + PAD;
- constexpr int OUT_H2_STRIDE = BLOCK_N / 2 + 1;
- half2 *out_smem = reinterpret_cast<half2 *>(smem_ptr);
+ if constexpr (EPILOGUE_MODE == 1 && BLOCK_N == 64) {
+ out_smem = reinterpret_cast<half2 *>(smem_ptr);
- for (int mm = 0; mm < 2; mm++) {
- const int lr_base = warp_id * 32 + mm * 16;
- const int lr_lane = lane_id >> 2;
- const int lc_lane = (lane_id & 3) << 1;
- const int lr = lr_base + lr_lane;
+ for (int mm = 0; mm < 2; mm++) {
+ float g1_reg[BLOCK_N / 2];
+ tcgen05_ld_16x256bx8(g1_reg, warp_id * 32 + mm * 16, 0);
+ asm volatile("tcgen05.wait::ld.sync.aligned;");
- if constexpr (BLOCK_N == 64) {
- half2 silu_h2[16];
- {
- float tmp0[32];
- tcgen05_ld_16x256bx8(tmp0, lr_base, 0);
- asm volatile("tcgen05.wait::ld.sync.aligned;");
+ float tmp[BLOCK_N / 2];
+ tcgen05_ld_16x256bx8(tmp, warp_id * 32 + mm * 16, OUT2_TMEM);
+ asm volatile("tcgen05.wait::ld.sync.aligned;");
- #pragma unroll
- for (int i = 0; i < 8; i++) {
- const float x0 = tmp0[i * 4 + 0];
- const float x1 = tmp0[i * 4 + 1];
- const float x2 = tmp0[i * 4 + 2];
- const float x3 = tmp0[i * 4 + 3];
+ #pragma unroll
+ for (int i = 0; i < BLOCK_N / 8; i++) {
+ const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
+ const int lc = i * 8 + (lane_id & 3) * 2;
+ const float2 x0 = float2{g1_reg[i * 4 + 0], g1_reg[i * 4 + 1]};
+ const float2 x8 = float2{g1_reg[i * 4 + 2], g1_reg[i * 4 + 3]};
- const float e0 = __expf(-x0);
- const float e1 = __expf(-x1);
- const float e2 = __expf(-x2);
- const float e3 = __expf(-x3);
- const float s0 = 1.0f / (1.0f + e0);
- const float s1 = 1.0f / (1.0f + e1);
- const float s2 = 1.0f / (1.0f + e2);
- const float s3 = 1.0f / (1.0f + e3);
+ float2 o0;
+ float2 o8;
+ float ex0, ex1, ex2, ex3;
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex0) : "f"((-x0.x) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex1) : "f"((-x0.y) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex2) : "f"((-x8.x) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex3) : "f"((-x8.y) * 1.4426950408889634f));
- const float y0 = x0 * s0;
- const float y1 = x1 * s1;
- const float y2 = x2 * s2;
- const float y3 = x3 * s3;
- silu_h2[i * 2 + 0] = __floats2half2_rn(y0, y1);
- silu_h2[i * 2 + 1] = __floats2half2_rn(y2, y3);
- }
- }
+ const float d0 = 1.0f + ex0;
+ const float d1 = 1.0f + ex1;
+ const float d2 = 1.0f + ex2;
+ const float d3 = 1.0f + ex3;
- float tmp[32];
- tcgen05_ld_16x256bx8(tmp, lr_base, OUT2_TMEM);
- asm volatile("tcgen05.wait::ld.sync.aligned;");
+ float s0, s1, s2, s3;
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s0) : "f"(d0));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s1) : "f"(d1));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s2) : "f"(d2));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s3) : "f"(d3));
- #pragma unroll
- for (int i = 0; i < 8; i++) {
- const float2 x0 = __half22float2(silu_h2[i * 2 + 0]);
- const float2 x8 = __half22float2(silu_h2[i * 2 + 1]);
- const float t0 = tmp[i * 4 + 0];
- const float t1 = tmp[i * 4 + 1];
- const float t2 = tmp[i * 4 + 2];
- const float t3 = tmp[i * 4 + 3];
+ o0.x = (x0.x * s0) * tmp[i * 4 + 0];
+ o0.y = (x0.y * s1) * tmp[i * 4 + 1];
+ o8.x = (x8.x * s2) * tmp[i * 4 + 2];
+ o8.y = (x8.y * s3) * tmp[i * 4 + 3];
- float2 o0;
- float2 o8;
- o0.x = x0.x * t0;
- o0.y = x0.y * t1;
- o8.x = x8.x * t2;
- o8.y = x8.y * t3;
+ const int hc = lc >> 1;
+ out_smem[(lr + 0) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 0)] = __float22half2_rn(o0);
+ out_smem[(lr + 8) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 8)] = __float22half2_rn(o8);
+ }
+ }
+ } else if constexpr (EPILOGUE_MODE == 2 && BLOCK_N == 128) {
+ out_smem = reinterpret_cast<half2 *>(smem_ptr);
- const int lc = i * 8 + lc_lane;
- out_smem[(lr + 0) * OUT_H2_STRIDE + (lc >> 1)] = __float22half2_rn(o0);
- out_smem[(lr + 8) * OUT_H2_STRIDE + (lc >> 1)] = __float22half2_rn(o8);
- }
- } else {
- #pragma unroll
- for (int pp = 0; pp < 2; pp++) {
- half2 silu_h2[16];
- {
- float tmp0[32];
- tcgen05_ld_16x256bx8(tmp0, lr_base, pp * (BLOCK_N / 2));
- asm volatile("tcgen05.wait::ld.sync.aligned;");
+ #pragma unroll
+ for (int mm = 0; mm < 2; mm++) {
+ const int tmem_row = warp_id * 32 + mm * 16;
- #pragma unroll
- for (int i = 0; i < 8; i++) {
- const float x0 = tmp0[i * 4 + 0];
- const float x1 = tmp0[i * 4 + 1];
- const float x2 = tmp0[i * 4 + 2];
- const float x3 = tmp0[i * 4 + 3];
+ #pragma unroll
+ for (int seg = 0; seg < 2; seg++) {
+ constexpr int SEG_N = 64;
+ constexpr int SEG_H2 = SEG_N / 2;
- const float e0 = __expf(-x0);
- const float e1 = __expf(-x1);
- const float e2 = __expf(-x2);
- const float e3 = __expf(-x3);
- const float s0 = 1.0f / (1.0f + e0);
- const float s1 = 1.0f / (1.0f + e1);
- const float s2 = 1.0f / (1.0f + e2);
- const float s3 = 1.0f / (1.0f + e3);
+ float g1_reg[SEG_H2];
+ tcgen05_ld_16x256bx8(g1_reg, tmem_row, seg * SEG_N);
+ asm volatile("tcgen05.wait::ld.sync.aligned;");
- const float y0 = x0 * s0;
- const float y1 = x1 * s1;
- const float y2 = x2 * s2;
- const float y3 = x3 * s3;
- silu_h2[i * 2 + 0] = __floats2half2_rn(y0, y1);
- silu_h2[i * 2 + 1] = __floats2half2_rn(y2, y3);
- }
- }
+ float tmp[SEG_H2];
+ tcgen05_ld_16x256bx8(tmp, tmem_row, OUT2_TMEM + seg * SEG_N);
+ asm volatile("tcgen05.wait::ld.sync.aligned;");
- float tmp[32];
- tcgen05_ld_16x256bx8(tmp, lr_base, OUT2_TMEM + pp * (BLOCK_N / 2));
- asm volatile("tcgen05.wait::ld.sync.aligned;");
+ #pragma unroll
+ for (int i = 0; i < SEG_N / 8; i++) {
+ const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
+ const int lc = seg * SEG_N + i * 8 + (lane_id & 3) * 2;
+ const float2 x0 = float2{g1_reg[i * 4 + 0], g1_reg[i * 4 + 1]};
+ const float2 x8 = float2{g1_reg[i * 4 + 2], g1_reg[i * 4 + 3]};
- const int lc_base = (pp << 6) + lc_lane;
- #pragma unroll
- for (int i = 0; i < 8; i++) {
- const float2 x0 = __half22float2(silu_h2[i * 2 + 0]);
- const float2 x8 = __half22float2(silu_h2[i * 2 + 1]);
- const float t0 = tmp[i * 4 + 0];
- const float t1 = tmp[i * 4 + 1];
- const float t2 = tmp[i * 4 + 2];
- const float t3 = tmp[i * 4 + 3];
+ float2 o0;
+ float2 o8;
+ float s;
+ s = fast_sigmoid(x0.x);
+ o0.x = (x0.x * s) * tmp[i * 4 + 0];
+ s = fast_sigmoid(x0.y);
+ o0.y = (x0.y * s) * tmp[i * 4 + 1];
+ s = fast_sigmoid(x8.x);
+ o8.x = (x8.x * s) * tmp[i * 4 + 2];
+ s = fast_sigmoid(x8.y);
+ o8.y = (x8.y * s) * tmp[i * 4 + 3];
- float2 o0;
- float2 o8;
- o0.x = x0.x * t0;
- o0.y = x0.y * t1;
- o8.x = x8.x * t2;
- o8.y = x8.y * t3;
+ const int hc = lc >> 1;
+ out_smem[(lr + 0) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 0)] = __float22half2_rn(o0);
+ out_smem[(lr + 8) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 8)] = __float22half2_rn(o8);
+ }
+ }
+ }
+ } else if constexpr ((EPILOGUE_MODE == 4 || EPILOGUE_MODE == 6 || EPILOGUE_MODE == 7 || EPILOGUE_MODE == 8) && BLOCK_N == 64) {
+ // 无 out_smem:直写 global(v2/v4),依赖 bar.sync 确保 tcgen05.dealloc 的安全时序。
+ #pragma unroll
+ for (int mm = 0; mm < 2; mm++) {
+ const int tmem_row = warp_id * 32 + mm * 16;
+ constexpr int SEG_N = 64;
+ constexpr int SEG_H2 = SEG_N / 2;
- const int lc = lc_base + i * 8;
- out_smem[(lr + 0) * OUT_H2_STRIDE + (lc >> 1)] = __float22half2_rn(o0);
- out_smem[(lr + 8) * OUT_H2_STRIDE + (lc >> 1)] = __float22half2_rn(o8);
+ float g1_reg[SEG_H2];
+ float tmp[SEG_H2];
+ tcgen05_ld_16x256bx8(g1_reg, tmem_row, 0);
+ tcgen05_ld_16x256bx8(tmp, tmem_row, OUT2_TMEM);
+ asm volatile("tcgen05.wait::ld.sync.aligned;");
+
+ const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
+ const int row0 = off_m + lr;
+ const int row8 = row0 + 8;
+ half2 *out0_h2 = reinterpret_cast<half2 *>(Out_ptr + row0 * N + off_n);
+ half2 *out8_h2 = reinterpret_cast<half2 *>(Out_ptr + row8 * N + off_n);
+
+ #pragma unroll
+ for (int j = 0; j < SEG_N / 16; j++) {
+ const int idx0 = j * 8;
+ const int idx1 = idx0 + 4;
+
+ float2 o0_0;
+ float2 o8_0;
+ if constexpr (EPILOGUE_MODE == 7) {
+ o0_0.x = tmp[idx0 + 0];
+ o0_0.y = tmp[idx0 + 1];
+ o8_0.x = tmp[idx0 + 2];
+ o8_0.y = tmp[idx0 + 3];
+ } else if constexpr (EPILOGUE_MODE == 8) {
+ o0_0.x = g1_reg[idx0 + 0];
+ o0_0.y = g1_reg[idx0 + 1];
+ o8_0.x = g1_reg[idx0 + 2];
+ o8_0.y = g1_reg[idx0 + 3];
+ } else {
+ const float2 x0 = float2{g1_reg[idx0 + 0], g1_reg[idx0 + 1]};
+ const float2 x8 = float2{g1_reg[idx0 + 2], g1_reg[idx0 + 3]};
+
+ float ex0, ex1, ex2, ex3;
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex0) : "f"((-x0.x) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex1) : "f"((-x0.y) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex2) : "f"((-x8.x) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex3) : "f"((-x8.y) * 1.4426950408889634f));
+
+ const float d0 = 1.0f + ex0;
+ const float d1 = 1.0f + ex1;
+ const float d2 = 1.0f + ex2;
+ const float d3 = 1.0f + ex3;
+
+ float s0, s1, s2, s3;
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s0) : "f"(d0));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s1) : "f"(d1));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s2) : "f"(d2));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s3) : "f"(d3));
+
+ const float p0 = x0.x * tmp[idx0 + 0];
+ const float p1 = x0.y * tmp[idx0 + 1];
+ const float p2 = x8.x * tmp[idx0 + 2];
+ const float p3 = x8.y * tmp[idx0 + 3];
+
+ o0_0.x = p0 * s0;
+ o0_0.y = p1 * s1;
+ o8_0.x = p2 * s2;
+ o8_0.y = p3 * s3;
+ }
+
+ float2 o0_1;
+ float2 o8_1;
+ if constexpr (EPILOGUE_MODE == 7) {
+ o0_1.x = tmp[idx1 + 0];
+ o0_1.y = tmp[idx1 + 1];
+ o8_1.x = tmp[idx1 + 2];
+ o8_1.y = tmp[idx1 + 3];
+ } else if constexpr (EPILOGUE_MODE == 8) {
+ o0_1.x = g1_reg[idx1 + 0];
+ o0_1.y = g1_reg[idx1 + 1];
+ o8_1.x = g1_reg[idx1 + 2];
+ o8_1.y = g1_reg[idx1 + 3];
+ } else {
+ const float2 x0 = float2{g1_reg[idx1 + 0], g1_reg[idx1 + 1]};
+ const float2 x8 = float2{g1_reg[idx1 + 2], g1_reg[idx1 + 3]};
+
+ float ex0, ex1, ex2, ex3;
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex0) : "f"((-x0.x) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex1) : "f"((-x0.y) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex2) : "f"((-x8.x) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex3) : "f"((-x8.y) * 1.4426950408889634f));
+
+ const float d0 = 1.0f + ex0;
+ const float d1 = 1.0f + ex1;
+ const float d2 = 1.0f + ex2;
+ const float d3 = 1.0f + ex3;
+
+ float s0, s1, s2, s3;
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s0) : "f"(d0));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s1) : "f"(d1));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s2) : "f"(d2));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s3) : "f"(d3));
+
+ const float p0 = x0.x * tmp[idx1 + 0];
+ const float p1 = x0.y * tmp[idx1 + 1];
+ const float p2 = x8.x * tmp[idx1 + 2];
+ const float p3 = x8.y * tmp[idx1 + 3];
+
+ o0_1.x = p0 * s0;
+ o0_1.y = p1 * s1;
+ o8_1.x = p2 * s2;
+ o8_1.y = p3 * s3;
+ }
+
+ const uint32_t u0_0 = half2_as_u32(__float22half2_rn(o0_0));
+ const uint32_t u1_0 = half2_as_u32(__float22half2_rn(o0_1));
+ const uint32_t u0_8 = half2_as_u32(__float22half2_rn(o8_0));
+ const uint32_t u1_8 = half2_as_u32(__float22half2_rn(o8_1));
+
+ const uint64_t both0 = ((uint64_t)u1_0 << 32) | (uint64_t)u0_0;
+ const uint64_t both8 = ((uint64_t)u1_8 << 32) | (uint64_t)u0_8;
+ const uint64_t peer0 = __shfl_xor_sync(0xFFFFFFFF, both0, 1);
+ const uint64_t peer8 = __shfl_xor_sync(0xFFFFFFFF, both8, 1);
+
+ const uint32_t peer_u0_0 = (uint32_t)peer0;
+ const uint32_t peer_u1_0 = (uint32_t)(peer0 >> 32);
+ const uint32_t peer_u0_8 = (uint32_t)peer8;
+ const uint32_t peer_u1_8 = (uint32_t)(peer8 >> 32);
+
+ const int t = lane_id & 3;
+ const int t_odd = t & 1;
+ const uint32_t lo0 = t_odd ? peer_u1_0 : u0_0;
+ const uint32_t hi0 = t_odd ? u1_0 : peer_u0_0;
+ const uint32_t lo8 = t_odd ? peer_u1_8 : u0_8;
+ const uint32_t hi8 = t_odd ? u1_8 : peer_u0_8;
+
+ if constexpr ((EPILOGUE_MODE == 4 || EPILOGUE_MODE == 7 || EPILOGUE_MODE == 8) || (EPILOGUE_MODE == 6 && !ENABLE_V4_EPILOGUE)) {
+ const int store_off = (t & 2) + ((t & 1) << 2);
+ const half2 *out0_p = out0_h2 + j * 8 + store_off;
+ const half2 *out8_p = out8_h2 + j * 8 + store_off;
+ stg_v2_b32(out0_p, lo0, hi0);
+ stg_v2_b32(out8_p, lo8, hi8);
+ } else {
+ const int base_lane = lane_id & ~3;
+ if (t == 0) {
+ const uint32_t lo0_p = __shfl_sync(0xFFFFFFFF, lo0, base_lane + 2);
+ const uint32_t hi0_p = __shfl_sync(0xFFFFFFFF, hi0, base_lane + 2);
+ const uint32_t lo8_p = __shfl_sync(0xFFFFFFFF, lo8, base_lane + 2);
+ const uint32_t hi8_p = __shfl_sync(0xFFFFFFFF, hi8, base_lane + 2);
+ const half2 *out0_p = out0_h2 + j * 8;
+ const half2 *out8_p = out8_h2 + j * 8;
+ stg_v4_b32(out0_p, lo0, hi0, lo0_p, hi0_p);
+ stg_v4_b32(out8_p, lo8, hi8, lo8_p, hi8_p);
+ } else if (t == 1) {
+ const uint32_t lo0_p = __shfl_sync(0xFFFFFFFF, lo0, base_lane + 3);
+ const uint32_t hi0_p = __shfl_sync(0xFFFFFFFF, hi0, base_lane + 3);
+ const uint32_t lo8_p = __shfl_sync(0xFFFFFFFF, lo8, base_lane + 3);
+ const uint32_t hi8_p = __shfl_sync(0xFFFFFFFF, hi8, base_lane + 3);
+ const half2 *out0_p = out0_h2 + j * 8 + 4;
+ const half2 *out8_p = out8_h2 + j * 8 + 4;
+ stg_v4_b32(out0_p, lo0, hi0, lo0_p, hi0_p);
+ stg_v4_b32(out8_p, lo8, hi8, lo8_p, hi8_p);
+ }
}
}
}
- }
+
+ asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
+ if (warp_id == 0) asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 4));
+ #if ENABLE_CLK64_DIAG
+ if (lane_id == 0) g_clk64_diag[diag_base + 2 + warp_id] = clock64() - t_epi;
+ #endif
+ return;
+ } else if constexpr ((EPILOGUE_MODE == 3 || EPILOGUE_MODE == 5 || EPILOGUE_MODE == 7 || EPILOGUE_MODE == 8) && BLOCK_N == 128) {
+ // 无 out_smem:直接写回 global,依赖 bar.sync 确保 tcgen05.dealloc 的安全时序。
+ #pragma unroll
+ for (int mm = 0; mm < 2; mm++) {
+ const int tmem_row = warp_id * 32 + mm * 16;
+ #pragma unroll
+ for (int seg = 0; seg < 2; seg++) {
+ constexpr int SEG_N = 64;
+ constexpr int SEG_H2 = SEG_N / 2;
+
+ float g1_reg[SEG_H2];
+ float tmp[SEG_H2];
+ tcgen05_ld_16x256bx8(g1_reg, tmem_row, seg * SEG_N);
+ tcgen05_ld_16x256bx8(tmp, tmem_row, OUT2_TMEM + seg * SEG_N);
+ asm volatile("tcgen05.wait::ld.sync.aligned;");
+
+ const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
+ const int row0 = off_m + lr;
+ const int row8 = row0 + 8;
+ half2 *out0_h2 = reinterpret_cast<half2 *>(Out_ptr + row0 * N + off_n) + seg * 32;
+ half2 *out8_h2 = reinterpret_cast<half2 *>(Out_ptr + row8 * N + off_n) + seg * 32;
+
+ #pragma unroll
+ for (int j = 0; j < SEG_N / 16; j++) {
+ const int idx0 = j * 8;
+ const int idx1 = idx0 + 4;
+
+ const float2 x0_0 = float2{g1_reg[idx0 + 0], g1_reg[idx0 + 1]};
+ const float2 x8_0 = float2{g1_reg[idx0 + 2], g1_reg[idx0 + 3]};
+ const float2 x0_1 = float2{g1_reg[idx1 + 0], g1_reg[idx1 + 1]};
+ const float2 x8_1 = float2{g1_reg[idx1 + 2], g1_reg[idx1 + 3]};
+
+ float ex0_0, ex0_1, ex1_0, ex1_1, ex2_0, ex2_1, ex3_0, ex3_1;
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex0_0) : "f"((-x0_0.x) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex0_1) : "f"((-x0_1.x) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex1_0) : "f"((-x0_0.y) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex1_1) : "f"((-x0_1.y) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex2_0) : "f"((-x8_0.x) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex2_1) : "f"((-x8_1.x) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex3_0) : "f"((-x8_0.y) * 1.4426950408889634f));
+ asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(ex3_1) : "f"((-x8_1.y) * 1.4426950408889634f));
+
+ const float d0_0 = 1.0f + ex0_0;
+ const float d0_1 = 1.0f + ex0_1;
+ const float d1_0 = 1.0f + ex1_0;
+ const float d1_1 = 1.0f + ex1_1;
+ const float d2_0 = 1.0f + ex2_0;
+ const float d2_1 = 1.0f + ex2_1;
+ const float d3_0 = 1.0f + ex3_0;
+ const float d3_1 = 1.0f + ex3_1;
+
+ float s0_0, s0_1, s1_0, s1_1, s2_0, s2_1, s3_0, s3_1;
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s0_0) : "f"(d0_0));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s0_1) : "f"(d0_1));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s1_0) : "f"(d1_0));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s1_1) : "f"(d1_1));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s2_0) : "f"(d2_0));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s2_1) : "f"(d2_1));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s3_0) : "f"(d3_0));
+ asm("rcp.approx.ftz.f32 %0, %1;" : "=f"(s3_1) : "f"(d3_1));
+
+ float2 o0_0;
+ float2 o8_0;
+ float2 o0_1;
+ float2 o8_1;
+
+ if constexpr (EPILOGUE_MODE == 7) {
+ o0_0.x = tmp[idx0 + 0];
+ o0_0.y = tmp[idx0 + 1];
+ o8_0.x = tmp[idx0 + 2];
+ o8_0.y = tmp[idx0 + 3];
+ o0_1.x = tmp[idx1 + 0];
+ o0_1.y = tmp[idx1 + 1];
+ o8_1.x = tmp[idx1 + 2];
+ o8_1.y = tmp[idx1 + 3];
+ } else if constexpr (EPILOGUE_MODE == 8) {
+ o0_0.x = x0_0.x;
+ o0_0.y = x0_0.y;
+ o8_0.x = x8_0.x;
+ o8_0.y = x8_0.y;
+ o0_1.x = x0_1.x;
+ o0_1.y = x0_1.y;
+ o8_1.x = x8_1.x;
+ o8_1.y = x8_1.y;
+ } else {
+ o0_0.x = (x0_0.x * tmp[idx0 + 0]) * s0_0;
+ o0_0.y = (x0_0.y * tmp[idx0 + 1]) * s1_0;
+ o8_0.x = (x8_0.x * tmp[idx0 + 2]) * s2_0;
+ o8_0.y = (x8_0.y * tmp[idx0 + 3]) * s3_0;
+ o0_1.x = (x0_1.x * tmp[idx1 + 0]) * s0_1;
+ o0_1.y = (x0_1.y * tmp[idx1 + 1]) * s1_1;
+ o8_1.x = (x8_1.x * tmp[idx1 + 2]) * s2_1;
+ o8_1.y = (x8_1.y * tmp[idx1 + 3]) * s3_1;
+ }
+
+ const uint32_t u0_0 = half2_as_u32(__float22half2_rn(o0_0));
+ const uint32_t u1_0 = half2_as_u32(__float22half2_rn(o0_1));
+ const uint32_t u0_8 = half2_as_u32(__float22half2_rn(o8_0));
+ const uint32_t u1_8 = half2_as_u32(__float22half2_rn(o8_1));
+
+ const uint64_t both0 = ((uint64_t)u1_0 << 32) | (uint64_t)u0_0;
+ const uint64_t both8 = ((uint64_t)u1_8 << 32) | (uint64_t)u0_8;
+ const uint64_t peer0 = __shfl_xor_sync(0xFFFFFFFF, both0, 1);
+ const uint64_t peer8 = __shfl_xor_sync(0xFFFFFFFF, both8, 1);
+
+ const uint32_t peer_u0_0 = (uint32_t)peer0;
+ const uint32_t peer_u1_0 = (uint32_t)(peer0 >> 32);
+ const uint32_t peer_u0_8 = (uint32_t)peer8;
+ const uint32_t peer_u1_8 = (uint32_t)(peer8 >> 32);
+
+ const int t = lane_id & 3;
+ const int t_odd = t & 1;
+ const uint32_t lo0 = t_odd ? peer_u1_0 : u0_0;
+ const uint32_t hi0 = t_odd ? u1_0 : peer_u0_0;
+ const uint32_t lo8 = t_odd ? peer_u1_8 : u0_8;
+ const uint32_t hi8 = t_odd ? u1_8 : peer_u0_8;
+
+ if constexpr ((EPILOGUE_MODE == 3 || EPILOGUE_MODE == 7 || EPILOGUE_MODE == 8) || (EPILOGUE_MODE == 5 && !ENABLE_V4_EPILOGUE)) {
+ const int store_off = (t & 2) + ((t & 1) << 2);
+ const half2 *out0_p = out0_h2 + j * 8 + store_off;
+ const half2 *out8_p = out8_h2 + j * 8 + store_off;
+ stg_v2_b32(out0_p, lo0, hi0);
+ stg_v2_b32(out8_p, lo8, hi8);
+ } else {
+ const int base_lane = lane_id & ~3;
+ if (t == 0) {
+ const uint32_t lo0_p = __shfl_sync(0xFFFFFFFF, lo0, base_lane + 2);
+ const uint32_t hi0_p = __shfl_sync(0xFFFFFFFF, hi0, base_lane + 2);
+ const uint32_t lo8_p = __shfl_sync(0xFFFFFFFF, lo8, base_lane + 2);
+ const uint32_t hi8_p = __shfl_sync(0xFFFFFFFF, hi8, base_lane + 2);
+ const half2 *out0_p = out0_h2 + j * 8;
+ const half2 *out8_p = out8_h2 + j * 8;
+ stg_v4_b32(out0_p, lo0, hi0, lo0_p, hi0_p);
+ stg_v4_b32(out8_p, lo8, hi8, lo8_p, hi8_p);
+ } else if (t == 1) {
+ const uint32_t lo0_p = __shfl_sync(0xFFFFFFFF, lo0, base_lane + 3);
+ const uint32_t hi0_p = __shfl_sync(0xFFFFFFFF, hi0, base_lane + 3);
+ const uint32_t lo8_p = __shfl_sync(0xFFFFFFFF, lo8, base_lane + 3);
+ const uint32_t hi8_p = __shfl_sync(0xFFFFFFFF, hi8, base_lane + 3);
+ const half2 *out0_p = out0_h2 + j * 8 + 4;
+ const half2 *out8_p = out8_h2 + j * 8 + 4;
+ stg_v4_b32(out0_p, lo0, hi0, lo0_p, hi0_p);
+ stg_v4_b32(out8_p, lo8, hi8, lo8_p, hi8_p);
+ }
+ }
+ }
+ }
+ }
+
+ asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
+ if (warp_id == 0) asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 4));
+ #if ENABLE_CLK64_DIAG
+ if (lane_id == 0) g_clk64_diag[diag_base + 2 + warp_id] = clock64() - t_epi;
+ #endif
+ return;
+ } else {
+ constexpr int G1_STRIDE = BLOCK_N + 8;
+ float *g1_smem = reinterpret_cast<float *>(smem_ptr);
+ out_smem = reinterpret_cast<half2 *>(g1_smem + BLOCK_M * G1_STRIDE);
+
+ for (int mm = 0; mm < 2; mm++) {
+ float tmp[BLOCK_N / 2];
+ if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + mm * 16, 0);
+ else tcgen05_ld_16x256bx16(tmp, warp_id * 32 + mm * 16, 0);
+ asm volatile("tcgen05.wait::ld.sync.aligned;");
+
+ #pragma unroll
+ for (int i = 0; i < BLOCK_N / 8; i++) {
+ const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
+ const int lc = i * 8 + (lane_id & 3) * 2;
+ reinterpret_cast<float2 *>(g1_smem + (lr + 0) * G1_STRIDE + lc)[0] = float2{tmp[i * 4 + 0], tmp[i * 4 + 1]};
+ reinterpret_cast<float2 *>(g1_smem + (lr + 8) * G1_STRIDE + lc)[0] = float2{tmp[i * 4 + 2], tmp[i * 4 + 3]};
+ }
+ }
+
+ for (int mm = 0; mm < 2; mm++) {
+ float tmp[BLOCK_N / 2];
+ if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + mm * 16, OUT2_TMEM);
+ else tcgen05_ld_16x256bx16(tmp, warp_id * 32 + mm * 16, OUT2_TMEM);
+ asm volatile("tcgen05.wait::ld.sync.aligned;");
+
+ #pragma unroll
+ for (int i = 0; i < BLOCK_N / 8; i++) {
+ const int lr = warp_id * 32 + mm * 16 + lane_id / 4;
+ const int lc = i * 8 + (lane_id & 3) * 2;
+ const float2 x0 = reinterpret_cast<const float2 *>(g1_smem + (lr + 0) * G1_STRIDE + lc)[0];
+ const float2 x8 = reinterpret_cast<const float2 *>(g1_smem + (lr + 8) * G1_STRIDE + lc)[0];
+
+ float2 o0;
+ float2 o8;
+ float s;
+ s = fast_sigmoid(x0.x);
+ o0.x = (x0.x * s) * tmp[i * 4 + 0];
+ s = fast_sigmoid(x0.y);
+ o0.y = (x0.y * s) * tmp[i * 4 + 1];
+ s = fast_sigmoid(x8.x);
+ o8.x = (x8.x * s) * tmp[i * 4 + 2];
+ s = fast_sigmoid(x8.y);
+ o8.y = (x8.y * s) * tmp[i * 4 + 3];
+
+ const int hc = lc >> 1;
+ out_smem[(lr + 0) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 0)] = __float22half2_rn(o0);
+ out_smem[(lr + 8) * OUT_H2_STRIDE + out_smem_col<SWZ>(hc, lr + 8)] = __float22half2_rn(o8);
+ }
+ }
+ }
+
asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
if (warp_id == 0) asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 4));
constexpr int WRITER_WARPS = BLOCK_M / WARP_SIZE;
- for (int r = warp_id; r < BLOCK_M; r += WRITER_WARPS) {
+ constexpr int ROWS_PER_WARP = BLOCK_M / WRITER_WARPS;
+ const int r_base = warp_id * ROWS_PER_WARP;
+ for (int rr = 0; rr < ROWS_PER_WARP; rr++) {
+ const int r = r_base + rr;
const int row = off_m + r;
half2 *dst = reinterpret_cast<half2 *>(Out_ptr + row * N + off_n);
const half2 *src = out_smem + r * OUT_H2_STRIDE;
if constexpr (BLOCK_N == 64) {
- dst[lane_id] = src[lane_id];
+ dst[lane_id] = src[out_smem_col<SWZ>(lane_id, r)];
} else {
- dst[lane_id] = src[lane_id];
- dst[lane_id + 32] = src[lane_id + 32];
+ dst[lane_id] = src[out_smem_col<SWZ>(lane_id, r)];
+ dst[lane_id + 32] = src[out_smem_col<SWZ>(lane_id + 32, r)];
}
}
}
⋯ 9 unchanged lines
) {
const int M = (int)A.size(0);
const int N = (int)B.size(0);
+ const int grid_m = M / BLOCK_M;
+ const int grid_n = N / BLOCK_N;
+ const bool prefer_cache_A = grid_n >= grid_m;
auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());
auto B_ptr = reinterpret_cast<const char *>(B.data_ptr());
⋯ 2 unchanged lines
auto C_ptr = reinterpret_cast<float *>(C.data_ptr());
CUtensorMap A_tmap, B_tmap;
- init_AB_tmap(&A_tmap, A_ptr, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K);
- init_AB_tmap(&B_tmap, B_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K);
+ init_AB_tmap(&A_tmap, A_ptr, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K, CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE);
+ init_AB_tmap(&B_tmap, B_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K, CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE);
- dim3 grid(1, (unsigned)((M / BLOCK_M) * (N / BLOCK_N)));
+ dim3 grid(1, (unsigned)(grid_m * grid_n));
const int tb_size = BLOCK_M + 2 * WARP_SIZE;
- const int AB_size = (BLOCK_M + BLOCK_N) * (BLOCK_K / 2);
- const int SFAB_size = 128 * (BLOCK_K / 16) * 2;
- const int smem_size = (AB_size + SFAB_size) * NUM_STAGES;
+ constexpr int AB_size = (BLOCK_M + BLOCK_N) * (BLOCK_K / 2);
+ constexpr int SFAB_size = 128 * (BLOCK_K / 16) * 2;
+ constexpr int smem_size = (AB_size + SFAB_size) * NUM_STAGES;
- auto kptr = gemm_f32_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>;
- if (smem_size > 48'000) cudaFuncSetAttribute(kptr, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
+ auto kptr0 = gemm_f32_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, 0>;
+ auto kptr1 = gemm_f32_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, 1>;
+ auto kptr = prefer_cache_A ? kptr1 : kptr0;
+ if constexpr (smem_size > 48'000) {
+ static int attr0 = 0;
+ static int attr1 = 0;
+ if (prefer_cache_A) {
+ if (__builtin_expect(attr1 == 0, 0)) {
+ cudaFuncSetAttribute(kptr1, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
+ attr1 = 1;
+ }
+ } else {
+ if (__builtin_expect(attr0 == 0, 0)) {
+ cudaFuncSetAttribute(kptr0, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
+ attr0 = 1;
+ }
+ }
+ }
kptr<<<grid, tb_size, smem_size>>>(A_tmap, B_tmap, SFA_ptr, SFB_ptr, C_ptr, M, N);
}
- template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>
+ template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES, int PAD, int SWZ>
static inline void launch_gemm_silu_mul(
const at::Tensor& A,
const at::Tensor& B,
⋯ 4 unchanged lines
) {
const int M = (int)A.size(0);
const int N = (int)B.size(0);
+ const int grid_m = M / BLOCK_M;
+ const int grid_n = N / BLOCK_N;
+ const bool prefer_cache_A = grid_n >= grid_m;
auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());
auto B_ptr = reinterpret_cast<const char *>(B.data_ptr());
⋯ 3 unchanged lines
auto Out_ptr = reinterpret_cast<half *>(out.data_ptr());
CUtensorMap A_tmap, B_tmap;
- init_AB_tmap(&A_tmap, A_ptr, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K);
- init_AB_tmap(&B_tmap, B_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K);
+ init_AB_tmap(&A_tmap, A_ptr, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K, CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE);
+ init_AB_tmap(&B_tmap, B_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K, CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE);
- dim3 grid(1, (unsigned)((M / BLOCK_M) * (N / BLOCK_N)));
+ dim3 grid(1, (unsigned)(grid_m * grid_n));
const int tb_size = BLOCK_M + 2 * WARP_SIZE;
- const int AB_size = (BLOCK_M + BLOCK_N) * (BLOCK_K / 2);
- const int SFAB_size = 128 * (BLOCK_K / 16) * 2;
- const int smem_size = (AB_size + SFAB_size) * NUM_STAGES;
+ constexpr int AB_size = (BLOCK_M + BLOCK_N) * (BLOCK_K / 2);
+ constexpr int SFAB_size = 128 * (BLOCK_K / 16) * 2;
+ constexpr int smem_stage = (AB_size + SFAB_size) * NUM_STAGES;
+ constexpr int OUT_H2_STRIDE = BLOCK_N / 2 + PAD;
+ constexpr int smem_scratch = BLOCK_M * OUT_H2_STRIDE * (int)sizeof(half2);
+ constexpr int smem_size = (smem_stage > smem_scratch) ? smem_stage : smem_scratch;
- auto kptr = gemm_silu_mul_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>;
- if (smem_size > 48'000) cudaFuncSetAttribute(kptr, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
+ auto kptr0 = gemm_silu_mul_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, 0, PAD, SWZ>;
+ auto kptr1 = gemm_silu_mul_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, 1, PAD, SWZ>;
+ auto kptr = prefer_cache_A ? kptr0 : kptr1;
+ if constexpr (smem_size > 48'000) {
+ static int attr0 = 0;
+ static int attr1 = 0;
+ if (prefer_cache_A) {
+ if (__builtin_expect(attr0 == 0, 0)) {
+ cudaFuncSetAttribute(kptr0, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
+ attr0 = 1;
+ }
⋯ diff truncated
scrolls · 1201 diff lines total

Best evidence level for this revision: reported

JSON