Skip to content
KernelIndex
Search⌘K

submission 401132

jon9509 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

refine.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-401132?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 group GEMMsuite of 4 cases
NVIDIA B200
32.7µs
#174 of 310
2026-01-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:761b92693b06da83749795ba08735b8de947c5a213e30dbda6398edb98e6f4e1
license declaredunknown
license concludedunknown
authorsjon9509
imported2026-08-15

Techniques

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

cluster"problem_shapes.get_host_problem_shape(), TileShape{}, AtomThrShapeMNK{}, ClusterShape{},\n args.hw_info, args.scheduler, scheduler_workspace",
fp4CUTLASS SM100 (B200) grouped GEMM for NVFP4 with blockwise scaling.
fused-epilogueusing EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;
warp-specialization"sm100_gemm_array_tma_warpspecialized.hpp",

Kernel source

refine.py1645 lines
import os
import time
import torch
import subprocess
from collections import OrderedDict
from torch.utils.cpp_extension import load_inline, include_paths


def custom_kernel(data):
    """
    CUTLASS SM100 (B200) grouped GEMM for NVFP4 with blockwise scaling.
    Based on example 81: blackwell_grouped_gemm_blockwise.cu pattern.
    """

    # -----------------------------
    # 0) Unpack input (robust)
    # -----------------------------
    if isinstance(data, (list, tuple)) and len(data) == 3:
        abc_tensors, sfasfb_tensors, problem_sizes = data
    elif isinstance(data, (list, tuple)) and len(data) == 4:
        abc_tensors, _, sfasfb_tensors, problem_sizes = data
    else:
        raise TypeError("Unexpected data format for custom_kernel")

    g = len(problem_sizes)
    if g == 0:
        return []

    # -----------------------------
    # 0.5) Optional Python cache (shape-keyed; ptr tensors updated in-place)
    # -----------------------------
    if not hasattr(custom_kernel, "_CACHE"):
        custom_kernel._CACHE = OrderedDict()

    def _tensor_key(t):
        return (
            tuple(int(x) for x in t.size()),
            tuple(int(x) for x in t.stride()),
            str(t.dtype),
            str(t.device),
        )

    def _build_cache_key(abc_tensors, sfasfb_tensors, problem_sizes):
        abc_sig = tuple((_tensor_key(a), _tensor_key(b), _tensor_key(c)) for a, b, c in abc_tensors)
        sfs_sig = tuple((_tensor_key(sfa), _tensor_key(sfb)) for sfa, sfb in sfasfb_tensors)
        ps_sig = tuple((int(m), int(n), int(k), int(l)) for (m, n, k, l) in problem_sizes)
        return (abc_sig, sfs_sig, ps_sig)

    def _bucket_id(problem_sizes, mod):
        # Deterministic FNV-1a over shapes
        h = 1469598103934665603
        for (m, n, k, l) in problem_sizes:
            for v in (m, n, k, l):
                h ^= int(v) & 0xFFFFFFFFFFFFFFFF
                h *= 1099511628211
                h &= 0xFFFFFFFFFFFFFFFF
        return int(h % mod)

    use_py_cache = int(os.environ.get("NVFP4_PY_CACHE", "1")) > 0
    use_echokey = int(os.environ.get("NVFP4_ECHOKEY", "0")) > 0
    echokey_debug = int(os.environ.get("NVFP4_ECHOKEY_DEBUG", "0")) > 0
    use_warmup = int(os.environ.get("NVFP4_ECHOKEY_WARMUP", "0")) > 0
    map_size = int(os.environ.get("NVFP4_ECHOKEY_MAP_SIZE", "16384"))
    map_size = max(1, map_size)
    cache_key = None
    cached_state = None

    if not hasattr(custom_kernel, "_ECHO_DUMMY"):
        custom_kernel._ECHO_DUMMY = (
            torch.empty(0, dtype=torch.int64, device=abc_tensors[0][2].device),
            torch.empty(0, dtype=torch.float32, device=abc_tensors[0][2].device),
        )
    if not hasattr(custom_kernel, "_POTENTIAL_MAP"):
        custom_kernel._POTENTIAL_MAP = {}
    if not hasattr(custom_kernel, "_SYNERGY_OUT"):
        custom_kernel._SYNERGY_OUT = {}

    # Defer cache lookup until after padding / L-expansion so ptrs are up to date.

    timing_py = int(os.environ.get("NVFP4_TIMING_PY", "0"))
    if timing_py:
        t_py0 = time.perf_counter()

    # -----------------------------
    # 1) Lazy-build extension
    # -----------------------------
    if not hasattr(custom_kernel, "_EXT"):

        cuda_root = os.environ.get("CUDA_HOME", "/usr/local/cuda")
        cuda_inc = os.path.join(cuda_root, "targets", "x86_64-linux", "include")
        cccl_inc = os.path.join(cuda_inc, "cccl")

        repo_path = "cutlass_repo"
        cutlass_inc = os.path.join(repo_path, "include")
        cutlass_util_inc = os.path.join(repo_path, "tools", "util", "include")
        header_check = os.path.join(cutlass_inc, "cutlass", "cutlass.h")

        if not os.path.exists(header_check):
            os.makedirs(repo_path, exist_ok=True)
            try:
                subprocess.run(
                    ["git", "clone", "--depth", "1", "https://github.com/NVIDIA/cutlass.git", repo_path],
                    check=True,
                    stdout=subprocess.PIPE,
                    stderr=subprocess.PIPE,
                )
            except subprocess.CalledProcessError as e:
                msg = e.stderr.decode(errors="ignore")
                raise RuntimeError(f"Could not clone CUTLASS. Details:\n{msg}")

        # Pin to the CUTLASS commit observed on the server.
        cutlass_commit = "2fafefb7b9c7e233c1cfb258f38139a6406d1511"
        try:
            subprocess.run(
                ["git", "-C", repo_path, "checkout", cutlass_commit],
                check=True,
                stdout=subprocess.PIPE,
                stderr=subprocess.PIPE,
            )
        except subprocess.CalledProcessError:
            subprocess.run(
                ["git", "-C", repo_path, "fetch", "--depth", "1", "origin", cutlass_commit],
                check=True,
                stdout=subprocess.PIPE,
                stderr=subprocess.PIPE,
            )
            subprocess.run(
                ["git", "-C", repo_path, "checkout", cutlass_commit],
                check=True,
                stdout=subprocess.PIPE,
                stderr=subprocess.PIPE,
            )

        # Patch the CUTLASS bug (apply even if repo already existed)
        kernel_file = os.path.join(
            repo_path,
            "include",
            "cutlass",
            "gemm",
            "kernel",
            "sm100_gemm_array_tma_warpspecialized.hpp",
        )
        if not os.path.exists(kernel_file):
            raise RuntimeError("CUTLASS kernel header not found to patch.")

        with open(kernel_file, "r") as f:
            content = f.read()

        patched = content.replace(
            "args.hw_info, scheduler_workspace",
            "args.hw_info, args.scheduler, scheduler_workspace",
        )
        patched = patched.replace(
            "problem_shapes.get_host_problem_shape(), TileShape{}, AtomThrShapeMNK{}, ClusterShape{},\n      args.hw_info, args.scheduler, scheduler_workspace",
            "problem_shapes.get_host_problem_shape(), TileShape{}, ClusterShape{},\n      args.hw_info, args.scheduler, scheduler_workspace",
        )
        patched = patched.replace(
            "grid_shape = TileScheduler::get_grid_shape(\n        params.scheduler,\n        params.problem_shape.get_host_problem_shape(),\n        TileShape{},\n        AtomThrShapeMNK{},\n        cluster_shape,\n        params.hw_info);",
            "grid_shape = TileScheduler::get_grid_shape(\n        params.scheduler,\n        params.problem_shape.get_host_problem_shape(),\n        TileShape{},\n        cluster_shape,\n        params.hw_info,\n        TileSchedulerArguments{});",
        )

        if patched != content:
            with open(kernel_file, "w") as f:
                f.write(patched)

        # Inject in-kernel freework inside the SM100 pipeline consumer wait (optional compile-time)
        pipeline_file = os.path.join(
            repo_path,
            "include",
            "cutlass",
            "pipeline",
            "sm100_pipeline.hpp",
        )
        if not os.path.exists(pipeline_file):
            raise RuntimeError("CUTLASS pipeline header not found to patch.")

        with open(pipeline_file, "r") as f:
            pcontent = f.read()

        old_wait = (
            "  void consumer_wait(uint32_t stage, uint32_t phase, ConsumerToken barrier_token) {\n"
            "    detail::pipeline_check_is_consumer(params_.role);\n"
            "    if (barrier_token == BarrierStatus::WaitAgain) {\n"
            "      full_barrier_ptr_[stage].wait(phase);\n"
            "    }\n"
            "  }\n"
        )
        new_wait = (
            "  void consumer_wait(uint32_t stage, uint32_t phase, ConsumerToken barrier_token) {\n"
            "    detail::pipeline_check_is_consumer(params_.role);\n"
            "    if (barrier_token == BarrierStatus::WaitAgain) {\n"
            "#if defined(NVFP4_INKERNEL_BURN) && (NVFP4_INKERNEL_BURN > 0)\n"
            "      auto do_freework = [&]() {\n"
            "        float acc = float(lane_idx_ + 1);\n"
            "        #pragma unroll 4\n"
            "        for (int i = 0; i < NVFP4_INKERNEL_BURN; ++i) {\n"
            "          acc = acc * 1.000001f + 0.000001f;\n"
            "        }\n"
            "        asm volatile(\"\" :: \"f\"(acc));\n"
            "      };\n"
            "      while (!full_barrier_ptr_[stage].try_wait(phase)) {\n"
            "        do_freework();\n"
            "      }\n"
            "#if defined(NVFP4_INKERNEL_BURN_MARKER)\n"
            "      asm volatile(\"/* NVFP4_INKERNEL_BURN_MARKER=%0 */\" :: \"n\"(NVFP4_INKERNEL_BURN_MARKER));\n"
            "#endif\n"
            "#else\n"
            "      full_barrier_ptr_[stage].wait(phase);\n"
            "#endif\n"
            "    }\n"
            "  }\n"
        )

        if old_wait in pcontent:
            pcontent = pcontent.replace(old_wait, new_wait)
            with open(pipeline_file, "w") as f:
                f.write(pcontent)

        if not os.path.exists(header_check):
            raise RuntimeError("CUTLASS headers not found after auto-heal.")

        # Include paths
        extra_includes = []
        for p in (cutlass_inc, cutlass_util_inc, cuda_inc, cccl_inc):
            if p and os.path.isdir(p):
                extra_includes.append(p)

        for p in list(dict.fromkeys(include_paths())):
            if os.path.isfile(os.path.join(p, "cutlass", "cutlass.h")):
                extra_includes.append(p)

        seen = set()
        extra_includes = [p for p in extra_includes if not (p in seen or seen.add(p))]

        cuda_src = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>

#include <vector>
#include <stdexcept>
#include <string>
#include <cstdio>
#include <cstdint>
#include <cstring>
#include <cstdlib>
#include <cmath>
#include <unordered_map>
#include <mutex>
#include <chrono>

// Core CUTLASS + CuTe
#include "cutlass/cutlass.h"
#include "cutlass/numeric_types.h"
#include "cutlass/numeric_conversion.h"
#include <cute/arch/config.hpp>
#if !defined(__CUDA_ARCH__)
#undef CUTE_ARCH_TMA_SM90_ENABLED
#undef CUTE_ARCH_TMA_SM100_ENABLED
#endif
#include "cute/tensor.hpp"

// Dispatch policies
#include "cutlass/gemm/dispatch_policy.hpp"
#include "cutlass/epilogue/dispatch_policy.hpp"

// GEMM infrastructure
#include "cutlass/gemm/group_array_problem_shape.hpp"
#include "cutlass/gemm/collective/collective_builder.hpp"
#include "cutlass/epilogue/collective/collective_builder.hpp"
#include "cutlass/gemm/device/gemm_universal_adapter.h"
#include "cutlass/gemm/kernel/gemm_universal.hpp"

// Packed stride helper
#include "cutlass/util/packed_stride.hpp"

using namespace cute;

static inline void check_cuda_tensor(const torch::Tensor& t, const char* name) {
  if (!t.defined()) throw std::runtime_error(std::string(name) + " is undefined");
  if (!t.is_cuda()) throw std::runtime_error(std::string(name) + " must be CUDA");
  if (!t.is_contiguous()) throw std::runtime_error(std::string(name) + " must be contiguous");
}

static inline void check_int32_tensor(const torch::Tensor& t, const char* name) {
  if (!t.defined()) throw std::runtime_error(std::string(name) + " is undefined");
  if (t.scalar_type() != torch::kInt32) throw std::runtime_error(std::string(name) + " must be int32");
  if (!t.is_contiguous()) throw std::runtime_error(std::string(name) + " must be contiguous");
  if (!(t.is_cuda() || t.is_cpu())) throw std::runtime_error(std::string(name) + " must be CPU or CUDA");
}

static inline void check_float_tensor_optional(const torch::Tensor& t, const char* name) {
  if (!t.defined()) return;
  if (!t.is_cuda()) throw std::runtime_error(std::string(name) + " must be CUDA");
  if (t.scalar_type() != torch::kFloat32) throw std::runtime_error(std::string(name) + " must be float32");
  if (!t.is_contiguous()) throw std::runtime_error(std::string(name) + " must be contiguous");
}

static inline void check_int64_tensor_optional(const torch::Tensor& t, const char* name) {
  if (!t.defined()) return;
  if (!t.is_cuda()) throw std::runtime_error(std::string(name) + " must be CUDA");
  if (t.scalar_type() != torch::kInt64) throw std::runtime_error(std::string(name) + " must be int64");
  if (!t.is_contiguous()) throw std::runtime_error(std::string(name) + " must be contiguous");
}

static inline std::string status_to_string(cutlass::Status st) {
  switch (st) {
    case cutlass::Status::kSuccess: return "kSuccess";
    case cutlass::Status::kErrorInvalidProblem: return "kErrorInvalidProblem";
    case cutlass::Status::kErrorNotSupported: return "kErrorNotSupported";
    default: return "cutlass::Status(" + std::to_string(int(st)) + ")";
  }
}

static inline int timing_mode() {
  static int mode = -1;
  if (mode == -1) {
    const char* env = std::getenv("NVFP4_TIMING");
    mode = env ? std::atoi(env) : 0;
  }
  return mode;
}

static inline int timing_abort() {
  static int mode = -1;
  if (mode == -1) {
    const char* env = std::getenv("NVFP4_TIMING_ABORT");
    mode = env ? std::atoi(env) : 0;
  }
  return mode;
}

static inline int timing_detail() {
  static int mode = -1;
  if (mode == -1) {
    const char* env = std::getenv("NVFP4_TIMING_DETAIL");
    mode = env ? std::atoi(env) : 0;
  }
  return mode;
}

#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)

// Block-scaled tensor core ops
using ArchTag       = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;

// Problem shape (L is not supported in this kernel; use 3D M,N,K)
using ProblemShape = cutlass::gemm::GroupProblemShape<Shape<int,int,int>>;
using UnderlyingProblemShape = typename ProblemShape::UnderlyingProblemShape;

// Element types
using ElementAccumulator = float;
using ElementCompute     = float;

// FROZEN: Element types below are validated. Do not change unless explicitly requested.
// Touching these will loop back on prior fixes.
using ElementInput = cutlass::float_e2m1_t;
using ElementA     = cutlass::nv_float4_t<ElementInput>;
using ElementB     = cutlass::nv_float4_t<ElementInput>;
using ElementC     = cutlass::half_t;
using ElementD     = ElementC;

// Layout tags
using LayoutA = cutlass::layout::RowMajor*;
using LayoutB = cutlass::layout::ColumnMajor*;
using LayoutC = cutlass::layout::RowMajor;
using LayoutD = cutlass::layout::RowMajor;

// Alignments
constexpr int AlignmentA = 32;
constexpr int AlignmentB = 32;
constexpr int AlignmentC = 8;
constexpr int AlignmentD = 8;

// Tile + cluster shape
using MmaTileShape = Shape<_128,_128,_128>;
using ClusterShape = Shape<_1,_1,_1>;


// FROZEN: Schedule tags below are validated. Do not change unless explicitly requested.
// Touching these will loop back on prior fixes.
using KernelSchedule   = cutlass::gemm::KernelPtrArrayTmaWarpSpecialized1SmNvf4Sm100;
using EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;

// Epilogue tile
using EpilogueTileShape = Shape<_128,_64>;

// Build epilogue first
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
  ArchTag, OperatorClass,
  MmaTileShape, ClusterShape,
  EpilogueTileShape,
  ElementAccumulator, ElementCompute,
  ElementC, LayoutC*, AlignmentC,
  ElementD, LayoutD*, AlignmentD,
  EpilogueSchedule
>::CollectiveOp;

// Build mainloop with element-pair tuple for block-scaled GEMM
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
  ArchTag, OperatorClass,
  ElementA, LayoutA, AlignmentA,
  ElementB, LayoutB, AlignmentB,
  ElementAccumulator,
  MmaTileShape, ClusterShape,
  cutlass::gemm::collective::StageCountAutoCarveout<
    static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))
  >,
  KernelSchedule
>::CollectiveOp;

// Get SFVecSize from the mainloop (this is the part that must match the kernel)
static constexpr int SFVecSize = int(CollectiveMainloop::SFVecSize);
using Sm1xxBlkScaledConfig = cutlass::detail::Sm1xxBlockScaledConfig<SFVecSize>;

// ScaleConfig + layouts (match mainloop expected types)
using LayoutSFA = typename cutlass::detail::LayoutSFAType<CollectiveMainloop>::type;
using LayoutSFB = typename cutlass::detail::LayoutSFBType<CollectiveMainloop>::type;
using InternalLayoutSFA = cute::remove_pointer_t<LayoutSFA>;
using InternalLayoutSFB = cute::remove_pointer_t<LayoutSFB>;
using ElementSF = typename cutlass::detail::ElementSFType<CollectiveMainloop>::type;

// Kernel
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
  ProblemShape,
  CollectiveMainloop,
  CollectiveEpilogue
>;

using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;

// Extract strides
using StrideA = typename CollectiveMainloop::StrideA;
using InternalStrideA = typename CollectiveMainloop::InternalStrideA;
using StrideB = typename CollectiveMainloop::StrideB;
using InternalStrideB = typename CollectiveMainloop::InternalStrideB;
using StrideC = typename Gemm::GemmKernel::InternalStrideC;
using StrideD = typename Gemm::GemmKernel::InternalStrideD;

