Skip to content
KernelIndex
Search⌘K

submission 340626

vesper15 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 dual GEMMsuite of 4 cases
NVIDIA B200
40.8µs
#284 of 420
2026-01-13

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5750db2af15ac4ea29c705c337e7979735be495e1913734e7dee82642837f321
license declaredunknown
license concludedunknown
authorsvesper15
imported2026-08-26

Techniques

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

fp4Optimized NVFP4 Dual GEMM with SiLU using custom CUDA kernels.
fused-epilogue__global__ void fused_silu_mul_kernel(
vector-width = float4const float4* __restrict__ x,

Kernel source

submission.py698 lines
import torch
from torch.utils.cpp_extension import load_inline

cuda_source = """
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>

// Optimized scale factor conversion kernel
// Maps from [M, K//16] to blocked format for torch._scaled_mm
__global__ void to_blocked_kernel(
    const uint8_t* __restrict__ input,
    uint8_t* __restrict__ output,
    const int M,
    const int K16,  // K // 16
    const int n_row_blocks,
    const int n_col_blocks
) {
    // Each thread handles one input element
    const int idx = blockIdx.x * blockDim.x + threadIdx.x;
    const int total_input = M * K16;

    if (idx < total_input) {
        const int m = idx / K16;
        const int k = idx % K16;

        // Compute output position based on blocked layout transformation
        const int row_block = m / 128;
        const int col_block = k / 4;
        const int in_row = m % 128;
        const int in_col = k % 4;
        const int i32 = in_row % 32;
        const int i4 = in_row / 32;

        const int block_idx = row_block * n_col_blocks + col_block;
        const int out_idx = block_idx * 512 + i32 * 16 + i4 * 4 + in_col;

        output[out_idx] = input[idx];
    }
}

torch::Tensor to_blocked_cuda(torch::Tensor input) {
    const int M = input.size(0);
    const int K16 = input.size(1);
    const int n_row_blocks = (M + 127) / 128;
    const int n_col_blocks = (K16 + 3) / 4;

    const int output_size = n_row_blocks * n_col_blocks * 512;
    auto output = torch::empty({output_size}, input.options());

    const int total = M * K16;
    const int threads = 512;
    const int blocks = (total + threads - 1) / threads;

    to_blocked_kernel<<<blocks, threads>>>(
        input.data_ptr<uint8_t>(),
        output.data_ptr<uint8_t>(),
        M, K16, n_row_blocks, n_col_blocks
    );

    return output;
}

// Fused epilogue: silu(x) * y -> fp16
// Vectorized for maximum throughput
__global__ void fused_silu_mul_kernel(
    const float* __restrict__ x,
    const float* __restrict__ y,
    half* __restrict__ out,
    const int size
) {
    const int idx = blockIdx.x * blockDim.x + threadIdx.x;
    const int stride = blockDim.x * gridDim.x;

    for (int i = idx; i < size; i += stride) {
        float x_val = x[i];
        float silu = x_val / (1.0f + __expf(-x_val));
        out[i] = __float2half(silu * y[i]);
    }
}

// Vectorized version - 4 elements per thread
__global__ void fused_silu_mul_vec4_kernel(
    const float4* __restrict__ x,
    const float4* __restrict__ y,
    half2* __restrict__ out,
    const int size4
) {
    const int lane = blockIdx.x * blockDim.x + threadIdx.x;
    const int stride = blockDim.x * gridDim.x;

    for (int i = lane; i < size4; i += stride * 2) {
        int idx = i;
        if (idx < size4) {
            float4 xv = x[idx];
            float4 yv = y[idx];

            float s0 = xv.x / (1.0f + __expf(-xv.x)) * yv.x;
            float s1 = xv.y / (1.0f + __expf(-xv.y)) * yv.y;
            float s2 = xv.z / (1.0f + __expf(-xv.z)) * yv.z;
            float s3 = xv.w / (1.0f + __expf(-xv.w)) * yv.w;

            out[idx * 2] = __floats2half2_rn(s0, s1);
            out[idx * 2 + 1] = __floats2half2_rn(s2, s3);
        }

        idx += stride;
        if (idx < size4) {
            float4 xv = x[idx];
            float4 yv = y[idx];

            float s0 = xv.x / (1.0f + __expf(-xv.x)) * yv.x;
            float s1 = xv.y / (1.0f + __expf(-xv.y)) * yv.y;
            float s2 = xv.z / (1.0f + __expf(-xv.z)) * yv.z;
            float s3 = xv.w / (1.0f + __expf(-xv.w)) * yv.w;

            out[idx * 2] = __floats2half2_rn(s0, s1);
            out[idx * 2 + 1] = __floats2half2_rn(s2, s3);
        }
    }
}

__global__ void fused_silu_mul_strided_kernel(
    const float* __restrict__ x,
    const float* __restrict__ y,
    half* __restrict__ out,
    const int rows,
    const int cols,
    const int ld_rows,
    const int ld_cols
) {
    const int idx = blockIdx.x * blockDim.x + threadIdx.x;
    const int stride = blockDim.x * gridDim.x;
    const int total = rows * cols;

    for (int i = idx; i < total; i += stride) {
        const int r = i / cols;
        const int c = i % cols;

        float x_val = x[i];
        float silu = x_val / (1.0f + __expf(-x_val));
        float y_val = y[i];
        out[r * ld_rows + c * ld_cols] = __float2half(silu * y_val);
    }
}

torch::Tensor fused_silu_mul(torch::Tensor x, torch::Tensor y) {
    auto out = torch::empty({x.size(0), x.size(1)},
        torch::dtype(torch::kFloat16).device(x.device()));

    const int size = x.numel();

    if (size % 4 == 0) {
        const int size4 = size / 4;
        const int threads = 256;
        const int blocks = (size4 + threads - 1) / threads;

        fused_silu_mul_vec4_kernel<<<blocks, threads>>>(
            reinterpret_cast<const float4*>(x.data_ptr<float>()),
            reinterpret_cast<const float4*>(y.data_ptr<float>()),
            reinterpret_cast<half2*>(out.data_ptr<at::Half>()),
            size4
        );
    } else {
        const int threads = 256;
        const int blocks = (size + threads - 1) / threads;

        fused_silu_mul_kernel<<<blocks, threads>>>(
            x.data_ptr<float>(),
            y.data_ptr<float>(),
            reinterpret_cast<half*>(out.data_ptr<at::Half>()),
            size
        );
    }

    return out;
}

void fused_silu_mul_inplace(torch::Tensor x, torch::Tensor y, torch::Tensor out) {
    TORCH_CHECK(x.is_cuda() && y.is_cuda() && out.is_cuda(), "All tensors must be CUDA");
    TORCH_CHECK(x.dtype() == torch::kFloat32 && y.dtype() == torch::kFloat32,
        "Inputs must be float32");
    TORCH_CHECK(out.dtype() == torch::kFloat16, "Output must be float16");
    TORCH_CHECK(x.sizes() == y.sizes(), "Input sizes must match");
    TORCH_CHECK(x.numel() == out.numel(), "Output must match input size");

    const int size = x.numel();

    const bool contiguous =
        out.stride(1) == 1 && out.stride(0) == out.size(1);

    if (contiguous && size % 4 == 0) {
        const int size4 = size / 4;
        const int threads = 256;
        const int blocks = (size4 + threads - 1) / threads;

        fused_silu_mul_vec4_kernel<<<blocks, threads>>>(
            reinterpret_cast<const float4*>(x.data_ptr<float>()),
            reinterpret_cast<const float4*>(y.data_ptr<float>()),
            reinterpret_cast<half2*>(out.data_ptr<at::Half>()),
            size4
        );
        return;
    }

    const int threads = 512;
    const int blocks = (size + threads - 1) / threads;

    if (contiguous) {
        fused_silu_mul_kernel<<<blocks, threads>>>(
            x.data_ptr<float>(),
            y.data_ptr<float>(),
            reinterpret_cast<half*>(out.data_ptr<at::Half>()),
            size
        );
    } else {
        fused_silu_mul_strided_kernel<<<blocks, threads>>>(
            x.data_ptr<float>(),
            y.data_ptr<float>(),
            reinterpret_cast<half*>(out.data_ptr<at::Half>()),
            out.size(0),
            out.size(1),
            static_cast<int>(out.stride(0)),
            static_cast<int>(out.stride(1))
        );
    }
}

__global__ void pack_fp8_permute_kernel(
    const uint8_t* __restrict__ input,
    uint8_t* __restrict__ output,
    const int dim0, const int dim1, const int dim2,
    const int order0, const int order1, const int order2
) {
    const int out_dim0 = (order0 == 0 ? dim0 : (order0 == 1 ? dim1 : dim2));
    const int out_dim1 = (order1 == 0 ? dim0 : (order1 == 1 ? dim1 : dim2));
    const int out_dim2 = (order2 == 0 ? dim0 : (order2 == 1 ? dim1 : dim2));

    const int total = out_dim0 * out_dim1 * out_dim2;
    const int tid = blockIdx.x * blockDim.x + threadIdx.x;
    const int stride = blockDim.x * gridDim.x;

    for (int idx = tid; idx < total; idx += stride) {
        int tmp = idx;
        const int o2 = tmp % out_dim2;
        tmp /= out_dim2;
        const int o1 = tmp % out_dim1;
        const int o0 = tmp / out_dim1;

        int coords[3];
        coords[order0] = o0;
        coords[order1] = o1;
        coords[order2] = o2;

        const int in_idx = coords[0] * dim1 * dim2 + coords[1] * dim2 + coords[2];
        output[idx] = input[in_idx];
    }
}

// Combined kernel that does everything in one launch
// Processes scale conversion + stores result
__global__ void process_scales_kernel(
    const uint8_t* __restrict__ sfa,
    const uint8_t* __restrict__ sfb1,
    const uint8_t* __restrict__ sfb2,
    uint8_t* __restrict__ scale_a,
    uint8_t* __restrict__ scale_b1,
    uint8_t* __restrict__ scale_b2,
    const int M, const int N,
    const int K16,
    const int n_row_blocks_a, const int n_col_blocks,
    const int n_row_blocks_b
) {
    const int tid = blockIdx.x * blockDim.x + threadIdx.x;
    const int stride = blockDim.x * gridDim.x;

    const int K16_vec = K16 >> 2;
    const int K16_rem = K16 & 3;
    const int base_k = K16_vec << 2;

    if (K16_vec > 0) {
        const int total_vec_a = M * K16_vec;
        for (int vec = tid; vec < total_vec_a; vec += stride) {
            const int m = vec / K16_vec;
            const int k4 = vec % K16_vec;

            const int row_block = m >> 7;
            const int in_row = m & 127;
            const int i32 = in_row & 31;
            const int i4 = in_row >> 5;
            const int block_idx = row_block * n_col_blocks + k4;
            const int out_idx = (block_idx << 9) + (i32 << 4) + (i4 << 2);

            const int base_in = m * K16 + (k4 << 2);
            const uint32_t val = *reinterpret_cast<const uint32_t*>(sfa + base_in);
            *reinterpret_cast<uint32_t*>(scale_a + out_idx) = val;
        }

        const int total_vec_b = N * K16_vec;
        for (int vec = tid; vec < total_vec_b; vec += stride) {
            const int n = vec / K16_vec;
            const int k4 = vec % K16_vec;

            const int row_block = n >> 7;
            const int in_row = n & 127;
            const int i32 = in_row & 31;
            const int i4 = in_row >> 5;
            const int block_idx = row_block * n_col_blocks + k4;
            const int out_idx = (block_idx << 9) + (i32 << 4) + (i4 << 2);

            const int base_in = n * K16 + (k4 << 2);
            const uint32_t val1 = *reinterpret_cast<const uint32_t*>(sfb1 + base_in);
            const uint32_t val2 = *reinterpret_cast<const uint32_t*>(sfb2 + base_in);
            *reinterpret_cast<uint32_t*>(scale_b1 + out_idx) = val1;
            *reinterpret_cast<uint32_t*>(scale_b2 + out_idx) = val2;
        }
    }

    if (K16_rem > 0) {
        const int total_rem_a = M * K16_rem;
        for (int rem = tid; rem < total_rem_a; rem += stride) {
            const int m = rem / K16_rem;
            const int k_offset = rem % K16_rem;
            const int k = base_k + k_offset;
            if (k >= K16) {
                continue;
            }

            const int row_block = m >> 7;
            const int in_row = m & 127;
            const int i32 = in_row & 31;
            const int i4 = in_row >> 5;
            const int col_block = k >> 2;
            const int in_col = k & 3;
            const int block_idx = row_block * n_col_blocks + col_block;
            const int out_idx = (block_idx << 9) + (i32 << 4) + (i4 << 2) + in_col;

            scale_a[out_idx] = sfa[m * K16 + k];
        }

        const int total_rem_b = N * K16_rem;
        for (int rem = tid; rem < total_rem_b; rem += stride) {
            const int n = rem / K16_rem;
            const int k_offset = rem % K16_rem;
            const int k = base_k + k_offset;
            if (k >= K16) {
                continue;
            }

            const int row_block = n >> 7;
            const int in_row = n & 127;
            const int i32 = in_row & 31;
            const int i4 = in_row >> 5;
            const int col_block = k >> 2;
            const int in_col = k & 3;
            const int block_idx = row_block * n_col_blocks + col_block;
            const int out_idx = (block_idx << 9) + (i32 << 4) + (i4 << 2) + in_col;

            scale_b1[out_idx] = sfb1[n * K16 + k];
            scale_b2[out_idx] = sfb2[n * K16 + k];
        }
    }
}

std::vector<torch::Tensor> process_all_scales(
    torch::Tensor sfa, torch::Tensor sfb1, torch::Tensor sfb2
) {
    const int M = sfa.size(0);
    const int N = sfb1.size(0);
    const int K16 = sfa.size(1);

    const int n_row_blocks_a = (M + 127) / 128;
    const int n_row_blocks_b = (N + 127) / 128;
    const int n_col_blocks = (K16 + 3) / 4;

    const int out_size_a = n_row_blocks_a * n_col_blocks * 512;
    const int out_size_b = n_row_blocks_b * n_col_blocks * 512;

    auto scale_a = torch::empty({out_size_a}, sfa.options());
    auto scale_b1 = torch::empty({out_size_b}, sfb1.options());
    auto scale_b2 = torch::empty({out_size_b}, sfb2.options());

    const int total = std::max(M * K16, N * K16);
    const int threads = 256;
    const int blocks = (total + threads - 1) / threads;

    process_scales_kernel<<<blocks, threads>>>(
        sfa.data_ptr<uint8_t>(),
        sfb1.data_ptr<uint8_t>(),
        sfb2.data_ptr<uint8_t>(),
        scale_a.data_ptr<uint8_t>(),
        scale_b1.data_ptr<uint8_t>(),
        scale_b2.data_ptr<uint8_t>(),
        M, N, K16,
        n_row_blocks_a, n_col_blocks, n_row_blocks_b
    );

    return {scale_a, scale_b1, scale_b2};
}

void pack_fp8_permute(
    torch::Tensor input,
    torch::Tensor output,
    int order0,
    int order1,
    int order2
) {
    TORCH_CHECK(input.dtype() == torch::kUInt8, "input must be uint8");
    TORCH_CHECK(output.dtype() == torch::kUInt8, "output must be uint8");
    TORCH_CHECK(input.dim() == 3, "input must be 3D");
    TORCH_CHECK(output.dim() == 3, "output must be 3D");

    const int dim0 = input.size(0);
    const int dim1 = input.size(1);
    const int dim2 = input.size(2);

    const int threads = 256;
    const int total = output.numel();
    const int blocks = (total + threads - 1) / threads;

    pack_fp8_permute_kernel<<<blocks, threads>>>(
        input.data_ptr<uint8_t>(),
        output.data_ptr<uint8_t>(),
        dim0, dim1, dim2,
        order0, order1, order2
    );
}

void process_all_scales_into(
    torch::Tensor sfa,
    torch::Tensor sfb1,
    torch::Tensor sfb2,
    torch::Tensor scale_a,
    torch::Tensor scale_b1,
    torch::Tensor scale_b2
) {
    const int M = sfa.size(0);
    const int N = sfb1.size(0);
    const int K16 = sfa.size(1);

    TORCH_CHECK(scale_a.dtype() == torch::kUInt8, "scale_a must be uint8");
    TORCH_CHECK(scale_b1.dtype() == torch::kUInt8, "scale_b1 must be uint8");
    TORCH_CHECK(scale_b2.dtype() == torch::kUInt8, "scale_b2 must be uint8");

    const int n_row_blocks_a = (M + 127) / 128;
    const int n_row_blocks_b = (N + 127) / 128;
    const int n_col_blocks = (K16 + 3) / 4;

    const int out_size_a = n_row_blocks_a * n_col_blocks * 512;
    const int out_size_b = n_row_blocks_b * n_col_blocks * 512;

    TORCH_CHECK(scale_a.numel() == out_size_a, "scale_a has incorrect size");
    TORCH_CHECK(scale_b1.numel() == out_size_b, "scale_b1 has incorrect size");
    TORCH_CHECK(scale_b2.numel() == out_size_b, "scale_b2 has incorrect size");

    const int total = std::max(M * K16, N * K16);
    const int threads = 256;
    const int blocks = (total + threads - 1) / threads;

    process_scales_kernel<<<blocks, threads>>>(
        sfa.data_ptr<uint8_t>(),
        sfb1.data_ptr<uint8_t>(),
        sfb2.data_ptr<uint8_t>(),
        scale_a.data_ptr<uint8_t>(),
        scale_b1.data_ptr<uint8_t>(),
        scale_b2.data_ptr<uint8_t>(),
        M, N, K16,
        n_row_blocks_a, n_col_blocks, n_row_blocks_b
    );
}
"""

cpp_source = """
#include <torch/extension.h>
#include <vector>

torch::Tensor to_blocked_cuda(torch::Tensor input);
torch::Tensor fused_silu_mul(torch::Tensor x, torch::Tensor y);
std::vector<torch::Tensor> process_all_scales(
    torch::Tensor sfa, torch::Tensor sfb1, torch::Tensor sfb2);
void process_all_scales_into(
    torch::Tensor sfa,
    torch::Tensor sfb1,
    torch::Tensor sfb2,
    torch::Tensor scale_a,
    torch::Tensor scale_b1,
    torch::Tensor scale_b2);
void fused_silu_mul_inplace(torch::Tensor x, torch::Tensor y, torch::Tensor out);
void pack_fp8_permute(
    torch::Tensor input,
    torch::Tensor output,
    int order0,
    int order1,
    int order2);
"""

# Compile CUDA kernels
cuda_module = load_inline(
    name='nvfp4_optimized',
    cpp_sources=cpp_source,
    cuda_sources=cuda_source,
    functions=[
        'to_blocked_cuda',
        'fused_silu_mul',
        'process_all_scales',
        'process_all_scales_into',
        'fused_silu_mul_inplace',
        'pack_fp8_permute'
    ],
    verbose=False,
    extra_cuda_cflags=['-O3', '--use_fast_math', '-lineinfo']
)

_SCALE_BUFFER_CACHE = {}
_FP8_LAYER_CACHE = {}
_SCALE_LAYER_CACHE = {}
_FP8_SLICE_CACHE = {}


def _device_cache_key(device):
    index = device.index if device.index is not None else -1
    return (device.type, index)

def _copy_fp8_slice_to_device(tensor_slice, device):
    bytes_gpu = tensor_slice.contiguous().view(torch.uint8).to(device=device, non_blocking=True)
    return bytes_gpu.view(tensor_slice.dtype)


def _get_cached_fp8_slice(tensor_slice, device, transpose=True):
    key = (
        _device_cache_key(device),
        int(tensor_slice.data_ptr()),
        tuple(tensor_slice.shape),
        tensor_slice._version,
        transpose,
    )
    cached = _FP8_SLICE_CACHE.get(key)
    if cached is not None:
        return cached

    gpu_tensor = _copy_fp8_slice_to_device(tensor_slice, device)
    if transpose:
        gpu_tensor = gpu_tensor.transpose(0, 1)
    _FP8_SLICE_CACHE[key] = gpu_tensor
    return gpu_tensor


def _prepare_uint8_layers(tensor, device, order):
    tensor_bytes = tensor.view(torch.uint8).contiguous()
    tensor_bytes = tensor_bytes.to(device=device, non_blocking=True)
    return tensor_bytes.view(tensor.shape).permute(*order).contiguous()


def _prepare_fp8_layers(tensor, device, order):
    return _prepare_uint8_layers(tensor, device, order).view(tensor.dtype)


def _prepare_scale_layers(tensor, device, order):
    return _prepare_uint8_layers(tensor, device, order)


def _get_cached_layers(tensor, device, order, cache, prepare_fn):
    key = (
        _device_cache_key(device),
        int(tensor.data_ptr()),
        tuple(tensor.shape),
        tuple(order),
    )
    version = tensor._version
    cached = cache.get(key)
    if cached and cached["version"] == version:
        return cached["tensor"]

    prepared = prepare_fn(tensor, device, order)
    cache[key] = {"tensor": prepared, "version": version}
    return prepared


def _get_cached_fp8_layers(tensor, device, order):
    return _get_cached_layers(tensor, device, order, _FP8_LAYER_CACHE, _prepare_fp8_layers)


def _get_cached_scale_layers(tensor, device, order):
    return _get_cached_layers(
        tensor, device, order, _SCALE_LAYER_CACHE, _prepare_scale_layers
    )


def _get_scale_buffers(device, M, N, K16):
    n_row_blocks_a = (M + 127) // 128
    n_row_blocks_b = (N + 127) // 128
    n_col_blocks = (K16 + 3) // 4

    scale_a_size = n_row_blocks_a * n_col_blocks * 512
    scale_b_size = n_row_blocks_b * n_col_blocks * 512

    key = (_device_cache_key(device), M, N, K16)
    cached = _SCALE_BUFFER_CACHE.get(key)
    if cached and cached["sizes"] == (scale_a_size, scale_b_size):
        return cached["buffers"]

    scale_a_buf = torch.empty(scale_a_size, dtype=torch.uint8, device=device)
    scale_b1_buf = torch.empty(scale_b_size, dtype=torch.uint8, device=device)
    scale_b2_buf = torch.empty(scale_b_size, dtype=torch.uint8, device=device)

    _SCALE_BUFFER_CACHE[key] = {
        "sizes": (scale_a_size, scale_b_size),
        "buffers": (scale_a_buf, scale_b1_buf, scale_b2_buf),
    }

    return scale_a_buf, scale_b1_buf, scale_b2_buf


def _run_single_layer(a, b1, b2, sfa, sfb1, sfb2, c_slice, device):
    fp8_dtype = sfa.dtype

    a_slice = _copy_fp8_slice_to_device(a[:, :, 0], device)
    sfa_slice = sfa[:, :, 0].view(torch.uint8).contiguous().to(device=device, non_blocking=True)
    sfb1_slice = sfb1[:, :, 0].view(torch.uint8).contiguous().to(device=device, non_blocking=True)
    sfb2_slice = sfb2[:, :, 0].view(torch.uint8).contiguous().to(device=device, non_blocking=True)

    scales = cuda_module.process_all_scales(sfa_slice, sfb1_slice, sfb2_slice)
    scale_a = scales[0].view(fp8_dtype)
    scale_b1 = scales[1].view(fp8_dtype)
    scale_b2 = scales[2].view(fp8_dtype)

    b1_t = _get_cached_fp8_slice(b1[:, :, 0], device, transpose=True)
    b2_t = _get_cached_fp8_slice(b2[:, :, 0], device, transpose=True)

    res1 = torch._scaled_mm(a_slice, b1_t, scale_a, scale_b1, out_dtype=torch.float32)
    res2 = torch._scaled_mm(a_slice, b2_t, scale_a, scale_b2, out_dtype=torch.float32)

    cuda_module.fused_silu_mul_inplace(res1, res2, c_slice)


def custom_kernel(data):
    """
    Optimized NVFP4 Dual GEMM with SiLU using custom CUDA kernels.
    """
    # Unpack data
    a, b1, b2, sfa, sfb1, sfb2, _, _, _, c = data

    device = a.device
    l = c.size(2)

    if l == 1:
        _run_single_layer(a, b1, b2, sfa, sfb1, sfb2, c[:, :, 0], device)
        return c

    # Prepare FP8 tensors in (L, ...) layouts once
    a_layers = _get_cached_fp8_layers(a, device, (2, 0, 1))
    b1_t_layers = _get_cached_fp8_layers(b1, device, (2, 1, 0))
    b2_t_layers = _get_cached_fp8_layers(b2, device, (2, 1, 0))

    sfa_layers = _get_cached_scale_layers(sfa, device, (2, 0, 1))
    sfb1_layers = _get_cached_scale_layers(sfb1, device, (2, 0, 1))
    sfb2_layers = _get_cached_scale_layers(sfb2, device, (2, 0, 1))

    # FP8 dtype for viewing scale buffers
    fp8_dtype = sfa.dtype

    # Reuse cached scale buffers sized for current shape
    M = a.size(0)
    N = sfb1.size(0)
    K16 = sfa.size(1)

    scale_a_buf, scale_b1_buf, scale_b2_buf = _get_scale_buffers(device, M, N, K16)

    scale_a_fp8 = scale_a_buf.view(fp8_dtype)
    scale_b1_fp8 = scale_b1_buf.view(fp8_dtype)
    scale_b2_fp8 = scale_b2_buf.view(fp8_dtype)

    for l_idx in range(l):
        sfa_slice = sfa_layers[l_idx]
        sfb1_slice = sfb1_layers[l_idx]
        sfb2_slice = sfb2_layers[l_idx]

        cuda_module.process_all_scales_into(
            sfa_slice,
            sfb1_slice,
            sfb2_slice,
            scale_a_buf,
            scale_b1_buf,
            scale_b2_buf
        )

        a_slice = a_layers[l_idx]
        b1_t = b1_t_layers[l_idx]
        b2_t = b2_t_layers[l_idx]

        res1 = torch._scaled_mm(a_slice, b1_t, scale_a_fp8, scale_b1_fp8, out_dtype=torch.float32)
        res2 = torch._scaled_mm(a_slice, b2_t, scale_a_fp8, scale_b2_fp8, out_dtype=torch.float32)

        c_slice = c[:, :, l_idx]
        cuda_module.fused_silu_mul_inplace(res1, res2, c_slice)

    return c
scrolls · 698 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