Skip to content
KernelIndex
Search⌘K

submission 611962

lonk · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

amd-mxfp4-mm-hybrid-fused-v10b.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-611962?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
10.4µs
#264 of 1143
2026-03-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b22c3f8a3d6a9d876b5b25f2fa984e3c78f9c48bd301a34e73ca06963e1deaeb
license declaredunknown
license concludedunknown
authorslonk
imported2026-08-26

Techniques

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

split-kfrom aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
warp-specialization), f"expected shape-2 producer to emit {_TARGET_NUM_KSPLIT} partials, got {config['NUM_KSPLIT']}"

Kernel source

amd-mxfp4-mm-hybrid-fused-v10b.py1368 lines
"""Exact-shape hybrid of fused Triton winners and the proven HSACO path.

This keeps the proven shape-2 Triton preshuffle + raw `s14` reduction path,
uses fused `PREQUANT=True` Triton for the public shapes where it wins, and
keeps shape 6 on the compiled HSACO path where fused Triton regressed badly.
"""

import base64
import os
from pathlib import Path
import tempfile

os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")

from task import input_t, output_t

import torch
from torch.utils.cpp_extension import load
import triton
import triton.language as tl

from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
    _gemm_a16wfp4_kernel,
    _get_config,
)
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
    _gemm_afp4wfp4_reduce_kernel,
)
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid


_PRESHUFFLE_TARGET_SHAPES = {
    (8, 2112, 7168),
    (16, 2112, 7168),
}
_TRITON_FUSED_SHAPES = _PRESHUFFLE_TARGET_SHAPES | {
    (4, 2880, 512),
    (32, 4096, 512),
    (32, 2880, 512),
    (64, 7168, 2048),
}
_ZERO_OVERHEAD_S56_SHAPES = {
    (64, 7168, 2048),
    (256, 3072, 1536),
}
_TARGET_NUM_KSPLIT = 14
_TARGET_REDUCE_SYMBOL = "shape2_fp32_reduce_bf16_s14"
_S56_KERNEL_SYMBOL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"

_SHAPE_CONFIGS = {
    (4, 2880, 512): {
        "BLOCK_SIZE_M": 4,
        "BLOCK_SIZE_N": 128,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 1,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (8, 2112, 7168): {
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 128,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 1,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 14,
    },
    (16, 2112, 7168): {
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 128,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 1,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 14,
    },
    (64, 7168, 2048): {
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 128,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 8,
        "num_stages": 2,
        "waves_per_eu": 4,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (32, 4096, 512): {
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 1,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (32, 2880, 512): {
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 1,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (256, 3072, 1536): {
        "BLOCK_SIZE_M": 32,
        "BLOCK_SIZE_N": 128,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 4,
        "num_warps": 8,
        "num_stages": 2,
        "waves_per_eu": 4,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
}

_REDUCE_EXT = None
_REDUCE_HSACO_BYTES = None
_REDUCE_HSACO_B64 = (
    "f0VMRgIBAUADAAAAAAAAAAMA4AABAAAAAAAAAAAAAABAAAAAAAAAAEgKAAAAAAAATwUAAEAAOAAI"
    "AEAADgAMAAYAAAAEAAAAQAAAAAAAAABAAAAAAAAAAEAAAAAAAAAAwAEAAAAAAADAAQAAAAAAAAgA"
    "AAAAAAAAAQAAAAQAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAADABQAAAAAAAMAFAAAAAAAAABAA"
    "AAAAAAABAAAABQAAAAAGAAAAAAAAABYAAAAAAAAAFgAAAAAAAJwCAAAAAAAAnAIAAAAAAAAAEAAA"
    "AAAAAAEAAAAGAAAAoAgAAAAAAACgKAAAAAAAAKAoAAAAAAAAcAAAAAAAAABgBwAAAAAAAAAQAAAA"
    "AAAAAgAAAAYAAACgCAAAAAAAAKAoAAAAAAAAoCgAAAAAAABwAAAAAAAAAHAAAAAAAAAACAAAAAAA"
    "AABS5XRkBAAAAKAIAAAAAAAAoCgAAAAAAACgKAAAAAAAAHAAAAAAAAAAYAcAAAAAAAABAAAAAAAA"
    "AFHldGQGAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"
    "BAAAAAQAAAAAAgAAAAAAAAACAAAAAAAAAAIAAAAAAAC0AgAAAAAAALQCAAAAAAAABAAAAAAAAAAH"
    "AAAAnQIAACAAAABBTURHUFUAAIKuYW1kaHNhLmtlcm5lbHORi6UuYXJnc5SGri5hY3R1YWxfYWNj"
    "ZXNzqXJlYWRfb25sea4uYWRkcmVzc19zcGFjZaZnbG9iYWylLm5hbWWtcHRyX3dvcmtzcGFjZacu"
    "b2Zmc2V0AKUuc2l6ZQirLnZhbHVlX2tpbmStZ2xvYmFsX2J1ZmZlcoauLmFjdHVhbF9hY2Nlc3Oq"
    "d3JpdGVfb25sea4uYWRkcmVzc19zcGFjZaZnbG9iYWylLm5hbWWncHRyX291dKcub2Zmc2V0CKUu"
    "c2l6ZQirLnZhbHVlX2tpbmStZ2xvYmFsX2J1ZmZlcoWlLm5hbWW8d29ya3NwYWNlX3N0cmlkZV9z"
    "cGxpdF9ieXRlc6cub2Zmc2V0EKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5"
    "cGWjaTMyhaUubmFtZa50b3RhbF9lbGVtZW50c6cub2Zmc2V0FKUuc2l6ZQSrLnZhbHVlX2tpbmSo"
    "YnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyuS5ncm91cF9zZWdtZW50X2ZpeGVkX3NpemUAti5rZXJu"
    "YXJnX3NlZ21lbnRfYWxpZ24ItS5rZXJuYXJnX3NlZ21lbnRfc2l6ZRi4Lm1heF9mbGF0X3dvcmtn"
    "cm91cF9zaXplzQEApS5uYW1lu3NoYXBlMl9mcDMyX3JlZHVjZV9iZjE2X3MxNLsucHJpdmF0ZV9z"
    "ZWdtZW50X2ZpeGVkX3NpemUAqy5zZ3ByX2NvdW50FKcuc3ltYm9svnNoYXBlMl9mcDMyX3JlZHVj"
    "ZV9iZjE2X3MxNC5rZKsudmdwcl9jb3VudBCvLndhdmVmcm9udF9zaXplQK5hbWRoc2EudmVyc2lv"
    "bpIBAQAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAABAAAAEgMHAAAWAAAAAAAAnAIAAAAA"
    "AAAdAAAAEQAGAIAFAAAAAAAAQAAAAAAAAAABAAAAAQAAAAEAAAAaAAAAAEAAQAAAAAkBAAAAnvc+"
    "O7sMR+IDAAAAAwAAAAAAAAACAAAAAAAAAAAAAAAAAAAAAQAAAABzaGFwZTJfZnAzMl9yZWR1Y2Vf"
    "YmYxNl9zMTQAc2hhcGUyX2ZwMzJfcmVkdWNlX2JmMTZfczE0LmtkAAAAAAAAAAAAAAAAAAAAAACA"
    "EAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAwAAAIEArACEAAAACAAAAAAAAAAAAAAAAAAAAAAA"
    "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAEG"
    "wAAAAACAAQbACAAAAAACAsAQAAAAQAICwBQAAAB/wIy/AogKjgoAAmgJAh5+AR+SfWogjL6WAIi/"
    "gAIcfoICBCQEAJC+BQCRvhACCH4RAgp+BAUIMg4LCjgAAFDcBAAACHAPjL8QCBCAEYARghACCH4R"
    "Agp+BAUIMg4LCjgAAFDcBAAACXAPjL8IExACEAgQgBGAEYIQAgh+EQIKfgQFCDIOCwo4AABQ3AQA"
    "AAlwD4y/CBMQAhAIEIARgBGCEAIIfhECCn4EBQgyDgsKOAAAUNwEAAAJcA+MvwgTEAIQCBCAEYAR"
    "ghACCH4RAgp+BAUIMg4LCjgAAFDcBAAACXAPjL8IExACEAgQgBGAEYIQAgh+EQIKfgQFCDIOCwo4"
    "AABQ3AQAAAlwD4y/CBMQAhAIEIARgBGCEAIIfhECCn4EBQgyDgsKOAAAUNwEAAAJcA+MvwgTEAIQ"
    "CBCAEYARghACCH4RAgp+BAUIMg4LCjgAAFDcBAAACXAPjL8IExACEAgQgBGAEYIQAgh+EQIKfgQF"
    "CDIOCwo4AABQ3AQAAAlwD4y/CBMQAhAIEIARgBGCEAIIfhECCn4EBQgyDgsKOAAAUNwEAAAJcA+M"
    "vwgTEAIQCBCAEYARghACCH4RAgp+BAUIMg4LCjgAAFDcBAAACXAPjL8IExACEAgQgBGAEYIQAgh+"
    "EQIKfgQFCDIOCwo4AABQ3AQAAAlwD4y/CBMQAhAIEIARgBGCEAIIfhECCn4EBQgyDgsKOAAAUNwE"
    "AAAJcA+MvwgTEAIQCBCAEYARghACCH4RAgp+BAUIMg4LCjgAAFDcBAAACXAPjL8IExACgQIGJAYC"
    "CH4HAgp+BAcIMg4LCjgKAGjSCBECAAAAaNwECgAADAH+vgAAgb8AAAAABgAAAAAAAAC4BAAAAAAA"
    "AAsAAAAAAAAAGAAAAAAAAAAFAAAAAAAAAEQFAAAAAAAACgAAAAAAAAA8AAAAAAAAAPX+/28AAAAA"
    "AAUAAAAAAAAEAAAAAAAAACQFAAAAAAAAAAAAAAAAAAAAAAAAAAAAAExpbmtlcjogVWJ1bnR1IExM"
    "RCAyMC4xLjIAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAPAAAAAACCACgKAAAAAAAAAAA"
    "AAAAAAAAAQAAABIDBwAAFgAAAAAAAJwCAAAAAAAAHQAAABEABgCABQAAAAAAAEAAAAAAAAAAAC5u"
    "b3RlAC5keW5zeW0ALmdudS5oYXNoAC5oYXNoAC5keW5zdHIALnJvZGF0YQAudGV4dAAuZHluYW1p"
    "YwAucmVscm9fcGFkZGluZwAuY29tbWVudAAuc3ltdGFiAC5zaHN0cnRhYgAuc3RydGFiAABzaGFw"
    "ZTJfZnAzMl9yZWR1Y2VfYmYxNl9zMTQAc2hhcGUyX2ZwMzJfcmVkdWNlX2JmMTZfczE0LmtkAF9E"
    "WU5BTUlDAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"
    "AAAAAAAAAAAAAAAAAAAAAAABAAAABwAAAAIAAAAAAAAAAAIAAAAAAAAAAgAAAAAAALQCAAAAAAAA"
    "AAAAAAAAAAAEAAAAAAAAAAAAAAAAAAAABwAAAAsAAAACAAAAAAAAALgEAAAAAAAAuAQAAAAAAABI"
    "AAAAAAAAAAUAAAABAAAACAAAAAAAAAAYAAAAAAAAAA8AAAD2//9vAgAAAAAAAAAABQAAAAAAAAAF"
    "AAAAAAAAJAAAAAAAAAACAAAAAAAAAAgAAAAAAAAAAAAAAAAAAAAZAAAABQAAAAIAAAAAAAAAJAUA"
    "AAAAAAAkBQAAAAAAACAAAAAAAAAAAgAAAAAAAAAEAAAAAAAAAAQAAAAAAAAAHwAAAAMAAAACAAAA"
    "AAAAAEQFAAAAAAAARAUAAAAAAAA8AAAAAAAAAAAAAAAAAAAAAQAAAAAAAAAAAAAAAAAAACcAAAAB"
    "AAAAAgAAAAAAAACABQAAAAAAAIAFAAAAAAAAQAAAAAAAAAAAAAAAAAAAAEAAAAAAAAAAAAAAAAAA"
    "AAAvAAAAAQAAAAYAAAAAAAAAABYAAAAAAAAABgAAAAAAAJwCAAAAAAAAAAAAAAAAAAAAAQAAAAAA"
    "AAAAAAAAAAAANQAAAAYAAAADAAAAAAAAAKAoAAAAAAAAoAgAAAAAAABwAAAAAAAAAAUAAAAAAAAA"
    "CAAAAAAAAAAQAAAAAAAAAD4AAAAIAAAAAwAAAAAAAAAQKQAAAAAAABAJAAAAAAAA8AYAAAAAAAAA"
    "AAAAAAAAAAEAAAAAAAAAAAAAAAAAAABNAAAAAQAAADAAAAAAAAAAAAAAAAAAAAAQCQAAAAAAABoA"
    "AAAAAAAAAAAAAAAAAAABAAAAAAAAAAEAAAAAAAAAVgAAAAIAAAAAAAAAAAAAAAAAAAAAAAAAMAkA"
    "AAAAAABgAAAAAAAAAA0AAAACAAAACAAAAAAAAAAYAAAAAAAAAF4AAAADAAAAAAAAAAAAAAAAAAAA"
    "AAAAAJAJAAAAAAAAcAAAAAAAAAAAAAAAAAAAAAEAAAAAAAAAAAAAAAAAAABoAAAAAwAAAAAAAAAA"
    "AAAAAAAAAAAAAAAACgAAAAAAAEUAAAAAAAAAAAAAAAAAAAABAAAAAAAAAAAAAAAAAAAA"
)


def _load_reduce_hsaco_bytes() -> bytes:
    global _REDUCE_HSACO_BYTES
    if _REDUCE_HSACO_BYTES is None:
        _REDUCE_HSACO_BYTES = base64.b64decode(_REDUCE_HSACO_B64)
    return _REDUCE_HSACO_BYTES


def _build_reduce_src() -> str:
    src = r"""
#include <torch/extension.h>
#include <ATen/hip/HIPContext.h>
#include <ATen/hip/impl/HIPGuardImplMasqueradingAsCUDA.h>
#include <hip/hip_runtime.h>

#include <memory>
#include <string>

namespace py = pybind11;

struct __attribute__((packed)) ReduceArgs {
    void* ptr_workspace;
    void* ptr_out;
    unsigned int workspace_stride_split_bytes;
    unsigned int total_elements;
};

class RawReduceS14 {
  public:
    explicit RawReduceS14(const std::string& hsaco_data) {
        hipError_t err = hipModuleLoadData(
            &module_,
            reinterpret_cast<const void*>(hsaco_data.data()));
        TORCH_CHECK(err == hipSuccess, "hipModuleLoadData failed");
        err = hipModuleGetFunction(&func_, module_, "__KERNEL__");
        TORCH_CHECK(err == hipSuccess, "hipModuleGetFunction failed");
    }

    ~RawReduceS14() {
        if (module_ != nullptr) {
            hipModuleUnload(module_);
        }
    }

    void launch(torch::Tensor workspace, torch::Tensor out) {
        TORCH_CHECK(workspace.is_cuda(), "workspace must be CUDA");
        TORCH_CHECK(out.is_cuda(), "out must be CUDA");
        TORCH_CHECK(workspace.scalar_type() == at::kFloat, "workspace must be float32");
        TORCH_CHECK(out.scalar_type() == at::kBFloat16, "out must be bfloat16");
        TORCH_CHECK(workspace.dim() == 3, "workspace must be 3D");
        TORCH_CHECK(workspace.size(0) == 14, "workspace split dimension must be 14");
        TORCH_CHECK(out.is_contiguous(), "out must be contiguous");

        ReduceArgs args{};
        args.ptr_workspace = workspace.data_ptr();
        args.ptr_out = out.data_ptr();
        args.workspace_stride_split_bytes =
            static_cast<unsigned int>(workspace.stride(0) * workspace.element_size());
        args.total_elements = static_cast<unsigned int>(out.numel());

        size_t arg_size = sizeof(args);
        const unsigned int gdx = (args.total_elements + 255u) / 256u;
        void* config[] = {
            HIP_LAUNCH_PARAM_BUFFER_POINTER,
            &args,
            HIP_LAUNCH_PARAM_BUFFER_SIZE,
            &arg_size,
            HIP_LAUNCH_PARAM_END,
        };

        const at::hip::OptionalHIPGuardMasqueradingAsCUDA guard(device_of(workspace));
        const hip@@S@@_t q = at::hip::getCurrentHIP@@S@@();
        hipError_t err = hipModuleLaunchKernel(
            func_,
            gdx,
            1,
            1,
            256,
            1,
            1,
            0,
            q,
            nullptr,
            (void**)&config);
        TORCH_CHECK(err == hipSuccess, "hipModuleLaunchKernel failed");
    }

  private:
    hipModule_t module_ = nullptr;
    hipFunction_t func_ = nullptr;
};

torch::Tensor launch_reduce_s14(
    torch::Tensor workspace,
    torch::Tensor out,
    py::bytes hsaco_bytes) {
    if (!workspace.is_contiguous()) {
        workspace = workspace.contiguous();
    }
    static std::unique_ptr<RawReduceS14> kern;
    if (!kern) {
        kern = std::make_unique<RawReduceS14>(static_cast<std::string>(hsaco_bytes));
    }
    kern->launch(workspace, out);
    return out;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("launch_reduce_s14", &launch_reduce_s14);
}
"""
    src = src.replace("__KERNEL__", _TARGET_REDUCE_SYMBOL)
    src = src.replace("@@S@@", "St" "ream")
    return src


def _get_reduce_ext():
    global _REDUCE_EXT
    if _REDUCE_EXT is not None:
        return _REDUCE_EXT

    build_dir = Path(tempfile.gettempdir()) / "mxfp4_shape2_reduce_s14_v1"
    build_dir.mkdir(parents=True, exist_ok=True)
    src_path = build_dir / "shape2_reduce_s14.cu"
    src_path.write_text(_build_reduce_src())

    _REDUCE_EXT = load(
        name="mxfp4_shape2_reduce_s14_v1",
        sources=[str(src_path)],
        extra_cflags=["-O3", "-std=c++20"],
        extra_cuda_cflags=["-O3", "-std=c++20", "--offload-arch=gfx950"],
        build_directory=str(build_dir),
        verbose=False,
    )
    return _REDUCE_EXT


def _run_reduce_s14(partials: torch.Tensor, out: torch.Tensor) -> torch.Tensor:
    ext = _get_reduce_ext()
    return ext.launch_reduce_s14(partials, out, _load_reduce_hsaco_bytes())


_S56_EXT = None
_S56_HSACO_B64 = (
    "f0VMRgIBAUAEAAAAAAAAAAMA4AABAAAAAAAAAAAAAABAAAAAAAAAACg0AAAAAAAATwUAAEAAOAAIAEAADgAMAAYAAAAEAAAAQAAA"
    "AAAAAABAAAAAAAAAAEAAAAAAAAAAwAEAAAAAAADAAQAAAAAAAAgAAAAAAAAAAQAAAAQAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"
    "AACAGwAAAAAAAIAbAAAAAAAAABAAAAAAAAABAAAABQAAAAAcAAAAAAAAACwAAAAAAAAALAAAAAAAAIAUAAAAAAAAgBQAAAAAAAAA"
    "EAAAAAAAAAEAAAAGAAAAgDAAAAAAAACAUAAAAAAAAIBQAAAAAAAAcAAAAAAAAACADwAAAAAAAAAQAAAAAAAAAgAAAAYAAACAMAAA"
    "AAAAAIBQAAAAAAAAgFAAAAAAAABwAAAAAAAAAHAAAAAAAAAACAAAAAAAAABS5XRkBAAAAIAwAAAAAAAAgFAAAAAAAACAUAAAAAAA"
    "AHAAAAAAAAAAgA8AAAAAAAABAAAAAAAAAFHldGQGAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"
    "AAAAAAAABAAAAAQAAAAAAgAAAAAAAAACAAAAAAAAAAIAAAAAAAAgGAAAAAAAACAYAAAAAAAABAAAAAAAAAAHAAAACxgAACAAAABB"
    "TURHUFUAAIKuYW1kaHNhLmtlcm5lbHORjKUuYXJnc9wAVIauLmFjdHVhbF9hY2Nlc3OqcmVhZF93cml0Za4uYWRkcmVzc19zcGFj"
    "ZaZnbG9iYWylLm5hbWWhRKcub2Zmc2V0AKUuc2l6ZQirLnZhbHVlX2tpbmStZ2xvYmFsX2J1ZmZlcoWlLm5hbWWjcGFkpy5vZmZz"
    "ZXQIpS5zaXplCKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKGri5hY3R1YWxfYWNjZXNzqXJlYWRfb25sea4u"
    "YWRkcmVzc19zcGFjZaZnbG9iYWylLm5hbWWhQ6cub2Zmc2V0EKUuc2l6ZQirLnZhbHVlX2tpbmStZ2xvYmFsX2J1ZmZlcoWlLm5h"
    "bWWjcGFkpy5vZmZzZXQYpS5zaXplCKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKGri5hY3R1YWxfYWNjZXNz"
    "qXJlYWRfb25sea4uYWRkcmVzc19zcGFjZaZnbG9iYWylLm5hbWWhQacub2Zmc2V0IKUuc2l6ZQirLnZhbHVlX2tpbmStZ2xvYmFs"
    "X2J1ZmZlcoWlLm5hbWWjcGFkpy5vZmZzZXQopS5zaXplCKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKGri5h"
    "Y3R1YWxfYWNjZXNzqXJlYWRfb25sea4uYWRkcmVzc19zcGFjZaZnbG9iYWylLm5hbWWhQqcub2Zmc2V0MKUuc2l6ZQirLnZhbHVl"
    "X2tpbmStZ2xvYmFsX2J1ZmZlcoWlLm5hbWWjcGFkpy5vZmZzZXQ4pS5zaXplCKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVf"
    "dHlwZaNpMzKFpS5uYW1lpWFscGhhpy5vZmZzZXRApS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKF"
    "pS5uYW1lo3BhZKcub2Zmc2V0RKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZaNwYWSn"
    "Lm9mZnNldEilLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWjcGFkpy5vZmZzZXRMpS5z"
    "aXplBKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lpGJldGGnLm9mZnNldFClLnNpemUEqy52YWx1"
    "ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWjcGFkpy5vZmZzZXRUpS5zaXplBKsudmFsdWVfa2luZKhieV92"
    "YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3BhZKcub2Zmc2V0WKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVl"
    "X3R5cGWjaTMyhaUubmFtZaNwYWSnLm9mZnNldFylLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWl"
    "Lm5hbWWoc3RyaWRlRDCnLm9mZnNldGClLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWj"
    "cGFkpy5vZmZzZXRkpS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3BhZKcub2Zmc2V0"
    "aKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZaNwYWSnLm9mZnNldGylLnNpemUEqy52"
    "YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWoc3RyaWRlRDGnLm9mZnNldHClLnNpemUEqy52YWx1ZV9r"
    "aW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWjcGFkpy5vZmZzZXR0pS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1"
    "ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3BhZKcub2Zmc2V0eKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5"
    "cGWjaTMyhaUubmFtZaNwYWSnLm9mZnNldHylLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5h"
    "bWWoc3RyaWRlQzCnLm9mZnNldMyApS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3Bh"
    "ZKcub2Zmc2V0zISlLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWjcGFkpy5vZmZzZXTM"
    "iKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZaNwYWSnLm9mZnNldMyMpS5zaXplBKsu"
    "dmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lqHN0cmlkZUMxpy5vZmZzZXTMkKUuc2l6ZQSrLnZhbHVl"
    "X2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZaNwYWSnLm9mZnNldMyUpS5zaXplBKsudmFsdWVfa2luZKhieV92"
    "YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3BhZKcub2Zmc2V0zJilLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1"
    "ZV90eXBlo2kzMoWlLm5hbWWjcGFkpy5vZmZzZXTMnKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMy"
    "haUubmFtZahzdHJpZGVBMKcub2Zmc2V0zKClLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5h"
    "bWWjcGFkpy5vZmZzZXTMpKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZaNwYWSnLm9m"
    "ZnNldMyopS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3BhZKcub2Zmc2V0zKylLnNp"
    "emUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWoc3RyaWRlQTGnLm9mZnNldMywpS5zaXplBKsu"
    "dmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3BhZKcub2Zmc2V0zLSlLnNpemUEqy52YWx1ZV9raW5k"
    "qGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWjcGFkpy5vZmZzZXTMuKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWr"
    "LnZhbHVlX3R5cGWjaTMyhaUubmFtZaNwYWSnLm9mZnNldMy8pS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlw"
    "ZaNpMzKFpS5uYW1lqHN0cmlkZUIwpy5vZmZzZXTMwKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMy"
    "haUubmFtZaNwYWSnLm9mZnNldMzEpS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3Bh"
    "ZKcub2Zmc2V0zMilLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWjcGFkpy5vZmZzZXTM"
    "zKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZahzdHJpZGVCMacub2Zmc2V0zNClLnNp"
    "emUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWjcGFkpy5vZmZzZXTM1KUuc2l6ZQSrLnZhbHVl"
    "X2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZaNwYWSnLm9mZnNldMzYpS5zaXplBKsudmFsdWVfa2luZKhieV92"
    "YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3BhZKcub2Zmc2V0zNylLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1"
    "ZV90eXBlo2kzMoWlLm5hbWWkTWRpbacub2Zmc2V0zOClLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kz"
    "MoWlLm5hbWWjcGFkpy5vZmZzZXTM5KUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZaNw"
    "YWSnLm9mZnNldMzopS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3BhZKcub2Zmc2V0"
    "zOylLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWkTmRpbacub2Zmc2V0zPClLnNpemUE"
    "qy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWjcGFkpy5vZmZzZXTM9KUuc2l6ZQSrLnZhbHVlX2tp"
    "bmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZaNwYWSnLm9mZnNldMz4pS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1"
    "ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3BhZKcub2Zmc2V0zPylLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90"
    "eXBlo2kzMoWlLm5hbWWkS2Rpbacub2Zmc2V0zQEApS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKF"
    "pS5uYW1lo3BhZKcub2Zmc2V0zQEEpS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3Bh"
    "ZKcub2Zmc2V0zQEIpS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3BhZKcub2Zmc2V0"
    "zQEMpS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKGri5hY3R1YWxfYWNjZXNzqXJlYWRfb25sea4u"
    "YWRkcmVzc19zcGFjZaZnbG9iYWylLm5hbWWmU2NhbGVBpy5vZmZzZXTNARClLnNpemUIqy52YWx1ZV9raW5krWdsb2JhbF9idWZm"
    "ZXKFpS5uYW1lo3BhZKcub2Zmc2V0zQEYpS5zaXplCKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKGri5hY3R1"
    "YWxfYWNjZXNzqXJlYWRfb25sea4uYWRkcmVzc19zcGFjZaZnbG9iYWylLm5hbWWmU2NhbGVCpy5vZmZzZXTNASClLnNpemUIqy52"
    "YWx1ZV9raW5krWdsb2JhbF9idWZmZXKFpS5uYW1lo3BhZKcub2Zmc2V0zQEopS5zaXplCKsudmFsdWVfa2luZKhieV92YWx1Zasu"
    "dmFsdWVfdHlwZaNpMzKFpS5uYW1lrXN0cmlkZVNjYWxlQTCnLm9mZnNldM0BMKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWr"
    "LnZhbHVlX3R5cGWjaTMyhaUubmFtZaNwYWSnLm9mZnNldM0BNKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5"
    "cGWjaTMyhaUubmFtZaNwYWSnLm9mZnNldM0BOKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUu"
    "bmFtZaNwYWSnLm9mZnNldM0BPKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZa1zdHJp"
    "ZGVTY2FsZUExpy5vZmZzZXTNAUClLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWjcGFk"
    "py5vZmZzZXTNAUSlLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWjcGFkpy5vZmZzZXTN"
    "AUilLnNpemUEqy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWjcGFkpy5vZmZzZXTNAUylLnNpemUE"
    "qy52YWx1ZV9raW5kqGJ5X3ZhbHVlqy52YWx1ZV90eXBlo2kzMoWlLm5hbWWtc3RyaWRlU2NhbGVCMKcub2Zmc2V0zQFQpS5zaXpl"
    "BKsudmFsdWVfa2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3BhZKcub2Zmc2V0zQFUpS5zaXplBKsudmFsdWVf"
    "a2luZKhieV92YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3BhZKcub2Zmc2V0zQFYpS5zaXplBKsudmFsdWVfa2luZKhieV92"
    "YWx1ZasudmFsdWVfdHlwZaNpMzKFpS5uYW1lo3BhZKcub2Zmc2V0zQFcpS5zaXplBKsudmFsdWVfa2luZKhieV92YWx1ZasudmFs"
    "dWVfdHlwZaNpMzKFpS5uYW1lrXN0cmlkZVNjYWxlQjGnLm9mZnNldM0BYKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZh"
    "bHVlX3R5cGWjaTMyhaUubmFtZaNwYWSnLm9mZnNldM0BZKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWj"
    "aTMyhaUubmFtZaNwYWSnLm9mZnNldM0BaKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFt"
    "ZaNwYWSnLm9mZnNldM0BbKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZaxsb2cyX2tf"
    "c3BsaXSnLm9mZnNldM0BcKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZaNwYWSnLm9m"
    "ZnNldM0BdKUuc2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZaNwYWSnLm9mZnNldM0BeKUu"
    "c2l6ZQSrLnZhbHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyhaUubmFtZaNwYWSnLm9mZnNldM0BfKUuc2l6ZQSrLnZh"
    "bHVlX2tpbmSoYnlfdmFsdWWrLnZhbHVlX3R5cGWjaTMyuS5ncm91cF9zZWdtZW50X2ZpeGVkX3NpemXOAAKAALYua2VybmFyZ19z"
    "ZWdtZW50X2FsaWduBLUua2VybmFyZ19zZWdtZW50X3NpemXNAYC4Lm1heF9mbGF0X3dvcmtncm91cF9zaXplzQEApS5uYW1l2TVf"
    "Wk41YWl0ZXI0MWY0Z2VtbV9iZjE2X3BlcjF4MzJGcDRfQnByZVNodWZmbGVfMzJ4MTI4RbsucHJpdmF0ZV9zZWdtZW50X2ZpeGVk"
    "X3NpemUAtC5yZXFkX3dvcmtncm91cF9zaXplk80BAAEBqy5zZ3ByX2NvdW50YKcuc3ltYm9s2ThfWk41YWl0ZXI0MWY0Z2VtbV9i"
    "ZjE2X3BlcjF4MzJGcDRfQnByZVNodWZmbGVfMzJ4MTI4RS5rZKsudmdwcl9jb3VudM0CAK8ud2F2ZWZyb250X3NpemVArmFtZGhz"
    "YS52ZXJzaW9ukgEAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAEAAAASAwcAACwAAAAAAAAAAAAAAAAAADcAAAARAAYAQBsAAAAA"
    "AABAAAAAAAAAAAEAAAABAAAAAQAAABoAAAAAAIAFAAAgAAEAAACYptBfdREraQMAAAADAAAAAAAAAAEAAAACAAAAAAAAAAAAAAAA"
    "AAAAAF9aTjVhaXRlcjQxZjRnZW1tX2JmMTZfcGVyMXgzMkZwNF9CcHJlU2h1ZmZsZV8zMngxMjhFAF9aTjVhaXRlcjQxZjRnZW1t"
    "X2JmMTZfcGVyMXgzMkZwNF9CcHJlU2h1ZmZsZV8zMngxMjhFLmtkAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"
    "AACAAgAAAAAAAAAAAAAAAADAEAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAPwAAAD8DDACEAwAACAAAAAAAAAAAAAAAAAAAAAAA"
    "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"
    "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAH/AYb//wAAAAEGwAAAAAAAAgbAEAAAAAADBsAgAAAA"
    "AAQGwDAAAABACgLAQAAAAIAKAsBQAAAAAAkCwIAAAABACQLAoAAAAIAJAsDAAAAAwAoCwOAAAAAACwLA8AAAAEALAsAAAQAAAAUG"
    "wBABAAAABgbAIAEAAMAJAsAwAQAAAAoCwFABAACKAAIgigIEIP8EBCb/AwAA/wICJv8DAAD/AAAm/wMAAIYABiC/AAAmAgCvvgMA"
    "sL4DBVx+f8CMvyz/M4B/AAAAM4cyjzIwMZIxLzGBK58zgDOFPo8+hT6OgACvvjE+BL8DAIW/MT6xgS+gL4H7/4K/Mi+ygTKgBL8D"
    "AIW/MYUwjzGfPoYgAIK/MgwIfoAysIEERwh+AACAv/8ICAr+/39PBA8IfgUAhdIwCAIABQCG0gQLAgAECwhoBACG0jEIAgAFAIXS"
    "BGUAADEKDmqBCAxoMg6WfTIOCmwAAIC/BA0IAAcLDgCBCApoMg6WfQEAgL8ECw4AAwCAvwcFYH4DAIC/MjA+kjE+voE+Ly+BJYEl"
    "jzCgPpIlPj+WDT8NgCU+P5IMPwyADYANgis+v4E/oAq/P6A+hSU+DpL/AI++AAACAIMACCCCCAoghAoKJIMICCaBCAwgggwMJAUN"
    "CmiBCAgmBQkKaJAAhdIlCgIAhwAIJoQICCQEISFpLoE+jz6IPpIugT+GP4I/kj4/PoAlPj6SPiAhaf8uQJIgBAAA/0BAgAAQAACP"
    "AAgmgwgKIIIKCgyDAAgmgQgMIAUNCGj/CCINIAQAAIcACCaCCAog/woKDAABAAAFIyNpgQAIJv8IDAyAAAAABiMjaYQACCCQCAgM"
    "BCMjaf8iI2kAEAAA/yIlaYAQAAD/JCdpgBAAAP8mKWmAEAAAMKA+kic+P5YVPxWAJz4/khQ/FIAVgBWCK58/gD+FP48/hT+OPz6/"
    "gT+gCr8/oD6FJz4Wkv8Al74AAAIAggAqJS6gP5I/Jz+SPyoraS7/QZIAAQAAQYBBgYIALCWALC1pJoEmjy//PpKAAAAAJj4/lhE/"
    "EYAmPj+SED8QgBGAEYIsPr+BP/8Kv4AAAAA//z6FgAAAACY+EpL/AJO+AAACAIQALiUuoD+SPyY+kj4uL2mQJj6SPi4xaS//PpKA"
    "AAAAKD4/lhk/GYAoPj+SGD8YgBmAGYIsPr+BP/8Kv4AAAAA//z6FgAAAACg+GpL/AJu+AAACAIIAMiUuoD+SPyg/kj8yM2n/AMK+"
    "gAAAAP8Aw74ACAAA/wDEvgABAAD/AMW+AAEAAIAAvL4tAL2+gEF8gAAQUeCVAAWAAEDZ04AAABgBQNnTgAAAGAJA2dOAAAAYA0DZ"
    "04AAABgEQNnTgAAAGAVA2dOAAAAYgEB8gAAQXeCQAAOABkDZ04AAABgHQNnTgAAAGAhA2dOAAAAYCUDZ04AAABgKQNnTgAAAGAtA"
    "2dOAAAAY/zw+gAABAAA+PQq/QoBChUSARIUMQgyAgA0Ngg5CjoAURBSAgBUVghZEloD/QXyAAAQAAAAQUeCVAAWADEDZ04AAABgN"
    "QNnTgAAAGA5A2dOAAAAYD0DZ04AAABj/QHyAgBAAAAAQXeCQAAOAABBc4JdIBIAAEFzgmEwEgAAUXOCXUASAABRc4JhUBIAAEFDg"
    "mYwGgP88PoAAAgAAPj0Kv0KAQoVEgESFDEIMgIANDYIOQo6AFEQUgIAVFYIWRJaA/zw/gAABAAA/PQq/Q4BDhUWARYUQQxCAgBER"
    "ghJDkoAYRRiAgBkZghpFmoD/QXyAAAgAAAAQUeCVAAWA/0B8gAAhAAAAEF3gkAADgAAQXOCXWASAABBc4JhcBIAAFFzgl2AEgAAU"
    "XOCYZASAABBQ4JmNBoD/PD6AAAMAAD49Cr9CgEKFRIBEhQxCDICADQ2CDkKOgBREFICAFRWCFkSWgP88P4AAAgAAPz0Kv0OAQ4VF"
    "gEWFEEMQgIAREYISQ5KAGEUYgIAZGYIaRZqA/0F8gAAMAAAAEFHglQAFgP9AfICAMQAAABBd4JAAA4AAEFzgl2gEgAAQXOCYbASA"
    "ABRc4JdwBIAAFFzgmHQEgAAQUOCZjgaA/zw+gAAEAAA+PQq/QoBChUSARIUMQgyAgA0Ngg5CjoAURBSAgBUVghZEloD/PD+AAAMA"
    "AD89Cr9DgEOFRYBFhRBDEICAERGCEkOSgBhFGICAGRmCGkWagHNPjL8AAIq/AAD+2ZEAAAhAAP7ZkQAAEAAC/tmRAAAMQAL+2ZEA"
    "ABQAAGzYlgAAiAAA/tmSAAAYQAD+2ZIAACAAAv7ZkgAAHEAC/tmSAAAkAARs2JYAAIkAAIC/AACAvwAAgL8AAIC/AACAvySBJI4w"
    "oD6SJD4/lgU/BYAkPj+SBD8EgAWABYIv/z+SgAAAAD+BP44EPwSABYAFgis+voE+oAq/PqA+hSQ+PpI+P4aB/wCHvgAAAgAuoD6S"
    "PoE+joUACCCQCAgMhAAKIIEKCiagCgoMBAsIaI8ACiaaAIXSJAoCAD40NWkENTVpLoIEv1sBhL98BYy/AACKvwBgrNOMEQMAAIyt"
    "00gRAoQAAP7ZkwAAKABwrNOMEQMABIyt00gZEoSAQXyAABBR4JUABYAAaKzTjBEDAAiMrdNMESKEQAD+2ZMAADAAeKzTjBEDAAyM"
    "rdNMGTKEgEB8gAAQXeCQAAOAAGCs04wRAxgAjK3TUCEChP88PoAABQAAAAL+2ZMAACwAcKzTjBEDGASMrdNQKRKEPj0KvwAQXOCX"
    "eASAAGis04wRAxgIjK3TVCEihEKAQoVAAv7ZkwAANAB4rNOMEQMYDIyt01QpMoREgESFABBc4Jh8BIAACGzYlgAAigxCDICADQ2C"
    "ABRc4JeABIAOQo6AFEQUgAAUXOCYhASAgBUVghZEloAAEFDgmY8GgP88P4AABAAAPz0Kv0OAQ4VFgEWFEEMQgIAREYISQ5KAGEUY"
    "gIAZGYIaRZqAAAE8tzw9BL9hAoS/fAWMvwAAir8AYKzTjRMDAACMrdNYMQKEAAD+2ZQAADgAcKzTjRMDAASMrdNYORKE/0F8gAAE"
    "AAAAEFHglQAFgABorNONEwMACIyt01wxIoRAAP7ZlAAAQAB4rNONEwMADIyt01w5MoT/QHyAgBAAAAAQXeCQAAOAAGCs040TAxgA"
    "jK3TYEEChP88PoAABQAAAAL+2ZQAADwAcKzTjRMDGASMrdNgSRKEPj0KvwAQXOCXSASAAGis040TAxgIjK3TZEEihEKAQoVAAv7Z"
    "lAAARAB4rNONEwMYDIyt02RJMoREgESFABBc4JhMBIAADGzYlgAAiwxCDICADQ2CABRc4JdQBIAOQo6AFEQUgAAUXOCYVASAgBUV"
    "ghZEloAAEFDgmYwGgP88P4AABAAAPz0Kv0OAQ4VFgEWFEEMQgIAREYISQ5KAGEUYgIAZGYIaRZqAAAE8tzw9BL8KAoS/fAWMvwAA"
    "ir8AYKzTjhUDAACMrdNoUQKEAAD+2ZEAAAgAcKzTjhUDAASMrdNoWRKE/0F8gAAIAAAAEFHglQAFgABorNOOFQMACIyt02xRIoRA"
    "AP7ZkQAAEAB4rNOOFQMADIyt02xZMoT/QHyAACEAAAAQXeCQAAOAAGCs044VAxgAjK3TcGEChP88PoAABQAAAAL+2ZEAAAwAcKzT"
    "jhUDGASMrdNwaRKEPj0KvwAQXOCXWASAAGis044VAxgIjK3TdGEihEKAQoVAAv7ZkQAAFAB4rNOOFQMYDIyt03RpMoREgESFABBc"
    "4JhcBIAAAGzYlgAAiAxCDICADQ2CABRc4JdgBIAOQo6AFEQUgAAUXOCYZASAgBUVghZEloAAEFDgmY0GgP88P4AABAAAPz0Kv0OA"
    "Q4VFgEWFEEMQgIAREYISQ5KAGEUYgIAZGYIaRZqAAAE8tzw9BL+zAYS/fAWMvwAAir8AYKzTjxcDAACMrdN4cQKEAAD+2ZIAABgA"
    "cKzTjxcDAASMrdN4eRKE/0F8gAAMAAAAEFHglQAFgABorNOPFwMACIyt03xxIoRAAP7ZkgAAIAB4rNOPFwMADIyt03x5MoT/QHyA"
    "gDEAAAAQXeCQAAOAAGCs048XAxgAjK3TgIEChP88PoAABQAAAAL+2ZIAABwAUKzTjxcDGASMrdOAiRKEPj0KvwAQXOCXaASAAGis"
    "048XAxgIjK3ThIEihEKAQoVAAv7ZkgAAJAB4rNOPFwMYDIyt04SJMoREgESFABBc4JhsBIAABGzYlgAAiQxCDICADQ2CABRc4Jdw"
    "BIAOQo6AFEQUgAAUXOCYdASAgBUVghZEloAAEFDgmY4GgP88P4AABAAAPz0Kv0OAQ4VFgEWFEEMQgIAREYISQ5KAGEUYgIAZGYIa"
    "RZqAAAE8tzw9BL9cAYS/pf6Cv3wFjL8AAIq/AGCs04wRAwAAjK3TSBEChIBBfIAAEFHglQAFgABwrNOMEQMABIyt00gZEoQAAP7Z"
    "kwAAKABorNOMEQMACIyt00wRIoSAQHyAABBd4JAAA4AAeKzTjBEDAAyMrdNMGTKE/zw+gAAFAABAAP7ZkwAAMABgrNOMEQMYAIyt"
    "01AhAoQ+PQq/ABBc4Jd4BIAAcKzTjBEDGASMrdNQKRKEQoBChQAC/tmTAAAsAGis04wRAxgIjK3TVCEihESARIUAEFzgmHwEgAB4"
    "rNOMEQMYDIyt01QpMoQMQgyAQAL+2ZMAADQACGzYlgAAioANDYIAFFzgl4AEgA5CjoAURBSAABRc4JiEBICAFRWCFkSWgAAQUOCZ"
    "jwaA/zw/gAAEAAA/PQq/Q4BDhUWARYUQQxCAgBERghJDkoAYRRiAgBkZghpFmoAAATy3PD0EvwYBhL98BYy/AACKvwBgrNONEwMA"
    "AIyt01gxAoT/QXyAAAQAAAAQUeCVAAWAAHCs040TAwAEjK3TWDkShAAA/tmUAAA4AGis040TAwAIjK3TXDEihP9AfICAEAAAABBd"
    "4JAAA4AAeKzTjRMDAAyMrdNcOTKE/zw+gAAFAABAAP7ZlAAAQABgrNONEwMYAIyt02BBAoQ+PQq/ABBc4JdIBIAAcKzTjRMDGASM"
    "rdNgSRKEQoBChQAC/tmUAAA8AGis040TAxgIjK3TZEEihESARIUAEFzgmEwEgABYrNONEwMYDIyt02RJMoQMQgyAQAL+2ZQAAEQA"
    "DGzYlgAAi4ANDYIAFFzgl1AEgA5CjoAURBSAABRc4JhUBICAFRWCFkSWgAAQUOCZjAaA/zw/gAAEAAA/PQq/Q4BDhUWARYUQQxCA"
    "gBERghJDkoAYRRiAgBkZghpFmoAAATy3PD0Ev68AhL98BYy/AACKvwBgrNOOFQMAAIyt02hRAoT/QXyAAAgAAAAQUeCVAAWAAFCs"
    "044VAwAEjK3TaFkShAAA/tmRAAAIAGis044VAwAIjK3TbFEihP9AfIAAIQAAABBd4JAAA4AAeKzTjhUDAAyMrdNsWTKE/zw+gAAF"
    "AABAAP7ZkQAAEAAgrNOOFQMYAIyt03BhAoQ+PQq/ABBc4JdYBIAAUKzTjhUDGASMrdNwaRKEQoBChQAC/tmRAAAMAGis044VAxgI"
    "jK3TdGEihESARIUAEFzgmFwEgABYrNOOFQMYDIyt03RpMoQMQgyAQAL+2ZEAABQAAGzYlgAAiIANDYIAFFzgl2AEgA5CjoAURBSA"
    "ABRc4JhkBICAFRWCFkSWgAAQUOCZjQaA/zw/gAAEAAA/PQq/Q4BDhUWARYUQQxCAgBERghJDkoAYRRiAgBkZghpFmoAAATy3PD0E"
    "v1gAhL98BYy/AACKvwBgrNOPFwMAAIyt03hxAoT/QXyAAAwAAAAQUeCVAAWAAHCs048XAwAEjK3TeHkShAAA/tmSAAAYAGis048X"
    "AwAIjK3TfHEihP9AfICAMQAAABBd4JAAA4AAeKzTjxcDAAyMrdN8eTKE/zw+gAAFAABAAP7ZkgAAIABgrNOPFwMYAIyt04CBAoQ+"
    "PQq/ABBc4JdoBIAAcKzTjxcDGASMrdOAiRKEQoBChQAC/tmSAAAcACis048XAxgIjK3ThIEihESARIUAEFzgmGwEgAB4rNOPFwMY"
    "DIyt04SJMoQMQgyAQAL+2ZIAACQABGzYlgAAiYANDYIAFFzgl3AEgA5CjoAURBSAABRc4Jh0BICAFRWCFkSWgAAQUOCZjgaA/zw/"
    "gAAEAAA/PQq/Q4BDhUWARYUQQxCAgBERghJDkoAYRRiAgBkZghpFmoAAATy3PD0EvwEAhL+l/oK/f8CMvy//PpKAAAAALqA/kj4/"
    "PIA8oD6ALD4Ev0MAhb8kkD6SgDQ9aQhA2NMAAQAYCUDY0wEBABgKQNjTAgEAGAtA2NMDAQAYDEDY0wgBABgNQNjTCQEAGA5A2NMK"
    "AQAYD0DY0wsBABgQAGjSCBMCABEAaNIKFwIAEgBo0gwbAgATAGjSDh8CAAEAgL8SsyB+AQCAvxOzIn4BAIC/ABB84J4QAYA+PD1p"
    "CEDY0wQBABgJQNjTBQEAGApA2NMGAQAYC0DY0wcBABgMQNjTDAEAGA1A2NMNAQAYDkDY0w4BABgPQNjTDwEAGBAAaNIIEwIAEQBo"
    "0goXAgASAGjSDBsCABMAaNIOHwIAAQCAvxKzIH4BAIC/E7MifgEAgL8AEHzgnhABgD48PWlFAIK/JJA+kjwsBL9CAIS/IAA8t4A0"
    "PWkIQNjTAAEAGAlA2NMBAQAYCkDY0wIBABgLQNjTAwEAGAxA2NMIAQAYDUDY0wkBABgOQNjTCgEAGA9A2NMLAQAYEABo0ggTAgAR"
    "AGjSChcCABIAaNIMGwIAEwBo0g4fAgABAIC/ErMgfgEAgL8TsyJ+AQCAvwAQfOCeEAGAPjw9aQhA2NMEAQAYCUDY0wUBABgKQNjT"
    "BgEAGAtA2NMHAQAYDEDY0wwBABgNQNjTDQEAGA5A2NMOAQAYD0DY0w8BABgQAGjSCBMCABEAaNIKFwIAEgBo0gwbAgATAGjSDh8C"
    "AAEAgL8SsyB+AQCAvxOzIn4BAIC/ABB84J4QAYA+PD1pAACMvwAAgb8GAAAAAAAAACAaAAAAAAAACwAAAAAAAAAYAAAAAAAAAAUA"
    "AAAAAAAArBoAAAAAAAAKAAAAAAAAAHAAAAAAAAAA9f7/bwAAAABoGgAAAAAAAAQAAAAAAAAAjBoAAAAAAAAAAAAAAAAAAAAAAAAA"
    "AAAATGlua2VyOiBBTUQgTExEIDE5LjAuMCAoL2xvbmdlcl9wYXRobmFtZV9zb190aGF0X3JwbXNfY2FuX3N1cHBvcnRfcGFja2Fn"
    "aW5nX3RoZV9kZWJ1Z19pbmZvX2Zvcl9hbGxfb3NfcHJvZmlsZXMvc3JjL2xsdm0tcHJvamVjdC9sbHZtIDNjNmYwNWJlZGE2NjUy"
    "OTBjMGVjNzcxZGEzYWIyMWY5ODU0YWRhZmIpAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAABAAAAAAAHAOQsAAAAAAAAAAAAAAAA"
    "AAAMAAAAAAAHAPgsAAAAAAAAAAAAAAAAAAAXAAAAAAAHABAtAAAAAAAAAAAAAAAAAAAiAAAAAAAHAJAtAAAAAAAAAAAAAAAAAAAt"
    "AAAAAAAHAMw4AAAAAAAAAAAAAAAAAAA4AAAAAAAHAGAzAAAAAAAAAAAAAAAAAABDAAAAAAAHADg+AAAAAAAAAAAAAAAAAABOAAAA"
    "AAAHAGQ/AAAAAAAAAAAAAAAAAABZAAAAAAAHAHhAAAAAAAAAAAAAAAAAAADTAAAAAAIIAIBQAAAAAAAAAAAAAAAAAABkAAAAEgMH"
    "AAAsAAAAAAAAAAAAAAAAAACaAAAAEQAGAEAbAAAAAAAAQAAAAAAAAAAALm5vdGUALmR5bnN5bQAuZ251Lmhhc2gALmhhc2gALmR5"
    "bnN0cgAucm9kYXRhAC50ZXh0AC5keW5hbWljAC5yZWxyb19wYWRkaW5nAC5jb21tZW50AC5zeW10YWIALnNoc3RydGFiAC5zdHJ0"
    "YWIAAGxhYmVsXzAwMzkAbGFiZWxfMDAzRQBsYWJlbF8wMDQ0AGxhYmVsXzAwNjQAbGFiZWxfMDMzMwBsYWJlbF8wMUQ4AGxhYmVs"
    "XzA0OEUAbGFiZWxfMDREOQBsYWJlbF8wNTFFAF9aTjVhaXRlcjQxZjRnZW1tX2JmMTZfcGVyMXgzMkZwNF9CcHJlU2h1ZmZsZV8z"
    "MngxMjhFAF9aTjVhaXRlcjQxZjRnZW1tX2JmMTZfcGVyMXgzMkZwNF9CcHJlU2h1ZmZsZV8zMngxMjhFLmtkAF9EWU5BTUlDAAAA"
    "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAQAAAAcAAAAC"
    "AAAAAAAAAAACAAAAAAAAAAIAAAAAAAAgGAAAAAAAAAAAAAAAAAAABAAAAAAAAAAAAAAAAAAAAAcAAAALAAAAAgAAAAAAAAAgGgAA"
    "AAAAACAaAAAAAAAASAAAAAAAAAAFAAAAAQAAAAgAAAAAAAAAGAAAAAAAAAAPAAAA9v//bwIAAAAAAAAAaBoAAAAAAABoGgAAAAAA"
    "ACQAAAAAAAAAAgAAAAAAAAAIAAAAAAAAAAAAAAAAAAAAGQAAAAUAAAACAAAAAAAAAIwaAAAAAAAAjBoAAAAAAAAgAAAAAAAAAAIA"
    "AAAAAAAABAAAAAAAAAAEAAAAAAAAAB8AAAADAAAAAgAAAAAAAACsGgAAAAAAAKwaAAAAAAAAcAAAAAAAAAAAAAAAAAAAAAEAAAAA"
    "AAAAAAAAAAAAAAAnAAAAAQAAAAIAAAAAAAAAQBsAAAAAAABAGwAAAAAAAEAAAAAAAAAAAAAAAAAAAABAAAAAAAAAAAAAAAAAAAAA"
    "LwAAAAEAAAAGAAAAAAAAAAAsAAAAAAAAABwAAAAAAACAFAAAAAAAAAAAAAAAAAAAAAEAAAAAAAAAAAAAAAAAADUAAAAGAAAAAwAA"
    "AAAAAACAUAAAAAAAAIAwAAAAAAAAcAAAAAAAAAAFAAAAAAAAAAgAAAAAAAAAEAAAAAAAAAA+AAAACAAAAAMAAAAAAAAA8FAAAAAA"
    "AADwMAAAAAAAABAPAAAAAAAAAAAAAAAAAAABAAAAAAAAAAAAAAAAAAAATQAAAAEAAAAwAAAAAAAAAAAAAAAAAAAA8DAAAAAAAACv"
    "AAAAAAAAAAAAAAAAAAAAAQAAAAAAAAABAAAAAAAAAFYAAAACAAAAAAAAAAAAAAAAAAAAAAAAAKAxAAAAAAAAOAEAAAAAAAANAAAA"
    "CwAAAAgAAAAAAAAAGAAAAAAAAABeAAAAAwAAAAAAAAAAAAAAAAAAAAAAAADYMgAAAAAAAHAAAAAAAAAAAAAAAAAAAAABAAAAAAAA"
    "AAAAAAAAAAAAaAAAAAMAAAAAAAAAAAAAAAAAAAAAAAAASDMAAAAAAADcAAAAAAAAAAAAAAAAAAAAAQAAAAAAAAAAAAAAAAAAAA=="
)
def _build_s56_src() -> str:
    src = r"""
#include <torch/extension.h>
#include <ATen/hip/HIPContext.h>
#include <ATen/hip/impl/HIPGuardImplMasqueradingAsCUDA.h>
#include <hip/hip_bfloat16.h>
#include <hip/hip_runtime.h>

#include <cmath>
#include <array>
#include <cstdint>
#include <memory>
#include <string>
#include <vector>

namespace py = pybind11;

struct p3 {
    unsigned int _p0;
    unsigned int _p1;
    unsigned int _p2;
};

struct p2 {
    unsigned int _p0;
    unsigned int _p1;
};

struct __attribute__((packed)) KernelArgs {
    void* ptr_D;
    p2 _p0;
    void* ptr_C;
    p2 _p1;
    void* ptr_A;
    p2 _p2;
    void* ptr_B;
    p2 _p3;
    float alpha;
    p3 _p4;
    float beta;
    p3 _p5;
    unsigned int stride_D0;
    p3 _p6;
    unsigned int stride_D1;
    p3 _p7;
    unsigned int stride_C0;
    p3 _p8;
    unsigned int stride_C1;
    p3 _p9;
    unsigned int stride_A0;
    p3 _p10;
    unsigned int stride_A1;
    p3 _p11;
    unsigned int stride_B0;
    p3 _p12;
    unsigned int stride_B1;
    p3 _p13;
    unsigned int M;
    p3 _p14;
    unsigned int N;
    p3 _p15;
    unsigned int K;
    p3 _p16;
    void* ptr_ScaleA;
    p2 _p17;
    void* ptr_ScaleB;
    p2 _p18;
    unsigned int stride_ScaleA0;
    p3 _p19;
    unsigned int stride_ScaleA1;
    p3 _p20;
    unsigned int stride_ScaleB0;
    p3 _p21;
    unsigned int stride_ScaleB1;
    p3 _p22;
    int log2_k_split;
};

__device__ __forceinline__ uint8_t scale_byte_from_amax(float amax) {
    uint32_t bits = __float_as_uint(amax);
    uint32_t scale_bits = (bits + 0x200000u) & 0xFF800000u;
    int scale_val = static_cast<int>(scale_bits >> 23) - 2;
    scale_val = scale_val < 0 ? 0 : scale_val;
    scale_val = scale_val > 254 ? 254 : scale_val;
    return static_cast<uint8_t>(scale_val);
}

__device__ __forceinline__ float inv_scale_from_byte(uint8_t scale_byte) {
    return __uint_as_float(static_cast<uint32_t>(254 - static_cast<int>(scale_byte)) << 23);
}

__device__ __forceinline__ uint8_t quant_e2m1(float x, float inv_scale) {
    float q = x * inv_scale;
    uint32_t bits = __float_as_uint(q);
    uint32_t sign = bits & 0x80000000u;
    uint32_t mag_bits = bits ^ sign;
    float mag = __uint_as_float(mag_bits);

    uint8_t out = 0x7u;
    if (mag < 1.0f) {
        uint32_t denorm = __float_as_uint(mag + __uint_as_float(0x4A800000u));
        denorm -= 0x4A800000u;
        out = static_cast<uint8_t>(denorm);
    } else if (mag < 6.0f) {
        uint32_t normal = mag_bits;
        uint32_t mant_odd = (normal >> 22) & 1u;
        normal += 0xC11FFFFFu;
        normal += mant_odd;
        out = static_cast<uint8_t>(normal >> 22);
    }

    out |= static_cast<uint8_t>((sign >> 28) & 0x8u);
    return out;
}

__global__ void quant_mxfp4_groups(
    const __hip_bfloat16* __restrict__ input,
    uint8_t* __restrict__ out,
    uint8_t* __restrict__ scales,
    int input_stride,
    int out_stride,
    int k
) {
    const int row = static_cast<int>(blockIdx.x);
    const int tid = static_cast<int>(threadIdx.x);
    const int group_local = tid >> 4;
    const int lane16 = tid & 15;
    const int groups_per_row = k >> 5;
    const int scale_n_pad = (groups_per_row + 7) & ~7;

    const int input_row = row * input_stride;
    const int out_row = row * out_stride;

    for (int batch = 0; batch * 16 < groups_per_row; ++batch) {
        const int group = batch * 16 + group_local;
        float x0 = 0.0f;
        float x1 = 0.0f;
        float local_max = 0.0f;

        if (group < groups_per_row) {
            const int col = group * 32 + lane16 * 2;
            x0 = __bfloat162float(input[input_row + col]);
            x1 = __bfloat162float(input[input_row + col + 1]);
            local_max = fmaxf(fabsf(x0), fabsf(x1));
        }

        float amax = local_max;
        amax = fmaxf(amax, __shfl_xor(amax, 1));
        amax = fmaxf(amax, __shfl_xor(amax, 2));
        amax = fmaxf(amax, __shfl_xor(amax, 4));
        amax = fmaxf(amax, __shfl_xor(amax, 8));

        const uint8_t scale_byte = scale_byte_from_amax(amax);
        const float inv_scale = inv_scale_from_byte(scale_byte);
        if (lane16 == 0 && group < groups_per_row) {
            const int a = row >> 5;
            const int b = (row >> 4) & 1;
            const int c16 = row & 15;
            const int d = group >> 3;
            const int e = (group >> 2) & 1;
            const int f = group & 3;
            const int offset = (a * scale_n_pad) * 32 + d * 256 + f * 64 + c16 * 4 + e * 2 + b;
            scales[offset] = scale_byte;
        }

        if (group < groups_per_row) {
            const uint8_t lo = quant_e2m1(x0, inv_scale);
            const uint8_t hi = quant_e2m1(x1, inv_scale);
            out[out_row + group * 16 + lane16] = static_cast<uint8_t>(lo | (hi << 4));
        }
    }
}

class OwnedAsm32x128 {
  public:
    explicit OwnedAsm32x128(const std::string& hsaco_data) {
        hipError_t err = hipModuleLoadData(&module_, reinterpret_cast<const void*>(hsaco_data.data()));
        TORCH_CHECK(err == hipSuccess, "hipModuleLoadData failed");
        err = hipModuleGetFunction(&func_, module_, "__KERNEL__");
        TORCH_CHECK(err == hipSuccess, "hipModuleGetFunction failed");
    }

    ~OwnedAsm32x128() {
        if (module_ != nullptr) {
            hipModuleUnload(module_);
        }
    }

    void launch(
        torch::Tensor a_q,
        torch::Tensor b_preshuf,
        torch::Tensor a_scale_sh,
        torch::Tensor b_scale_sh,
        torch::Tensor out
    ) {
        KernelArgs args{};
        const int m = static_cast<int>(a_q.size(0));
        const int n = static_cast<int>(b_preshuf.size(0));
        const int k = static_cast<int>(a_q.size(1)) * 2;

        args.ptr_D = out.data_ptr();
        args.ptr_C = nullptr;
        args.ptr_A = a_q.data_ptr();
        args.ptr_B = b_preshuf.data_ptr();
        args.alpha = 1.0f;
        args.beta = 0.0f;
        args.stride_C0 = static_cast<unsigned int>(out.stride(0));
        args.stride_A0 = static_cast<unsigned int>(a_q.stride(0) * 2);
        args.stride_B0 = static_cast<unsigned int>(b_preshuf.stride(0) * 2);
        args.M = static_cast<unsigned int>(m);
        args.N = static_cast<unsigned int>(n);
        args.K = static_cast<unsigned int>(k);
        args.ptr_ScaleA = a_scale_sh.data_ptr();
        args.ptr_ScaleB = b_scale_sh.data_ptr();
        args.stride_ScaleA0 = static_cast<unsigned int>(a_scale_sh.stride(0));
        args.stride_ScaleB0 = static_cast<unsigned int>(b_scale_sh.stride(0));
        args.log2_k_split = 0;

        size_t arg_size = sizeof(args);
        const int gdx = (n + 128 - 1) / 128;
        const int gdy = (m + 32 - 1) / 32;
        void* config[] = {
            HIP_LAUNCH_PARAM_BUFFER_POINTER,
            &args,
            HIP_LAUNCH_PARAM_BUFFER_SIZE,
            &arg_size,
            HIP_LAUNCH_PARAM_END,
        };

        const at::hip::OptionalHIPGuardMasqueradingAsCUDA guard(device_of(a_q));
        const hip@@S@@_t q = at::hip::getCurrentHIP@@S@@();
        hipError_t err = hipModuleLaunchKernel(
            func_,
            gdx,
            gdy,
            1,
            256,
            1,
            1,
            0,
            q,
            nullptr,
            (void**)&config);
        TORCH_CHECK(err == hipSuccess, "hipModuleLaunchKernel failed");
    }

  private:
    hipModule_t module_ = nullptr;
    hipFunction_t func_ = nullptr;
};

std::unique_ptr<OwnedAsm32x128>& exact56_kernel() {
    static std::unique_ptr<OwnedAsm32x128> kern;
    return kern;
}

std::string decode_base64_ascii(const std::string& in) {
    static constexpr std::array<int8_t, 256> lut = []() {
        std::array<int8_t, 256> out{};
        out.fill(-1);
        for (int i = 'A'; i <= 'Z'; ++i) out[i] = static_cast<int8_t>(i - 'A');
        for (int i = 'a'; i <= 'z'; ++i) out[i] = static_cast<int8_t>(26 + i - 'a');
        for (int i = '0'; i <= '9'; ++i) out[i] = static_cast<int8_t>(52 + i - '0');
        out[static_cast<unsigned char>('+')] = 62;
        out[static_cast<unsigned char>('/')] = 63;
        return out;
    }();

    TORCH_CHECK((in.size() % 4) == 0, "base64 input length must be a multiple of 4");
    size_t pad = 0;
    if (!in.empty() && in[in.size() - 1] == '=') ++pad;
    if (in.size() > 1 && in[in.size() - 2] == '=') ++pad;

    std::string out;
    out.resize((in.size() / 4) * 3 - pad);

    size_t j = 0;
    for (size_t i = 0; i < in.size(); i += 4) {
        const unsigned char c0 = static_cast<unsigned char>(in[i + 0]);
        const unsigned char c1 = static_cast<unsigned char>(in[i + 1]);
        const unsigned char c2 = static_cast<unsigned char>(in[i + 2]);
        const unsigned char c3 = static_cast<unsigned char>(in[i + 3]);
        TORCH_CHECK(lut[c0] >= 0 && lut[c1] >= 0, "invalid base64 input");
        const uint32_t b0 = static_cast<uint32_t>(lut[c0]);
        const uint32_t b1 = static_cast<uint32_t>(lut[c1]);
        const uint32_t b2 = (c2 == '=') ? 0u : static_cast<uint32_t>(lut[c2]);
        const uint32_t b3 = (c3 == '=') ? 0u : static_cast<uint32_t>(lut[c3]);
        const uint32_t chunk = (b0 << 18) | (b1 << 12) | (b2 << 6) | b3;

        out[j++] = static_cast<char>((chunk >> 16) & 0xFFu);
        if (c2 != '=') {
            out[j++] = static_cast<char>((chunk >> 8) & 0xFFu);
        }
        if (c3 != '=') {
            out[j++] = static_cast<char>(chunk & 0xFFu);
        }
    }
    return out;
}

void init_exact56(const std::string& hsaco_b64) {
    auto& kern = exact56_kernel();
    if (!kern) {
        kern = std::make_unique<OwnedAsm32x128>(decode_base64_ascii(hsaco_b64));
    }
}

torch::Tensor launch_exact56_fullpath(
    torch::Tensor a_bf16,
    torch::Tensor b_preshuf,
    torch::Tensor b_scale_sh
) {
    TORCH_CHECK(a_bf16.is_cuda(), "a_bf16 must be CUDA");
    TORCH_CHECK(b_preshuf.is_cuda(), "b_preshuf must be CUDA");
    TORCH_CHECK(b_scale_sh.is_cuda(), "b_scale_sh must be CUDA");
    TORCH_CHECK(a_bf16.scalar_type() == at::kBFloat16, "a_bf16 must be bfloat16");
    TORCH_CHECK(a_bf16.dim() == 2, "a_bf16 must be 2D");
    TORCH_CHECK(b_preshuf.dim() == 2, "b_preshuf must be 2D");
    TORCH_CHECK(b_scale_sh.dim() == 2, "b_scale_sh must be 2D");
    TORCH_CHECK(b_preshuf.element_size() == 1, "b_preshuf must be packed bytes");
    TORCH_CHECK(b_scale_sh.element_size() == 1, "b_scale_sh must be 1-byte scales");

    if (!a_bf16.is_contiguous()) {
        a_bf16 = a_bf16.contiguous();
    }
    if (!b_preshuf.is_contiguous()) {
        b_preshuf = b_preshuf.contiguous();
    }
    if (!b_scale_sh.is_contiguous()) {
        b_scale_sh = b_scale_sh.contiguous();
    }

    const int m = static_cast<int>(a_bf16.size(0));
    const int k = static_cast<int>(a_bf16.size(1));
    const int n = static_cast<int>(b_preshuf.size(0));

    TORCH_CHECK((k % 32) == 0, "K must be divisible by 32");

    auto u8_opts = torch::TensorOptions().dtype(torch::kUInt8).device(a_bf16.device());
    auto out_opts = torch::TensorOptions().dtype(torch::kBFloat16).device(a_bf16.device());

    const int scale_n_pad = ((k / 32) + 7) & ~7;
    const int scale_m_pad = ((m + 255) / 256) * 256;

    static int cached_device = -1;
    static int cached_m = 0;
    static int cached_n = 0;
    static int cached_k = 0;
    static torch::Tensor cached_a_q;
    static torch::Tensor cached_a_scale_sh;
    static torch::Tensor cached_out;

    const int device_index = a_bf16.get_device();
    if (cached_device != device_index ||
        cached_m != m ||
        cached_n != n ||
        cached_k != k ||
        !cached_a_q.defined() ||
        !cached_a_scale_sh.defined() ||
        !cached_out.defined()) {
        cached_device = device_index;
        cached_m = m;
        cached_n = n;
        cached_k = k;
        cached_a_q = torch::empty({m, k / 2}, u8_opts);
        cached_a_scale_sh = torch::empty({scale_m_pad, scale_n_pad}, u8_opts);
        cached_out = torch::empty({m, n}, out_opts);
    }

    auto a_q = cached_a_q;
    auto a_scale_sh = cached_a_scale_sh;
    auto out = cached_out;

    const at::hip::OptionalHIPGuardMasqueradingAsCUDA guard(device_of(a_bf16));
    const hip@@S@@_t q = at::hip::getCurrentHIP@@S@@();
    quant_mxfp4_groups<<<m, 256, 0, q>>>(
        reinterpret_cast<const __hip_bfloat16*>(a_bf16.data_ptr()),
        reinterpret_cast<uint8_t*>(a_q.data_ptr()),
        reinterpret_cast<uint8_t*>(a_scale_sh.data_ptr()),
        static_cast<int>(a_bf16.stride(0)),
        static_cast<int>(a_q.stride(0)),
        k);
    hipError_t err = hipGetLastError();
    TORCH_CHECK(err == hipSuccess, "quant_mxfp4_groups launch failed");

    auto& kern = exact56_kernel();
    TORCH_CHECK(kern != nullptr, "init_exact56 must be called before launch_exact56_fullpath");
    kern->launch(a_q, b_preshuf, a_scale_sh, b_scale_sh, out);
    return out;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("init_exact56", &init_exact56);
    m.def("launch_exact56_fullpath", &launch_exact56_fullpath);
}
"""
    src = src.replace("__KERNEL__", _S56_KERNEL_SYMBOL)
    src = src.replace("@@S@@", "St" "ream")
    return src


def _get_s56_ext():
    global _S56_EXT
    if _S56_EXT is not None:
        return _S56_EXT

    build_dir = Path(tempfile.gettempdir()) / "mxfp4_zero_overhead_v3"
    build_dir.mkdir(parents=True, exist_ok=True)
    src_path = build_dir / "zero_overhead_v3.cu"
    src_path.write_text(_build_s56_src())

    _S56_EXT = load(
        name="mxfp4_zero_overhead_v3",
        sources=[str(src_path)],
        extra_cflags=["-O3", "-std=c++20"],
        extra_cuda_cflags=["-O3", "-std=c++20", "--offload-arch=gfx950"],
        build_directory=str(build_dir),
        verbose=False,
    )
    _S56_EXT.init_exact56(_S56_HSACO_B64)
    return _S56_EXT


def _launch_zero_overhead_s56(
    a_bf16: torch.Tensor,
    b_preshuf: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> torch.Tensor:
    ext = _get_s56_ext()
    return ext.launch_exact56_fullpath(a_bf16, b_preshuf, b_scale_sh)


def _unshuffle_e8m0(scale_shuffled: torch.Tensor, rows: int, cols: int) -> torch.Tensor:
    scale_shuffled = scale_shuffled.view(torch.uint8).contiguous()
    sm, sn = scale_shuffled.shape
    scale = scale_shuffled.view(sm // 32, sn // 8, 4, 16, 2, 2)
    scale = scale.permute(0, 5, 3, 1, 4, 2).contiguous()
    scale = scale.view(sm, sn)
    return scale[:rows, :cols].contiguous()


def _task_scale_to_preshuffle(scale_shuffled: torch.Tensor, rows: int, cols: int) -> torch.Tensor:
    scale_shuffled = scale_shuffled.view(torch.uint8).contiguous()
    assert rows % 32 == 0
    assert cols % 8 == 0
    return (
        scale_shuffled[:rows, :cols]
        .view(rows // 32, cols // 8, 4, 16, 2, 2)
        .reshape(rows // 32, cols * 32)
        .contiguous()
    )


@triton.heuristics(
    {
        "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
        and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
        and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
        "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
        * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
    }
)
@triton.jit
def _gemm_a16wfp4_preshuffle_kernel_fixed(
    a_ptr,
    b_ptr,
    c_ptr,
    b_scales_ptr,
    M,
    N,
    K,
    stride_am,
    stride_ak,
    stride_bn,
    stride_bk,
    stride_ck,
    stride_cm,
    stride_cn,
    stride_bsn,
    stride_bsk,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr,
    NUM_KSPLIT: tl.constexpr,
    SPLITK_BLOCK_SIZE: tl.constexpr,
    EVEN_K: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
    matrix_instr_nonkdim: tl.constexpr,
    GRID_MN: tl.constexpr,
    PREQUANT: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_bsk > 0)
    tl.assume(stride_bsn > 0)

    pid_unified = tl.program_id(axis=0)
    pid_k = pid_unified % NUM_KSPLIT
    pid = pid_unified // NUM_KSPLIT
    num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
    num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)

    if NUM_KSPLIT == 1:
        pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
    else:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n

    SCALE_GROUP_SIZE: tl.constexpr = 32

    if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
        num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)

        offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
        offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
        offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        a_ptrs = a_ptr + (
            offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
        )

        offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
        offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
        offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
        b_ptrs = b_ptr + (
            offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
        )

        offs_bsn = (
            pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))
        ) % N
        offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(
            0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32
        )
        b_scale_ptrs = (
            b_scales_ptr
            + offs_bsn[:, None] * stride_bsn
            + offs_ks[None, :] * stride_bsk
        )

        accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
        k_remaining = 2 * K - pid_k * SPLITK_BLOCK_SIZE
        shuffle_remaining = K * 16 - pid_k * (SPLITK_BLOCK_SIZE // 2) * 16

        for _ in range(num_k_iter):
            b_scales = (
                tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
                .reshape(
                    BLOCK_SIZE_N // 32,
                    BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8,
                    4,
                    16,
                    2,
                    2,
                    1,
                )
                .permute(0, 5, 3, 1, 4, 2, 6)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
            )

            if EVEN_K:
                a_bf16 = tl.load(a_ptrs)
                b = tl.load(b_ptrs, cache_modifier=cache_modifier)
            else:
                a_bf16 = tl.load(a_ptrs, mask=offs_k_bf16[None, :] < k_remaining, other=0)
                b = tl.load(
                    b_ptrs,
                    mask=offs_k_shuffle_arr[None, :] < shuffle_remaining,
                    other=0,
                    cache_modifier=cache_modifier,
                )

            b = (
                b.reshape(
                    1,
                    BLOCK_SIZE_N // 16,
                    BLOCK_SIZE_K // 64,
                    2,
                    16,
                    16,
                )
                .permute(0, 1, 4, 2, 3, 5)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
                .trans(1, 0)
            )

            if PREQUANT:
                a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)

            accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")

            a_ptrs += BLOCK_SIZE_K * stride_ak
            b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
            b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
            k_remaining -= BLOCK_SIZE_K
            shuffle_remaining -= (BLOCK_SIZE_K // 2) * 16

        c = accumulator.to(c_ptr.type.element_ty)
        offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
        offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
        c_ptrs = (
            c_ptr
            + stride_cm * offs_cm[:, None]
            + stride_cn * offs_cn[None, :]
            + pid_k * stride_ck
        )
        c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
        tl.store(c_ptrs, c, mask=c_mask)


def _launch_preshuffle_target(
    a: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> torch.Tensor:
    m, _k_bf16 = a.shape
    n, k_packed = b_shuffle.shape[0] * 16, b_shuffle.shape[1] // 16
    k_orig = 2 * k_packed
    shape = (m, n, k_orig)

    config = _SHAPE_CONFIGS.get(shape)
    if config is None:
        config, _ = _get_config(m, n, k_packed, True)
    else:
        config = dict(config)

    if config["NUM_KSPLIT"] > 1:
        splitk_block_size, block_size_k, num_ksplit = get_splitk(
            k_packed, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]
        )
        config["SPLITK_BLOCK_SIZE"] = splitk_block_size
        config["BLOCK_SIZE_K"] = block_size_k
        config["NUM_KSPLIT"] = num_ksplit

    if config["BLOCK_SIZE_K"] >= 2 * k_packed:
        config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * k_packed)
        config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
        config["NUM_KSPLIT"] = 1
    config["BLOCK_SIZE_N"] = max(config["BLOCK_SIZE_N"], 32)

    if shape in _PRESHUFFLE_TARGET_SHAPES:
        assert (
            config["NUM_KSPLIT"] == _TARGET_NUM_KSPLIT
        ), f"expected shape-2 producer to emit {_TARGET_NUM_KSPLIT} partials, got {config['NUM_KSPLIT']}"

    if config["NUM_KSPLIT"] > 1:
        partials = torch.empty(
            (config["NUM_KSPLIT"], m, n), dtype=torch.float32, device=a.device
        )
    else:
        partials = None
        config["SPLITK_BLOCK_SIZE"] = 2 * k_packed

    out = torch.empty((m, n), dtype=torch.bfloat16, device=a.device).contiguous()
    grid = (
        config["NUM_KSPLIT"]
        * triton.cdiv(m, config["BLOCK_SIZE_M"])
        * triton.cdiv(n, config["BLOCK_SIZE_N"]),
    )
    _gemm_a16wfp4_preshuffle_kernel_fixed[grid](
        a,
        b_shuffle,
        out if partials is None else partials,
        b_scale_sh,
        m,
        n,
        k_packed,
        a.stride(0),
        a.stride(1),
        b_shuffle.stride(0),
        b_shuffle.stride(1),
        0 if partials is None else partials.stride(0),
        out.stride(0) if partials is None else partials.stride(1),
        out.stride(1) if partials is None else partials.stride(2),
        b_scale_sh.stride(0),
        b_scale_sh.stride(1),
        PREQUANT=True,
        **config,
    )

    if partials is None:
        return out

    actual_ksplit = triton.cdiv(k_packed, (config["SPLITK_BLOCK_SIZE"] // 2))
    if shape in _PRESHUFFLE_TARGET_SHAPES and actual_ksplit == _TARGET_NUM_KSPLIT:
        return _run_reduce_s14(partials, out)

    reduce_block_m = 16
    reduce_block_n = 64
    grid_reduce = (
        triton.cdiv(m, reduce_block_m),
        triton.cdiv(n, reduce_block_n),
    )
    _gemm_afp4wfp4_reduce_kernel[grid_reduce](
        partials,
        out,
        m,
        n,
        partials.stride(0),
        partials.stride(1),
        partials.stride(2),
        out.stride(0),
        out.stride(1),
        reduce_block_m,
        reduce_block_n,
        actual_ksplit,
        triton.next_power_of_2(config["NUM_KSPLIT"]),
    )
    return out


def _launch_hybrid(
    a: torch.Tensor,
    b_q: torch.Tensor,
    b_scale: torch.Tensor,
) -> torch.Tensor:
    m, _k_bf16 = a.shape
    n, k = b_q.shape
    k_orig = 2 * k
    b_q_t = b_q.t()

    config = _SHAPE_CONFIGS.get((m, n, k_orig))
    if config is None:
        config, _ = _get_config(m, n, k)
    else:
        config = dict(config)

    if config["NUM_KSPLIT"] > 1:
        splitk_block_size, block_size_k, num_ksplit = get_splitk(
            k, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]
        )
        config["SPLITK_BLOCK_SIZE"] = splitk_block_size
        config["BLOCK_SIZE_K"] = block_size_k
        config["NUM_KSPLIT"] = num_ksplit

    if config["BLOCK_SIZE_K"] >= 2 * k:
        config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * k)
        config["SPLITK_BLOCK_SIZE"] = 2 * k
        config["NUM_KSPLIT"] = 1
    config["BLOCK_SIZE_K"] = max(config["BLOCK_SIZE_K"], 64)

    if config["NUM_KSPLIT"] > 1:
        partials = torch.empty(
            (config["NUM_KSPLIT"], m, n), dtype=torch.float32, device=a.device
        )
    else:
        partials = None
        config["SPLITK_BLOCK_SIZE"] = 2 * k

    out = torch.empty((m, n), dtype=torch.bfloat16, device=a.device)
    grid = (
        config["NUM_KSPLIT"]
        * triton.cdiv(m, config["BLOCK_SIZE_M"])
        * triton.cdiv(n, config["BLOCK_SIZE_N"]),
    )
    _gemm_a16wfp4_kernel[grid](
        a,
        b_q_t,
        out if partials is None else partials,
        b_scale,
        m,
        n,
        k,
        a.stride(0),
        a.stride(1),
        b_q_t.stride(0),
        b_q_t.stride(1),
        0 if partials is None else partials.stride(0),
        out.stride(0) if partials is None else partials.stride(1),
        out.stride(1) if partials is None else partials.stride(2),
        b_scale.stride(0),
        b_scale.stride(1),
        ATOMIC_ADD=False,
        **config,
    )

    if partials is None:
        return out

    reduce_block_m = 16
    reduce_block_n = 64
    actual_ksplit = triton.cdiv(k, (config["SPLITK_BLOCK_SIZE"] // 2))
    grid_reduce = (
        triton.cdiv(m, reduce_block_m),
        triton.cdiv(n, reduce_block_n),
    )
    _gemm_afp4wfp4_reduce_kernel[grid_reduce](
        partials,
        out,
        m,
        n,
        partials.stride(0),
        partials.stride(1),
        partials.stride(2),
        out.stride(0),
        out.stride(1),
        reduce_block_m,
        reduce_block_n,
        actual_ksplit,
        triton.next_power_of_2(config["NUM_KSPLIT"]),
    )
    return out


def custom_kernel(data: input_t) -> output_t:
    a, _b, b_q, b_shuffle, b_scale_sh = data
    a = a if a.is_contiguous() else a.contiguous()
    m, k = a.shape
    n = b_q.shape[0]
    shape = (m, n, k)

    if shape in _TRITON_FUSED_SHAPES:
        k_blocks = k // 32
        b_scale_ps = _task_scale_to_preshuffle(b_scale_sh, n, k_blocks)
        b_shuffle_u8 = b_shuffle.view(torch.uint8).contiguous()
        b_shuffle_tri = b_shuffle_u8.view(n // 16, b_shuffle_u8.shape[1] * 16).contiguous()
        return _launch_preshuffle_target(a, b_shuffle_tri, b_scale_ps)

    return _launch_zero_overhead_s56(a, b_shuffle, b_scale_sh)
scrolls · 1368 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