struct CacheEntry {
  at::Tensor dev_buffer;
  at::Tensor workspace;
  size_t total_bytes = 0;
  size_t off_sA = 0;
  size_t off_sB = 0;
  size_t off_sC = 0;
  size_t off_sD = 0;
  size_t off_lSFA = 0;
  size_t off_lSFB = 0;
  size_t off_ps = 0;
  int device = -1;
  int64_t g = 0;
  std::vector<UnderlyingProblemShape> ps_host_shape;
};

static std::unordered_map<std::string, CacheEntry> g_cache;
static std::mutex g_cache_mutex;
static std::mutex g_last_mutex;
static int g_last_device = -1;
static int64_t g_last_g = 0;
static std::vector<int32_t> g_last_ps_host;
static int64_t g_last_ps_ptr = 0;
static std::string g_last_key;

static std::unordered_map<int, cutlass::KernelHardwareInfo> g_hw_cache;
static std::mutex g_hw_mutex;

static inline std::string make_cache_key(int device, int64_t g, const std::vector<int32_t>& ps_host) {
  std::string key;
  key.reserve(size_t(16 + ps_host.size() * 6));
  key += std::to_string(device);
  key.push_back(':');
  key += std::to_string((long long)g);
  for (int32_t v : ps_host) {
    key.push_back(',');
    key += std::to_string(int(v));
  }
  return key;
}

static inline std::string make_cache_key_ptr(int device, int64_t g, int64_t ps_ptr, int64_t nvals) {
  std::string key;
  key.reserve(64);
  key += std::to_string(device);
  key.push_back(':');
  key += std::to_string((long long)g);
  key += ":ptr=" + std::to_string((long long)ps_ptr);
  key += ":n=" + std::to_string((long long)nvals);
  return key;
}



__global__ void sf_float_to_ue4m3_kernel(const float* in, ElementSF* out, int64_t n) {
  int64_t idx = int64_t(blockIdx.x) * blockDim.x + threadIdx.x;
  if (idx < n) {
    cutlass::NumericConverter<ElementSF, float, cutlass::FloatRoundStyle::round_to_nearest> cvt;
    out[idx] = cvt(in[idx]);
  }
}

__global__ void pad_copy_fp4_kernel(
    uint8_t* dst,
    const uint8_t* src,
    int64_t total) {
  int64_t idx = int64_t(blockIdx.x) * blockDim.x + threadIdx.x;
  if (idx < total) {
    dst[idx] = src[idx];
  }
}

__global__ void pad_copy_c_half_kernel(
    cutlass::half_t* dst,
    const cutlass::half_t* src,
    int64_t m,
    int64_t n,
    int64_t dst_stride_m,
    int64_t src_stride_m) {
  int64_t idx = int64_t(blockIdx.x) * blockDim.x + threadIdx.x;
  int64_t total = m * n;
  if (idx < total) {
    int64_t row = idx / n;
    int64_t col = idx - row * n;
    dst[row * dst_stride_m + col] = src[row * src_stride_m + col];
  }
}

__global__ void synergy_prologue_kernel(int64_t* timestamps) {
  if (threadIdx.x == 0) {
    uint64_t t;
    asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(t));
    timestamps[0] = (int64_t)t;
  }
}

__global__ void synergy_epilogue_kernel(
    int64_t* timestamps,
    float* potential_map,
    int64_t map_size,
    int64_t bucket_id,
    int64_t total_work_items,
    int burn) {
  float acc0 = 0.0f;
  if (burn > 0) {
    float acc = float(threadIdx.x + 1);
    #pragma unroll 4
    for (int i = 0; i < burn; ++i) {
      acc = acc * 1.000001f + 0.000001f;
    }
    if (threadIdx.x == 0) {
      acc0 = acc;
      if (timestamps) {
        timestamps[0] ^= (int64_t)(acc);
      }
    }
  }

  if (threadIdx.x == 0) {
    uint64_t t1;
    asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(t1));
    int64_t t0 = timestamps[0];
    int64_t dt = (int64_t)t1 - t0;

    float resonance = 0.0f;
    int64_t cascade = 0;
    if (potential_map && map_size > 0 && total_work_items > 0) {
      int64_t idx = bucket_id % map_size;
      if (idx < 0) idx += map_size;
      float old_v = potential_map[idx];
      float time_per_item = float(dt) / float(total_work_items);
      float base_v = (old_v == 0.0f) ? time_per_item : old_v;
      float expected = base_v * float(total_work_items);
      float tau_res = 100000.0f;
      float dilation = float(dt) - expected;
      resonance = tanhf(dilation / tau_res);
      if (fabsf(resonance) > 0.95f) {
        cascade = 1;
      }
      float alpha = 0.002f;
      float new_v = (old_v == 0.0f) ? time_per_item : ((1.0f - alpha) * old_v + alpha * time_per_item);
      potential_map[idx] = new_v;
    }

    timestamps[1] = (int64_t)t1;
    timestamps[2] = dt;
    timestamps[3] = (int64_t)(resonance * 1000000.0f);
    timestamps[4] = cascade;
    if (burn > 0) {
      timestamps[4] += burn;
      if (potential_map && map_size > 0) {
        potential_map[0] += acc0 * 1.0e-7f;
      }
    }
  }
}

torch::Tensor convert_sf_float_to_ue4m3(torch::Tensor in) {
  check_cuda_tensor(in, "sf_in");
  if (in.scalar_type() != torch::kFloat32) {
    throw std::runtime_error("convert_sf_float_to_ue4m3 expects float32 input");
  }

  auto out = torch::empty_like(in, torch::TensorOptions().dtype(torch::kUInt8).device(in.device()));
  int64_t n = in.numel();
  if (n == 0) return out;

  int threads = 256;
  int blocks = int((n + threads - 1) / threads);
  sf_float_to_ue4m3_kernel<<<blocks, threads, 0, 0>>>(
      in.data_ptr<float>(),
      reinterpret_cast<ElementSF*>(out.data_ptr<uint8_t>()),
      n);
  return out;
}

torch::Tensor pad_a_fp4(torch::Tensor a, int64_t m_pad) {
  check_cuda_tensor(a, "a");
  auto sizes = a.sizes();
  if (sizes.size() != 3) {
    throw std::runtime_error("pad_a_fp4 expects a 3D tensor");
  }
  int64_t m = sizes[0];
  int64_t k2 = sizes[1];
  int64_t l = sizes[2];
  if (m_pad < m) {
    throw std::runtime_error("pad_a_fp4 expects m_pad >= m");
  }
  auto out = torch::empty({m_pad, k2, l}, a.options());
  int64_t total = m * k2 * l;
  int64_t total_bytes = int64_t(out.numel()) * int64_t(sizeof(ElementA));
  int64_t valid_bytes = total * int64_t(sizeof(ElementA));
  if (valid_bytes < total_bytes) {
    uint8_t* base = reinterpret_cast<uint8_t*>(out.data_ptr());
    cudaMemset(base + valid_bytes, 0, size_t(total_bytes - valid_bytes));
  }
  if (total == 0) return out;
  int threads = 256;
  int blocks = int((total + threads - 1) / threads);
  pad_copy_fp4_kernel<<<blocks, threads, 0, 0>>>(
      reinterpret_cast<uint8_t*>(out.data_ptr()),
      reinterpret_cast<const uint8_t*>(a.data_ptr()),
      total);
  return out;
}

torch::Tensor pad_c_half(torch::Tensor c, int64_t m_pad, int64_t n_pad) {
  check_cuda_tensor(c, "c");
  auto sizes = c.sizes();
  if (sizes.size() != 3) {
    throw std::runtime_error("pad_c_half expects a 3D tensor");
  }
  int64_t m = sizes[0];
  int64_t n = sizes[1];
  int64_t l = sizes[2];
  if (l != 1) {
    throw std::runtime_error("pad_c_half expects L=1");
  }
  if (m_pad < m || n_pad < n) {
    throw std::runtime_error("pad_c_half expects m_pad >= m and n_pad >= n");
  }
  auto out = torch::empty({m_pad, n_pad, l}, c.options());
  int64_t total_elems = out.numel();
  if (total_elems > 0) {
    cudaMemset(out.data_ptr(), 0, size_t(total_elems) * sizeof(ElementC));
  }
  int64_t total = m * n;
  if (total == 0) return out;
  int threads = 256;
  int blocks = int((total + threads - 1) / threads);
  pad_copy_c_half_kernel<<<blocks, threads, 0, 0>>>(
      reinterpret_cast<cutlass::half_t*>(out.data_ptr()),
      reinterpret_cast<const cutlass::half_t*>(c.data_ptr()),
      m,
      n,
      n_pad,
      n);
  return out;
}

void grouped_nvfp4_tma_gemm(
    torch::Tensor problem_sizes,  // int32 [g,3]
    torch::Tensor ptr_A,          // int64 [g]
    torch::Tensor ptr_B,          // int64 [g]
    torch::Tensor ptr_SFA,        // int64 [g]
    torch::Tensor ptr_SFB,        // int64 [g]
    torch::Tensor ptr_C,          // int64 [g]
    torch::Tensor ptr_D,          // int64 [g]
    torch::Tensor synergy_out,    // int64 [>=5] or empty
    torch::Tensor potential_map,  // float32 [map_size] or empty
    int64_t bucket_id,            // map index (unused for now)
    int64_t total_work_items      // scalar (unused for now)
) {
  check_int32_tensor(problem_sizes, "problem_sizes");
  check_cuda_tensor(ptr_A, "ptr_A");
  check_cuda_tensor(ptr_B, "ptr_B");
  check_cuda_tensor(ptr_SFA, "ptr_SFA");
  check_cuda_tensor(ptr_SFB, "ptr_SFB");
  check_cuda_tensor(ptr_C, "ptr_C");
  check_cuda_tensor(ptr_D, "ptr_D");
  check_int64_tensor_optional(synergy_out, "synergy_out");
  check_float_tensor_optional(potential_map, "potential_map");

  int64_t g = problem_sizes.size(0);
  if (g <= 0) return;

  if (synergy_out.defined() && synergy_out.numel() > 0 && synergy_out.numel() < 5) {
    throw std::runtime_error("synergy_out must have at least 5 elements");
  }
  if (synergy_out.defined() && synergy_out.numel() > 0) {
    if (!potential_map.defined() || potential_map.numel() == 0) {
      throw std::runtime_error("potential_map must be provided when synergy_out is set");
    }
  }

  int tmode = timing_mode();
  int tdetail = (tmode > 0) ? timing_detail() : 0;
  auto t0 = std::chrono::high_resolution_clock::now();
  auto t_ps0 = t0;
  auto t_ps1 = t0;
  auto t_cache0 = t0;
  auto t_cache1 = t0;
  auto t_build0 = t0;
  auto t_build1 = t0;
  auto t_upload0 = t0;
  auto t_upload1 = t0;
  auto t_can0 = t0;
  auto t_can1 = t0;
  auto t_ws0 = t0;
  auto t_ws1 = t0;
  auto t_sync0 = t0;
  auto t_sync1 = t0;

  // problem_sizes: int32 [g,3] (CPU or CUDA)
  bool ps_on_cuda = problem_sizes.is_cuda();
  std::vector<int32_t> ps_host;
  auto copy_ps_host = [&]() {
    if (!ps_host.empty()) return;
    ps_host.resize(size_t(g) * 3);
    if (ps_on_cuda) {
      cudaMemcpy(ps_host.data(),
                 problem_sizes.data_ptr<int32_t>(),
                 sizeof(int32_t) * size_t(g) * 3,
                 cudaMemcpyDeviceToHost);
    } else {
      std::memcpy(ps_host.data(),
                  problem_sizes.data_ptr<int32_t>(),
                  sizeof(int32_t) * size_t(g) * 3);
    }
  };
  if (tmode > 0) t_ps0 = std::chrono::high_resolution_clock::now();
  if (!ps_on_cuda) {
    copy_ps_host();
  }
  if (tmode > 0) t_ps1 = std::chrono::high_resolution_clock::now();

  // Build per-group strides/layouts/problem-shapes into a single aligned host buffer
  auto bytes = torch::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA);

  constexpr size_t kAlign = 128;
  auto align_up = [](size_t n, size_t a) { return (n + a - 1) & ~(a - 1); };
  auto max_align = [](size_t a, size_t b) { return a > b ? a : b; };

  size_t sz_sA   = size_t(g) * sizeof(InternalStrideA);
  size_t sz_sB   = size_t(g) * sizeof(InternalStrideB);
  size_t sz_sC   = size_t(g) * sizeof(StrideC);
  size_t sz_sD   = size_t(g) * sizeof(StrideD);
  size_t sz_lSFA = size_t(g) * sizeof(InternalLayoutSFA);
  size_t sz_lSFB = size_t(g) * sizeof(InternalLayoutSFB);
  size_t sz_ps   = size_t(g) * sizeof(UnderlyingProblemShape);

  size_t base_align = kAlign;
  base_align = max_align(base_align, alignof(InternalStrideA));
  base_align = max_align(base_align, alignof(InternalStrideB));
  base_align = max_align(base_align, alignof(StrideC));
  base_align = max_align(base_align, alignof(StrideD));
  base_align = max_align(base_align, alignof(InternalLayoutSFA));
  base_align = max_align(base_align, alignof(InternalLayoutSFB));
  base_align = max_align(base_align, alignof(UnderlyingProblemShape));

  int device = -1;
  cudaGetDevice(&device);
  cutlass::KernelHardwareInfo hw_info;
  {
    std::lock_guard<std::mutex> lock(g_hw_mutex);
    auto it = g_hw_cache.find(device);
    if (it != g_hw_cache.end()) {
      hw_info = it->second;
    } else {
      cudaDeviceProp prop{};
      cudaGetDeviceProperties(&prop, device);
      hw_info.device_id = device;
      hw_info.sm_count  = prop.multiProcessorCount;
      g_hw_cache.emplace(device, hw_info);
    }
  }

  typename GemmKernel::TileScheduler::Arguments scheduler_args{};

  if (tmode > 0) t_cache0 = std::chrono::high_resolution_clock::now();
  CacheEntry* entry = nullptr;
  std::string key;
  bool last_hit = false;
  int64_t ps_ptr_val = ps_on_cuda ? (int64_t)problem_sizes.data_ptr<int32_t>() : 0;
  {
    std::lock_guard<std::mutex> lock(g_last_mutex);
    if (g_last_device == device && g_last_g == g) {
      if (ps_on_cuda) {
        if (g_last_ps_ptr == ps_ptr_val) {
          key = g_last_key;
          last_hit = true;
        }
      } else {
        if (g_last_ps_host == ps_host) {
          key = g_last_key;
          last_hit = true;
        }
      }
    }
  }
  if (last_hit) {
    std::lock_guard<std::mutex> lock(g_cache_mutex);
    auto it = g_cache.find(key);
    if (it != g_cache.end()) {
      entry = &it->second;
    } else {
      last_hit = false;
    }
  }
  if (!last_hit) {
    if (ps_on_cuda) {
      key = make_cache_key_ptr(device, g, ps_ptr_val, g * 3);
    } else {
      key = make_cache_key(device, g, ps_host);
    }
    std::lock_guard<std::mutex> lock(g_cache_mutex);
    auto it = g_cache.find(key);
    if (it != g_cache.end()) {
      entry = &it->second;
    }
  }
  if (tmode > 0) t_cache1 = std::chrono::high_resolution_clock::now();

  std::vector<uint8_t> host_buffer;
  uint8_t* h_ptr = nullptr;
  UnderlyingProblemShape* h_ps = nullptr;
  if (entry == nullptr) {
    if (ps_on_cuda && ps_host.empty()) {
      if (tmode > 0) t_ps0 = std::chrono::high_resolution_clock::now();
      copy_ps_host();
      if (tmode > 0) t_ps1 = std::chrono::high_resolution_clock::now();
    }
    if (tmode > 0) t_build0 = std::chrono::high_resolution_clock::now();
    size_t off = 0;
    auto alloc = [&](size_t sz, size_t align) {
      off = align_up(off, align);
      size_t cur = off;
      off += sz;
      return cur;
    };

    size_t off_sA   = alloc(sz_sA,   max_align(kAlign, alignof(InternalStrideA)));
    size_t off_sB   = alloc(sz_sB,   max_align(kAlign, alignof(InternalStrideB)));
    size_t off_sC   = alloc(sz_sC,   max_align(kAlign, alignof(StrideC)));
    size_t off_sD   = alloc(sz_sD,   max_align(kAlign, alignof(StrideD)));
    size_t off_lSFA = alloc(sz_lSFA, max_align(kAlign, alignof(InternalLayoutSFA)));
    size_t off_lSFB = alloc(sz_lSFB, max_align(kAlign, alignof(InternalLayoutSFB)));
    size_t off_ps   = alloc(sz_ps,   max_align(kAlign, alignof(UnderlyingProblemShape)));
    size_t total_bytes = align_up(off, kAlign);

    host_buffer.resize(total_bytes + base_align);
    uintptr_t base = reinterpret_cast<uintptr_t>(host_buffer.data());
    uintptr_t aligned_base = (base + (base_align - 1)) & ~(uintptr_t(base_align - 1));
    h_ptr = reinterpret_cast<uint8_t*>(aligned_base);

    auto* h_sA   = reinterpret_cast<InternalStrideA*>(h_ptr + off_sA);
    auto* h_sB   = reinterpret_cast<InternalStrideB*>(h_ptr + off_sB);
    auto* h_sC   = reinterpret_cast<StrideC*>(h_ptr + off_sC);
    auto* h_sD   = reinterpret_cast<StrideD*>(h_ptr + off_sD);
    auto* h_lSFA = reinterpret_cast<InternalLayoutSFA*>(h_ptr + off_lSFA);
    auto* h_lSFB = reinterpret_cast<InternalLayoutSFB*>(h_ptr + off_lSFB);
    h_ps         = reinterpret_cast<UnderlyingProblemShape*>(h_ptr + off_ps);

    for (int64_t i = 0; i < g; ++i) {
      int M = ps_host[size_t(i)*3 + 0];
      int N = ps_host[size_t(i)*3 + 1];
      int K = ps_host[size_t(i)*3 + 2];

      // A/B are stored packed as [M, K/2, 1] and [N, K/2, 1] for NVFP4.
      // Use K-major packed strides to match reference layout: (m,k,1) stride (k,1,m*k).
      h_sA[i] = cutlass::make_cute_packed_stride(InternalStrideA{}, cute::make_shape(M, K, 1));
      h_sB[i] = cutlass::make_cute_packed_stride(InternalStrideB{}, cute::make_shape(N, K, 1));
      h_sC[i] = cutlass::make_cute_packed_stride(StrideC{}, cute::make_shape(M, N, 1));
      h_sD[i] = cutlass::make_cute_packed_stride(StrideD{}, cute::make_shape(M, N, 1));

      h_lSFA[i] = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(M, N, K, 1));
      h_lSFB[i] = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(M, N, K, 1));
      h_ps[i]   = UnderlyingProblemShape{M, N, K};
    }
    if (tmode > 0) t_build1 = std::chrono::high_resolution_clock::now();

    CacheEntry new_entry;
    new_entry.total_bytes = total_bytes;
    new_entry.off_sA = off_sA;
    new_entry.off_sB = off_sB;
    new_entry.off_sC = off_sC;
    new_entry.off_sD = off_sD;
    new_entry.off_lSFA = off_lSFA;
    new_entry.off_lSFB = off_lSFB;
    new_entry.off_ps = off_ps;
    new_entry.device = device;
    new_entry.g = g;
    new_entry.ps_host_shape.resize(size_t(g));
    for (int64_t i = 0; i < g; ++i) {
      int M = ps_host[size_t(i)*3 + 0];
      int N = ps_host[size_t(i)*3 + 1];
      int K = ps_host[size_t(i)*3 + 2];
      new_entry.ps_host_shape[size_t(i)] = UnderlyingProblemShape{M, N, K};
    }

    // One device buffer for strides/layouts/problem shapes
    new_entry.dev_buffer = torch::empty({(long long)total_bytes}, bytes);
    if (tmode > 0) t_upload0 = std::chrono::high_resolution_clock::now();
    cudaMemcpy(new_entry.dev_buffer.data_ptr(), h_ptr, total_bytes, cudaMemcpyHostToDevice);
    if (tmode > 0) t_upload1 = std::chrono::high_resolution_clock::now();

    {
      std::lock_guard<std::mutex> lock(g_cache_mutex);
      auto it = g_cache.emplace(key, std::move(new_entry)).first;
      entry = &it->second;
    }
    {
      std::lock_guard<std::mutex> lock(g_last_mutex);
      g_last_device = device;
      g_last_g = g;
      if (ps_on_cuda) {
        g_last_ps_ptr = ps_ptr_val;
        g_last_ps_host.clear();
      } else {
        g_last_ps_ptr = 0;
        g_last_ps_host = ps_host;
      }
      g_last_key = key;
    }
  }

  auto d_ptr = reinterpret_cast<uint8_t*>(entry->dev_buffer.data_ptr<uint8_t>());
  auto ps_ptr = reinterpret_cast<UnderlyingProblemShape*>(d_ptr + entry->off_ps);
  auto ps_host_ptr = ps_on_cuda ? nullptr : entry->ps_host_shape.data();

  auto pA   = reinterpret_cast<typename Gemm::ElementA const**>(ptr_A.data_ptr<int64_t>());
  auto pB   = reinterpret_cast<typename Gemm::ElementB const**>(ptr_B.data_ptr<int64_t>());
  auto pSFA = reinterpret_cast<ElementSF const**>(ptr_SFA.data_ptr<int64_t>());
  auto pSFB = reinterpret_cast<ElementSF const**>(ptr_SFB.data_ptr<int64_t>());
  auto pC   = reinterpret_cast<typename Gemm::ElementC const**>(ptr_C.data_ptr<int64_t>());
  auto pD   = reinterpret_cast<typename Gemm::EpilogueOutputOp::ElementOutput**>(ptr_D.data_ptr<int64_t>());

  auto sA   = reinterpret_cast<InternalStrideA*>(d_ptr + entry->off_sA);
  auto sB   = reinterpret_cast<InternalStrideB*>(d_ptr + entry->off_sB);
  auto sC   = reinterpret_cast<StrideC*>(d_ptr + entry->off_sC);
  auto sD   = reinterpret_cast<StrideD*>(d_ptr + entry->off_sD);
  auto lSFA = reinterpret_cast<InternalLayoutSFA*>(d_ptr + entry->off_lSFA);
  auto lSFB = reinterpret_cast<InternalLayoutSFB*>(d_ptr + entry->off_lSFB);

  typename Gemm::Arguments args{
    cutlass::gemm::GemmUniversalMode::kGrouped,
    { (int)g, ps_ptr, ps_host_ptr },
    { pA, StrideA(sA), pB, StrideB(sB), pSFA, LayoutSFA(lSFA), pSFB, LayoutSFB(lSFB) },
    { {}, pC, sC, pD, sD },

    hw_info,
    scheduler_args
  };
  args.epilogue.thread.alpha = 1.0f;
  args.epilogue.thread.beta = 0.0f;

  Gemm gemm;
  std::vector<InternalStrideA> h_sA_dbg;
  std::vector<InternalStrideB> h_sB_dbg;
  std::vector<StrideC> h_sC_dbg;
  std::vector<StrideD> h_sD_dbg;
  std::vector<InternalLayoutSFA> h_lSFA_dbg;
  std::vector<InternalLayoutSFB> h_lSFB_dbg;
  bool have_debug = false;
  auto build_debug = [&]() {
    if (have_debug) return;
    have_debug = true;
    if (ps_host.empty()) {
      copy_ps_host();
    }
    if (ps_host.empty()) {
      return;
    }
    h_sA_dbg.resize(size_t(g));
    h_sB_dbg.resize(size_t(g));
    h_sC_dbg.resize(size_t(g));
    h_sD_dbg.resize(size_t(g));
    h_lSFA_dbg.resize(size_t(g));
    h_lSFB_dbg.resize(size_t(g));
    for (int64_t i = 0; i < g; ++i) {
      int M = ps_host[size_t(i)*3 + 0];
      int N = ps_host[size_t(i)*3 + 1];
      int K = ps_host[size_t(i)*3 + 2];
      h_sA_dbg[size_t(i)] = cutlass::make_cute_packed_stride(InternalStrideA{}, cute::make_shape(M, K, 1));
      h_sB_dbg[size_t(i)] = cutlass::make_cute_packed_stride(InternalStrideB{}, cute::make_shape(N, K, 1));
      h_sC_dbg[size_t(i)] = cutlass::make_cute_packed_stride(StrideC{}, cute::make_shape(M, N, 1));
      h_sD_dbg[size_t(i)] = cutlass::make_cute_packed_stride(StrideD{}, cute::make_shape(M, N, 1));
      h_lSFA_dbg[size_t(i)] = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(M, N, K, 1));
      h_lSFB_dbg[size_t(i)] = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(M, N, K, 1));
    }
  };

  if (tmode > 0) t_can0 = std::chrono::high_resolution_clock::now();
  cutlass::Status st = gemm.can_implement(args);
  cutlass::Status st_nohost = st;
  cutlass::Status st_host = st;
  if (ps_on_cuda) {
    auto args_nohost = args;
    args_nohost.problem_shape = { (int)g, ps_ptr, nullptr };
    st_nohost = gemm.can_implement(args_nohost);
    if (st_nohost == cutlass::Status::kSuccess) {
      args = args_nohost;
      st = st_nohost;
    } else if (!entry->ps_host_shape.empty()) {
      auto args_host = args;
      args_host.problem_shape = { (int)g, ps_ptr, entry->ps_host_shape.data() };
      st_host = gemm.can_implement(args_host);
      if (st_host == cutlass::Status::kSuccess) {
        args = args_host;
        st = st_host;
      }
    }
  }
  if (tmode > 0) t_can1 = std::chrono::high_resolution_clock::now();
  if (st != cutlass::Status::kSuccess) {
    // Probe: can_implement with host problem shape pointer cleared (matches reference host_problem_shapes_available=false).
    auto args_nohost = args;
    args_nohost.problem_shape = { (int)g, ps_ptr, nullptr };
    st_nohost = gemm.can_implement(args_nohost);
    std::string msg = "can_implement failed: " + status_to_string(st);
    msg += " (" + std::string(cutlassGetStatusString(st)) + ")";
    // Always include minimal shape/type info to aid debugging without 
    msg += " | g=" + std::to_string((long long)g);
    msg += " SFVecSize=" + std::to_string(SFVecSize);
    msg += " ElementSF_size=" + std::to_string(sizeof(ElementSF));
    msg += " ElementA_size=" + std::to_string(sizeof(ElementA));
    msg += " ElementB_size=" + std::to_string(sizeof(ElementB));
    msg += " ElementC_size=" + std::to_string(sizeof(ElementC));
    if (g > 0) {
      if (ps_host.empty()) {
        copy_ps_host();
      }
      if (!ps_host.empty()) {
        int M = ps_host[0];
        int N = ps_host[1];
        int K = ps_host[2];
        msg += " | group0 M=" + std::to_string(M)
            + " N=" + std::to_string(N)
            + " K=" + std::to_string(K);
        msg += " K_div2=" + std::to_string(K / 2);
      } else {
        msg += " | group0_unavailable";
      }
    }

    build_debug();

    // Per-group isolation probe to find the first failing group (attach to thrown msg).
    std::vector<int64_t> ptrA_host;
    std::vector<int64_t> ptrB_host;
    std::vector<int64_t> ptrSFA_host;
    std::vector<int64_t> ptrSFB_host;
    std::vector<int64_t> ptrC_host;
    std::vector<int64_t> ptrD_host;
    ptrA_host.resize(size_t(g));
    ptrB_host.resize(size_t(g));
    ptrSFA_host.resize(size_t(g));
    ptrSFB_host.resize(size_t(g));
    ptrC_host.resize(size_t(g));
    ptrD_host.resize(size_t(g));
    cudaMemcpy(ptrA_host.data(),  ptr_A.data_ptr<int64_t>(),  sizeof(int64_t) * size_t(g), cudaMemcpyDeviceToHost);
    cudaMemcpy(ptrB_host.data(),  ptr_B.data_ptr<int64_t>(),  sizeof(int64_t) * size_t(g), cudaMemcpyDeviceToHost);
    cudaMemcpy(ptrSFA_host.data(), ptr_SFA.data_ptr<int64_t>(), sizeof(int64_t) * size_t(g), cudaMemcpyDeviceToHost);
    cudaMemcpy(ptrSFB_host.data(), ptr_SFB.data_ptr<int64_t>(), sizeof(int64_t) * size_t(g), cudaMemcpyDeviceToHost);
    cudaMemcpy(ptrC_host.data(),  ptr_C.data_ptr<int64_t>(),  sizeof(int64_t) * size_t(g), cudaMemcpyDeviceToHost);
    cudaMemcpy(ptrD_host.data(),  ptr_D.data_ptr<int64_t>(),  sizeof(int64_t) * size_t(g), cudaMemcpyDeviceToHost);

    if (ps_host.empty()) {
      copy_ps_host();
    }
    if (ps_host.empty()) {
      msg += " | ps_host_unavailable";
      throw std::runtime_error(msg);
    }

    bool found = false;
    for (int64_t i = 0; i < g; ++i) {
      auto ps_ptr_i = ps_ptr + i;
      auto ps_host_ptr_i = ps_host_ptr ? (ps_host_ptr + i) : nullptr;
      auto pA_i   = pA + i;
      auto pB_i   = pB + i;
      auto pSFA_i = pSFA + i;
      auto pSFB_i = pSFB + i;
      auto pC_i   = pC + i;
      auto pD_i   = pD + i;
      auto sA_i   = sA + i;
      auto sB_i   = sB + i;
      auto sC_i   = sC + i;
      auto sD_i   = sD + i;
      auto lSFA_i = lSFA + i;
      auto lSFB_i = lSFB + i;

      typename Gemm::Arguments args1{
        cutlass::gemm::GemmUniversalMode::kGrouped,
        { 1, ps_ptr_i, ps_host_ptr_i },
        { pA_i, StrideA(sA_i), pB_i, StrideB(sB_i), pSFA_i, LayoutSFA(lSFA_i), pSFB_i, LayoutSFB(lSFB_i) },
        { {}, pC_i, sC_i, pD_i, sD_i },
        hw_info,
        scheduler_args
      };
      args1.epilogue.thread.alpha = 1.0f;
      args1.epilogue.thread.beta = 0.0f;

      auto st1 = gemm.can_implement(args1);
      if (st1 != cutlass::Status::kSuccess) {
        int M = ps_host[size_t(i)*3 + 0];
        int N = ps_host[size_t(i)*3 + 1];
        int K = ps_host[size_t(i)*3 + 2];
        msg += " | first_bad_group=" + std::to_string((long long)i);
        msg += " M=" + std::to_string(M);
        msg += " N=" + std::to_string(N);
        msg += " K=" + std::to_string(K);
        msg += " Mmod128=" + std::to_string(M % 128);
        msg += " Nmod128=" + std::to_string(N % 128);
        msg += " Kmod128=" + std::to_string(K % 128);
        int64_t reqA = int64_t(AlignmentA) * int64_t(sizeof(ElementA));
        int64_t reqB = int64_t(AlignmentB) * int64_t(sizeof(ElementB));
        int64_t reqC = int64_t(AlignmentC) * int64_t(sizeof(ElementC));
        int64_t reqD = int64_t(AlignmentD) * int64_t(sizeof(ElementD));
        msg += " A_ptr_mod16=" + std::to_string((long long)(ptrA_host[size_t(i)] % 16));
        msg += " B_ptr_mod16=" + std::to_string((long long)(ptrB_host[size_t(i)] % 16));
        msg += " SFA_ptr_mod16=" + std::to_string((long long)(ptrSFA_host[size_t(i)] % 16));
        msg += " SFB_ptr_mod16=" + std::to_string((long long)(ptrSFB_host[size_t(i)] % 16));
        msg += " | AlignA_bytes=" + std::to_string((long long)reqA)
            + " A_ptr_mod_req=" + std::to_string((long long)(ptrA_host[size_t(i)] % reqA));
        msg += " AlignB_bytes=" + std::to_string((long long)reqB)
            + " B_ptr_mod_req=" + std::to_string((long long)(ptrB_host[size_t(i)] % reqB));
        msg += " AlignC_bytes=" + std::to_string((long long)reqC)
            + " C_ptr_mod_req=" + std::to_string((long long)(ptrC_host[size_t(i)] % reqC));
        msg += " AlignD_bytes=" + std::to_string((long long)reqD)
            + " D_ptr_mod_req=" + std::to_string((long long)(ptrD_host[size_t(i)] % reqD));
        msg += " | StrideA=("
            + std::to_string((long long)cute::get<0>(h_sA_dbg[size_t(i)])) + ","
            + std::to_string((long long)cute::get<1>(h_sA_dbg[size_t(i)])) + ","
            + std::to_string((long long)cute::get<2>(h_sA_dbg[size_t(i)])) + ")";
        msg += " StrideB=("
            + std::to_string((long long)cute::get<0>(h_sB_dbg[size_t(i)])) + ","
            + std::to_string((long long)cute::get<1>(h_sB_dbg[size_t(i)])) + ","
            + std::to_string((long long)cute::get<2>(h_sB_dbg[size_t(i)])) + ")";
        msg += " StrideC=("
            + std::to_string((long long)cute::get<0>(h_sC_dbg[size_t(i)])) + ","
            + std::to_string((long long)cute::get<1>(h_sC_dbg[size_t(i)])) + ","
            + std::to_string((long long)cute::get<2>(h_sC_dbg[size_t(i)])) + ")";
        msg += " StrideD=("
            + std::to_string((long long)cute::get<0>(h_sD_dbg[size_t(i)])) + ","
            + std::to_string((long long)cute::get<1>(h_sD_dbg[size_t(i)])) + ","
            + std::to_string((long long)cute::get<2>(h_sD_dbg[size_t(i)])) + ")";
        msg += " | ExpStride2(A,B,C,D)=(0,0,0,0)";
        auto sfa_layout_i = h_lSFA_dbg[size_t(i)];
        auto sfb_layout_i = h_lSFB_dbg[size_t(i)];
        msg += " | SFA_size=" + std::to_string((long long)cute::size(sfa_layout_i));
        msg += " SFA_compact=" + std::to_string((long long)cute::size(cute::filter_zeros(sfa_layout_i)));
        msg += " SFB_size=" + std::to_string((long long)cute::size(sfb_layout_i));
        msg += " SFB_compact=" + std::to_string((long long)cute::size(cute::filter_zeros(sfb_layout_i)));
        msg += " | TileMNK=("
            + std::to_string((long long)cute::size<0>(MmaTileShape{})) + ","
            + std::to_string((long long)cute::size<1>(MmaTileShape{})) + ","
            + std::to_string((long long)cute::size<2>(MmaTileShape{})) + ")";
        msg += " Cluster=("
            + std::to_string((long long)cute::size<0>(ClusterShape{})) + ","
            + std::to_string((long long)cute::size<1>(ClusterShape{})) + ","
            + std::to_string((long long)cute::size<2>(ClusterShape{})) + ")";
        found = true;
        break;
      }
    }
    if (!found) {
      msg += " | per_group_probe=all_pass";
    }
    msg += " | can_impl_nohost=" + status_to_string(st_nohost);

    throw std::runtime_error(msg);
  }

  if (tmode > 0) t_ws0 = std::chrono::high_resolution_clock::now();
  size_t workspace_size = Gemm::get_workspace_size(args);
  if (!entry->workspace.defined() || entry->workspace.numel() < (long long)workspace_size) {
    entry->workspace = torch::empty({(long long)workspace_size}, bytes);
  }
  if (tmode > 0) t_ws1 = std::chrono::high_resolution_clock::now();

  bool do_synergy = synergy_out.defined() && synergy_out.numel() >= 5 &&
                    potential_map.defined() && potential_map.numel() > 0;
  int64_t map_size = do_synergy ? potential_map.numel() : 0;
  float kernel_ms = 0.0f;
  int epi_threads = 32;
  const char* epi_env = std::getenv("NVFP4_EPILOGUE_THREADS");
  if (epi_env) {
    int v = std::atoi(epi_env);
    if (v > 0) epi_threads = v;
  }

  int epi_burn = 200;
  const char* epi_burn_env = std::getenv("NVFP4_EPILOGUE_BURN");
  if (epi_burn_env) {
    int v = std::atoi(epi_burn_env);
    if (v > 0) epi_burn = v;
  }

  if (tmode > 0) {
    cudaEvent_t ev_start, ev_stop;
    cudaEventCreate(&ev_start);
    cudaEventCreate(&ev_stop);
    if (do_synergy) {
      synergy_prologue_kernel<<<1, 1, 0, 0>>>(
          synergy_out.data_ptr<int64_t>());
    }
    cudaEventRecord(ev_start, 0);
    st = gemm(args, entry->workspace.data_ptr(), 0);
    cudaEventRecord(ev_stop, 0);
    cudaEventSynchronize(ev_stop);
    cudaEventElapsedTime(&kernel_ms, ev_start, ev_stop);
    cudaEventDestroy(ev_start);
    cudaEventDestroy(ev_stop);
    if (do_synergy) {
      synergy_epilogue_kernel<<<1, epi_threads, 0, 0>>>(
          synergy_out.data_ptr<int64_t>(),
          potential_map.data_ptr<float>(),
          map_size,
          bucket_id,
          total_work_items,
          epi_burn);
    }
  } else {
    if (do_synergy) {
      synergy_prologue_kernel<<<1, 1, 0, 0>>>(
          synergy_out.data_ptr<int64_t>());
    }
    st = gemm(args, entry->workspace.data_ptr(), 0);
    if (do_synergy) {
      synergy_epilogue_kernel<<<1, epi_threads, 0, 0>>>(
          synergy_out.data_ptr<int64_t>(),
          potential_map.data_ptr<float>(),
          map_size,
          bucket_id,
          total_work_items,
          epi_burn);
    }
  }
  if (st != cutlass::Status::kSuccess) {
    throw std::runtime_error("GEMM launch failed: " + status_to_string(st));
  }

  if (tmode > 0) {
    t_sync0 = std::chrono::high_resolution_clock::now();
    cudaDeviceSynchronize();
    t_sync1 = std::chrono::high_resolution_clock::now();
    auto t1 = std::chrono::high_resolution_clock::now();
    double total_ms = std::chrono::duration<double, std::milli>(t1 - t0).count();
    double setup_ms = total_ms - double(kernel_ms);
    std::fprintf(stderr,
                 "[nvfp4] g=%lld setup_ms=%.3f kernel_ms=%.3f total_ms=%.3f workspace_bytes=%lld\n",
                 (long long)g, setup_ms, kernel_ms, total_ms, (long long)workspace_size);
    if (tdetail > 0) {
      auto ms = [](auto a, auto b) {
        return std::chrono::duration<double, std::milli>(b - a).count();
      };
      double ps_ms = ms(t_ps0, t_ps1);
      double cache_ms = ms(t_cache0, t_cache1);
      double build_ms = ms(t_build0, t_build1);
      double upload_ms = ms(t_upload0, t_upload1);
      double can_ms = ms(t_can0, t_can1);
      double ws_ms = ms(t_ws0, t_ws1);
      double sync_ms = ms(t_sync0, t_sync1);
      std::fprintf(stderr,
                   "[nvfp4_detail] ps_ms=%.3f cache_ms=%.3f build_ms=%.3f upload_ms=%.3f can_ms=%.3f ws_ms=%.3f sync_ms=%.3f\n",
                   ps_ms, cache_ms, build_ms, upload_ms, can_ms, ws_ms, sync_ms);
    }
    std::fflush(stderr);
    if (timing_abort() > 0) {
      throw std::runtime_error("NVFP4_TIMING_ABORT");
    }
  }
}

