Skip to content
KernelIndex
Search⌘K

claude-opus-4-1-20250805 / cuda6c53f4

claude-opus-4-1-20250805_cuda_6c53f4 · claude-opus-4-1-20250805 · cuda · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 183 lines, Apache-2.0, pinned at da91508.

main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-20250805-cuda-6c53f4?include=source"
interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, fp8_e4m3, int32

Benchmark evidence

No published measurement for this revision.

No evidence · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:c2097132be03260122cb02b5dbd84a3286715b94f92e6433c823a3c101ce0213
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20

Kernel source

main.cpp183 lines
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <vector>
#include "kernel.h"

torch::Tensor run(
    torch::Tensor routing_logits,
    torch::Tensor routing_bias,
    torch::Tensor hidden_states,
    torch::Tensor hidden_states_scale,
    torch::Tensor gemm1_weights,
    torch::Tensor gemm1_weights_scale,
    torch::Tensor gemm2_weights,
    torch::Tensor gemm2_weights_scale,
    int64_t local_expert_offset,
    double routed_scaling_factor) {
    
    // Input validation
    TORCH_CHECK(routing_logits.is_cuda(), "routing_logits must be on CUDA");
    TORCH_CHECK(routing_bias.is_cuda(), "routing_bias must be on CUDA");
    TORCH_CHECK(hidden_states.is_cuda(), "hidden_states must be on CUDA");
    TORCH_CHECK(hidden_states_scale.is_cuda(), "hidden_states_scale must be on CUDA");
    TORCH_CHECK(gemm1_weights.is_cuda(), "gemm1_weights must be on CUDA");
    TORCH_CHECK(gemm1_weights_scale.is_cuda(), "gemm1_weights_scale must be on CUDA");
    TORCH_CHECK(gemm2_weights.is_cuda(), "gemm2_weights must be on CUDA");
    TORCH_CHECK(gemm2_weights_scale.is_cuda(), "gemm2_weights_scale must be on CUDA");
    
    // Get dimensions
    const int seq_len = routing_logits.size(0);
    const int num_experts = routing_logits.size(1);
    const int hidden_size = hidden_states.size(1);
    const int num_local_experts = gemm1_weights.size(0);
    
    // Validate dimensions
    TORCH_CHECK(num_experts == 256, "num_experts must be 256, got ", num_experts);
    TORCH_CHECK(num_local_experts == 32, "num_local_experts must be 32, got ", num_local_experts);
    TORCH_CHECK(hidden_size == 7168, "hidden_size must be 7168, got ", hidden_size);
    
    // Validate shapes
    TORCH_CHECK(routing_bias.numel() >= num_experts, "routing_bias shape mismatch");
    TORCH_CHECK(hidden_states_scale.size(0) == 56 && hidden_states_scale.size(1) == seq_len,
                "hidden_states_scale shape mismatch");
    TORCH_CHECK(gemm1_weights.size(1) == 4096 && gemm1_weights.size(2) == 7168,
                "gemm1_weights shape mismatch");
    TORCH_CHECK(gemm2_weights.size(1) == 7168 && gemm2_weights.size(2) == 2048,
                "gemm2_weights shape mismatch");
    
    // Ensure contiguous tensors
    routing_logits = routing_logits.contiguous();
    routing_bias = routing_bias.contiguous().view({num_experts});
    hidden_states = hidden_states.contiguous();
    hidden_states_scale = hidden_states_scale.contiguous();
    gemm1_weights = gemm1_weights.contiguous();
    gemm1_weights_scale = gemm1_weights_scale.contiguous();
    gemm2_weights = gemm2_weights.contiguous();
    gemm2_weights_scale = gemm2_weights_scale.contiguous();
    
    // Convert tensors to appropriate types
    if (routing_logits.scalar_type() != torch::kFloat32) {
        routing_logits = routing_logits.to(torch::kFloat32);
    }
    if (routing_bias.scalar_type() != torch::kBFloat16) {
        routing_bias = routing_bias.to(torch::kBFloat16);
    }
    if (hidden_states_scale.scalar_type() != torch::kFloat32) {
        hidden_states_scale = hidden_states_scale.to(torch::kFloat32);
    }
    if (gemm1_weights_scale.scalar_type() != torch::kFloat32) {
        gemm1_weights_scale = gemm1_weights_scale.to(torch::kFloat32);
    }
    if (gemm2_weights_scale.scalar_type() != torch::kFloat32) {
        gemm2_weights_scale = gemm2_weights_scale.to(torch::kFloat32);
    }
    
    // Handle FP8 tensors - convert to uint8 view
    torch::Tensor hidden_states_uint8, gemm1_weights_uint8, gemm2_weights_uint8;
    
    auto convert_fp8_to_uint8 = [](torch::Tensor tensor) -> torch::Tensor {
        // Check if tensor is already uint8
        if (tensor.scalar_type() == torch::kUInt8) {
            return tensor;
        }
        
        // Handle FP8 E4M3FN type
        if (tensor.scalar_type() == torch::kFloat8_e4m3fn) {
            // Reinterpret FP8 data as uint8
            return tensor.view(torch::kUInt8);
        }
        
        // Handle int8/char types
        if (tensor.scalar_type() == torch::kInt8 || tensor.scalar_type() == torch::kChar) {
            // Reinterpret as uint8
            return tensor.view(torch::kUInt8);
        }
        
        // For other types (e.g., float for testing), quantize to uint8
        if (tensor.scalar_type() == torch::kFloat32 || 
            tensor.scalar_type() == torch::kFloat16 ||
            tensor.scalar_type() == torch::kBFloat16) {
            // Simple quantization for testing
            auto float_tensor = tensor.to(torch::kFloat32);
            auto abs_max = float_tensor.abs().max();
            if (abs_max.item<float>() == 0.0f) {
                return torch::zeros_like(tensor, torch::TensorOptions().dtype(torch::kUInt8));
            }
            float scale = 127.0f / abs_max.item<float>();
            auto quantized = (float_tensor * scale + 128.0f).round().clamp(0, 255);
            return quantized.to(torch::kUInt8);
        }
        
        // Default: try to convert directly
        return tensor.to(torch::kUInt8);
    };
    
    hidden_states_uint8 = convert_fp8_to_uint8(hidden_states);
    gemm1_weights_uint8 = convert_fp8_to_uint8(gemm1_weights);
    gemm2_weights_uint8 = convert_fp8_to_uint8(gemm2_weights);
    
    // Create output tensor
    auto options = torch::TensorOptions()
        .dtype(torch::kBFloat16)
        .device(hidden_states.device());
    torch::Tensor output = torch::zeros({seq_len, hidden_size}, options);
    
    // Get CUDA stream
    cudaStream_t stream = at::cuda::getCurrentCUDAStream();
    
    // Get data pointers
    const float* routing_logits_ptr = routing_logits.data_ptr<float>();
    const __nv_bfloat16* routing_bias_ptr = reinterpret_cast<const __nv_bfloat16*>(
        routing_bias.data_ptr<at::BFloat16>());
    const uint8_t* hidden_states_ptr = hidden_states_uint8.data_ptr<uint8_t>();
    const float* hidden_states_scale_ptr = hidden_states_scale.data_ptr<float>();
    const uint8_t* gemm1_weights_ptr = gemm1_weights_uint8.data_ptr<uint8_t>();
    const float* gemm1_weights_scale_ptr = gemm1_weights_scale.data_ptr<float>();
    const uint8_t* gemm2_weights_ptr = gemm2_weights_uint8.data_ptr<uint8_t>();
    const float* gemm2_weights_scale_ptr = gemm2_weights_scale.data_ptr<float>();
    __nv_bfloat16* output_ptr = reinterpret_cast<__nv_bfloat16*>(
        output.data_ptr<at::BFloat16>());
    
    // Launch kernels
    launch_moe_kernels(
        routing_logits_ptr,
        routing_bias_ptr,
        hidden_states_ptr,
        hidden_states_scale_ptr,
        gemm1_weights_ptr,
        gemm1_weights_scale_ptr,
        gemm2_weights_ptr,
        gemm2_weights_scale_ptr,
        static_cast<int>(local_expert_offset),
        static_cast<float>(routed_scaling_factor),
        output_ptr,
        seq_len,
        stream
    );
    
    // Synchronize to ensure completion
    cudaStreamSynchronize(stream);
    
    // Check for CUDA errors
    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess) {
        TORCH_CHECK(false, "CUDA kernel error: ", cudaGetErrorString(err));
    }
    
    return output;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run", &run, "MoE FP8 block scale DS routing kernel (B200 optimized)",
          py::arg("routing_logits"),
          py::arg("routing_bias"),
          py::arg("hidden_states"),
          py::arg("hidden_states_scale"),
          py::arg("gemm1_weights"),
          py::arg("gemm1_weights_scale"),
          py::arg("gemm2_weights"),
          py::arg("gemm2_weights_scale"),
          py::arg("local_expert_offset"),
          py::arg("routed_scaling_factor"));
}
scrolls · 183 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON