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