#else

void grouped_nvfp4_tma_gemm(
    torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor,
    torch::Tensor, torch::Tensor, int64_t, int64_t
) {
  throw std::runtime_error("SM100 support not enabled");
}

#endif

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("grouped_nvfp4_tma_gemm", &grouped_nvfp4_tma_gemm, "SM100 grouped NVFP4 GEMM");
  m.def("convert_sf_float_to_ue4m3", &convert_sf_float_to_ue4m3, "Convert float scale factors to ue4m3");
  m.def("pad_a_fp4", &pad_a_fp4, "Pad A for FP4 with zero fill");
  m.def("pad_c_half", &pad_c_half, "Pad C for fp16 with zero fill");
}
"""

        burn_macro = int(os.environ.get("NVFP4_INKERNEL_BURN", "100000"))
        build_name = f"ek_nvfp4_sm100_ex81_pattern_v3_burn{burn_macro}_probe28"
        print(f"[nvfp4] build_name={build_name} NVFP4_INKERNEL_BURN={burn_macro}", file=os.sys.stderr, flush=True)

        extra_cuda_cflags = [
            "-O3",
            "--use_fast_math",
            "--std=c++17",
            "--expt-relaxed-constexpr",
            "--expt-extended-lambda",
            "-lineinfo",
            "-gencode=arch=compute_100a,code=sm_100a",
            "-gencode=arch=compute_100a,code=compute_100a",
            "-DCUTLASS_ARCH_MMA_SM100_SUPPORTED",
            "-DCUTLASS_ARCH_MMA_SM100A_ENABLED",
            "-UCUTLASS_ENABLE_SYNCLOG",
        ]
        if burn_macro > 0:
            extra_cuda_cflags.append(f"-DNVFP4_INKERNEL_BURN={burn_macro}")
            extra_cuda_cflags.append(f"-DNVFP4_INKERNEL_BURN_MARKER={burn_macro}")
        extra_cflags = ["-O3", "-std=c++17"]

        custom_kernel._EXT = load_inline(
            name=build_name,
            cpp_sources="",
            cuda_sources=cuda_src,
            functions=None,
            extra_cuda_cflags=extra_cuda_cflags,
            extra_cflags=extra_cflags,
            with_cuda=True,
            verbose=False,
            extra_include_paths=extra_includes,
        )

    ext = custom_kernel._EXT

    # -----------------------------
    # 2) Build pointer arrays (3D kernel; expand L into groups)
    # -----------------------------
    dev = abc_tensors[0][2].device

    abc_padded = []
    sfs_out = []
    problem_sizes_out = []
    c_pad_info = []  # (orig_c, c_pad, m, n, l_idx)

    def _maybe_contig(t):
        return t if t.is_contiguous() else t.contiguous()

    graph_safe = True
    for i in range(g):
        a, b, c = abc_tensors[i]
        sfa, sfb = sfasfb_tensors[i][0], sfasfb_tensors[i][1]
        m, n, k, l = problem_sizes[i]

        if l <= 0:
            raise RuntimeError(f"Invalid L={l} for group {i}")
        if l != 1:
            raise RuntimeError(f"Only L=1 is supported in this fast path (got L={l} for group {i})")

        if sfa.dtype == torch.float32:
            graph_safe = False
            sfa = ext.convert_sf_float_to_ue4m3(sfa)
        elif sfa.dtype == torch.float8_e4m3fnuz:
            pass
        elif sfa.dtype == torch.float8_e4m3fn:
            pass
        elif sfa.dtype != torch.uint8:
            raise TypeError("SFA tensors must be float32, float8_e4m3fnuz, float8_e4m3fn, or uint8 (ue4m3 storage).")
        if sfb.dtype == torch.float32:
            graph_safe = False
            sfb = ext.convert_sf_float_to_ue4m3(sfb)
        elif sfb.dtype == torch.float8_e4m3fnuz:
            pass
        elif sfb.dtype == torch.float8_e4m3fn:
            pass
        elif sfb.dtype != torch.uint8:
            raise TypeError("SFB tensors must be float32, float8_e4m3fnuz, float8_e4m3fn, or uint8 (ue4m3 storage).")
        if not sfa.is_cuda or not sfb.is_cuda:
            raise TypeError("SFA/SFB tensors must be CUDA tensors.")

        a_i = _maybe_contig(a)
        b_i = _maybe_contig(b)
        c_i = _maybe_contig(c)
        if a_i.data_ptr() != a.data_ptr() or b_i.data_ptr() != b.data_ptr() or c_i.data_ptr() != c.data_ptr():
            graph_safe = False
        sfa_i = sfa
        sfb_i = sfb

        div_sfa = (k // 16)
        div_sfb = (k // 16)
        if div_sfa <= 0 or (sfa_i.numel() % div_sfa) != 0:
            raise RuntimeError(f"SFA size mismatch for group {i}: sfa.numel={sfa_i.numel()}")
        if div_sfb <= 0 or (sfb_i.numel() % div_sfb) != 0:
            raise RuntimeError(f"SFB size mismatch for group {i}: sfb.numel={sfb_i.numel()}")
        m_eff = sfa_i.numel() // div_sfa
        n_eff = sfb_i.numel() // div_sfb

        if m_eff < m:
            raise RuntimeError(f"SFA implies smaller M for group {i}: m_eff={m_eff} m={m}")
        if n_eff < n:
            raise RuntimeError(f"SFB implies smaller N for group {i}: n_eff={n_eff} n={n}")

        skip_pad = int(os.environ.get("NVFP4_SKIP_PAD", "1"))
        if m_eff != m:
            if skip_pad > 0:
                m_eff = m
            else:
                graph_safe = False
                a_i = ext.pad_a_fp4(a_i, m_eff)

        if m_eff != m or n_eff != n:
            if skip_pad > 0:
                n_eff = n
            else:
                graph_safe = False
                c_i = ext.pad_c_half(c_i, m_eff, n_eff)
                c_pad_info.append((abc_tensors[i][2], c_i, m, n, None))

        abc_padded.append((a_i, b_i, c_i))
        sfs_out.append((sfa_i, sfb_i))
        problem_sizes_out.append((int(m_eff), int(n_eff), int(k)))

    if timing_py:
        t_py1 = time.perf_counter()

    def _build_cache_key_prepared(abc_tensors, sfasfb_tensors, problem_sizes):
        abc_sig = tuple((_tensor_key(a), _tensor_key(b), _tensor_key(c)) for a, b, c in abc_tensors)
        sfs_sig = tuple((_tensor_key(sfa), _tensor_key(sfb)) for sfa, sfb in sfasfb_tensors)
        ps_sig = tuple((int(m), int(n), int(k)) for (m, n, k) in problem_sizes)
        return (abc_sig, sfs_sig, ps_sig)

    cache_key = _build_cache_key_prepared(abc_padded, sfs_out, problem_sizes_out) if use_py_cache else None
    cached_state = custom_kernel._CACHE.get(cache_key) if use_py_cache else None
    if use_py_cache and use_echokey and cached_state is not None and cached_state.get("echokey_pack") is None:
        cached_state = None

    A_ptrs = [int(abc_padded[i][0].data_ptr()) for i in range(len(problem_sizes_out))]
    B_ptrs = [int(abc_padded[i][1].data_ptr()) for i in range(len(problem_sizes_out))]
    C_ptrs = [int(abc_padded[i][2].data_ptr()) for i in range(len(problem_sizes_out))]
    D_ptrs = C_ptrs

    SFA_ptrs = [int(sfs_out[i][0].data_ptr()) for i in range(len(problem_sizes_out))]
    SFB_ptrs = [int(sfs_out[i][1].data_ptr()) for i in range(len(problem_sizes_out))]

    def _update_ptr_tensor(dst, src_list):
        src = torch.as_tensor(src_list, device="cpu", dtype=torch.int64)
        dst.copy_(src, non_blocking=True)

    if cached_state is not None:
        ps = cached_state["ps"]
        ptr_A = cached_state["ptr_A"]
        ptr_B = cached_state["ptr_B"]
        ptr_SFA = cached_state["ptr_SFA"]
        ptr_SFB = cached_state["ptr_SFB"]
        ptr_C = cached_state["ptr_C"]
        ptr_D = cached_state["ptr_D"]

        prev_ptrs = cached_state.get("ptr_host", None)
        cur_ptrs = (A_ptrs, B_ptrs, SFA_ptrs, SFB_ptrs, C_ptrs, D_ptrs)
        if prev_ptrs != cur_ptrs:
            _update_ptr_tensor(ptr_A, A_ptrs)
            _update_ptr_tensor(ptr_B, B_ptrs)
            _update_ptr_tensor(ptr_SFA, SFA_ptrs)
            _update_ptr_tensor(ptr_SFB, SFB_ptrs)
            _update_ptr_tensor(ptr_C, C_ptrs)
            _update_ptr_tensor(ptr_D, D_ptrs)
            cached_state["ptr_host"] = cur_ptrs
        custom_kernel._CACHE.move_to_end(cache_key)
    else:
        ptr_A = torch.tensor(A_ptrs, device=dev, dtype=torch.int64)
        ptr_B = torch.tensor(B_ptrs, device=dev, dtype=torch.int64)
        ptr_SFA = torch.tensor(SFA_ptrs, device=dev, dtype=torch.int64)
        ptr_SFB = torch.tensor(SFB_ptrs, device=dev, dtype=torch.int64)
        ptr_C = torch.tensor(C_ptrs, device=dev, dtype=torch.int64)
        ptr_D = torch.tensor(D_ptrs, device=dev, dtype=torch.int64)
        ps = torch.tensor(
            [(int(m), int(n), int(k)) for (m, n, k) in problem_sizes_out],
            device=dev,
            dtype=torch.int32,
        )

    if timing_py:
        t_py2 = time.perf_counter()

    echokey_pack = None
    if use_echokey:
        if cached_state is not None and cached_state.get("echokey_pack") is not None:
            echokey_pack = cached_state["echokey_pack"]
        else:
            bucket_id = _bucket_id(problem_sizes, map_size)
            total_work_items = 0
            for (m, n, k) in problem_sizes_out:
                total_work_items += int(m) * int(n) * int(k)
            if dev not in custom_kernel._POTENTIAL_MAP:
                custom_kernel._POTENTIAL_MAP[dev] = torch.zeros(map_size, dtype=torch.float32, device=dev)
            if dev not in custom_kernel._SYNERGY_OUT:
                custom_kernel._SYNERGY_OUT[dev] = torch.zeros(5, dtype=torch.int64, device=dev)
            potential_map = custom_kernel._POTENTIAL_MAP[dev]
            synergy_out = custom_kernel._SYNERGY_OUT[dev]
            echokey_pack = {
                "bucket_id": bucket_id,
                "total_work_items": total_work_items,
                "potential_map": potential_map,
                "synergy_out": synergy_out,
                "warmed": False,
            }

    if use_py_cache:
        custom_kernel._CACHE[cache_key] = {
            "ps": ps,
            "ptr_A": ptr_A,
            "ptr_B": ptr_B,
            "ptr_SFA": ptr_SFA,
            "ptr_SFB": ptr_SFB,
            "ptr_C": ptr_C,
            "ptr_D": ptr_D,
            "ptr_host": (A_ptrs, B_ptrs, SFA_ptrs, SFB_ptrs, C_ptrs, D_ptrs),
            "echokey_pack": echokey_pack,
        }
        custom_kernel._CACHE.move_to_end(cache_key)
        max_cache = int(os.environ.get("NVFP4_PY_CACHE_SIZE", "8"))
        while len(custom_kernel._CACHE) > max_cache:
            custom_kernel._CACHE.popitem(last=False)

    # -----------------------------
    # 3) Run grouped GEMM
    # -----------------------------
    if timing_py:
        torch.cuda.synchronize(dev)
        t_call0 = time.perf_counter()
    if use_echokey:
        synergy_out = echokey_pack["synergy_out"]
        potential_map = echokey_pack["potential_map"]
        bucket_id = int(echokey_pack["bucket_id"])
        total_work_items = int(echokey_pack["total_work_items"])
    else:
        synergy_out, potential_map = custom_kernel._ECHO_DUMMY
        bucket_id = 0
        total_work_items = 0

    use_graph = int(os.environ.get("NVFP4_USE_GRAPH", "0")) > 0
    use_manual_launch = True
    if use_graph and use_py_cache and graph_safe:
        if not hasattr(custom_kernel, "_GRAPH_CACHE"):
            custom_kernel._GRAPH_CACHE = {}
        graph_key = (cache_key, "graph", tuple(A_ptrs), tuple(B_ptrs), tuple(C_ptrs), tuple(SFA_ptrs), tuple(SFB_ptrs))
        g_entry = custom_kernel._GRAPH_CACHE.get(graph_key)
        if g_entry is None:
            g_entry = {"graph": None, "armed": False, "disabled": False}
            custom_kernel._GRAPH_CACHE[graph_key] = g_entry

        if g_entry["disabled"]:
            pass
        elif g_entry["graph"] is None:
            if not g_entry["armed"]:
                # Arm on first encounter; capture on next call to avoid skewing first timing.
                g_entry["armed"] = True
            else:
                # Warmup to ensure allocations are done
                ext.grouped_nvfp4_tma_gemm(
                    ps, ptr_A, ptr_B, ptr_SFA, ptr_SFB, ptr_C, ptr_D,
                    synergy_out, potential_map, bucket_id, total_work_items
                )
                torch.cuda.synchronize(dev)
                graph = torch.cuda.CUDAGraph()
                torch.cuda.set_device(dev)
                import warnings
                with warnings.catch_warnings(record=True) as w:
                    warnings.simplefilter("always")
                    with torch.cuda.graph(graph):
                        ext.grouped_nvfp4_tma_gemm(
                            ps, ptr_A, ptr_B, ptr_SFA, ptr_SFB, ptr_C, ptr_D,
                            synergy_out, potential_map, bucket_id, total_work_items
                        )
                if w:
                    g_entry["disabled"] = True
                else:
                    g_entry["graph"] = graph
                    graph.replay()
                    use_manual_launch = False
        else:
            g_entry["graph"].replay()
            use_manual_launch = False

    if use_manual_launch:
        if use_echokey:
            if use_warmup and not echokey_pack.get("warmed", False):
                ext.grouped_nvfp4_tma_gemm(
                    ps, ptr_A, ptr_B, ptr_SFA, ptr_SFB, ptr_C, ptr_D,
                    synergy_out, potential_map, bucket_id, total_work_items
                )
                echokey_pack["warmed"] = True
            ext.grouped_nvfp4_tma_gemm(
                ps, ptr_A, ptr_B, ptr_SFA, ptr_SFB, ptr_C, ptr_D,
                synergy_out, potential_map, bucket_id, total_work_items
            )
        else:
            ext.grouped_nvfp4_tma_gemm(
                ps, ptr_A, ptr_B, ptr_SFA, ptr_SFB, ptr_C, ptr_D,
                synergy_out, potential_map, bucket_id, total_work_items
            )
    if use_echokey and echokey_debug:
        torch.cuda.synchronize(dev)
        dbg = echokey_pack["synergy_out"].detach().cpu().tolist()
        print(f"[echokey] bucket={echokey_pack['bucket_id']} dt={dbg[2]}", file=os.sys.stderr, flush=True)
    if timing_py:
        torch.cuda.synchronize(dev)
        t_call1 = time.perf_counter()

    for orig_c, c_pad, m, n, l_idx in c_pad_info:
        if l_idx is None:
            orig_c.copy_(c_pad[:m, :n, :])
        else:
            orig_c[:, :, l_idx].copy_(c_pad[:m, :n, 0])

    if timing_py:
        t_py3 = time.perf_counter()
        prep_ms = (t_py1 - t_py0) * 1e3
        ptr_ms = (t_py2 - t_py1) * 1e3
        call_ms = (t_call1 - t_call0) * 1e3
        post_ms = (t_py3 - t_call1) * 1e3
        total_ms = (t_py3 - t_py0) * 1e3
        print(
            f"[nvfp4_py] prep_ms={prep_ms:.3f} ptr_ms={ptr_ms:.3f} "
            f"call_ms={call_ms:.3f} post_ms={post_ms:.3f} total_ms={total_ms:.3f}",
            file=os.sys.stderr,
            flush=True,
        )

    return [abc_tensors[i][2] for i in range(g)]
scrolls · 1645 lines total

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

Best evidence level for this revision: reported

JSON