gemini-2.5-pro_cuda_c02672
gemini-2.5-pro · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 104 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-cuda-c02672?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:c00c440257fc6d45ebb960cb4f54a07cf867599feb98057a3b525c4f59e3fc65
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
params.hidden_states = reinterpret_cast<const __nv_fp8_e4m3*>(hidden_states.data_ptr());Kernel source
main.cpp104 lines
#include <torch/extension.h>
#include "kernel.h"
#include <vector>
#include <string>
// Helper to check tensor properties
void check_tensor(const torch::Tensor& tensor, const std::string& name, torch::ScalarType dtype, const std::vector<int64_t>& shape) {
TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(tensor.scalar_type() == dtype, name, " has incorrect dtype, expected ", dtype, " but got ", tensor.scalar_type());
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(tensor.dim() == shape.size(), name, " has incorrect number of dimensions");
for (size_t i = 0; i < shape.size(); ++i) {
if (shape[i] != -1) {
TORCH_CHECK(tensor.size(i) == shape[i], name, " has incorrect shape at dim ", i, ", expected ", shape[i], " but got ", tensor.size(i));
}
}
}
torch::Tensor moe_fp8_block_scale_ds_routing(
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,
int local_expert_offset,
float routed_scaling_factor
) {
// --- Input Validation ---
const int seq_len = routing_logits.size(0);
check_tensor(routing_logits, "routing_logits", torch::kFloat32, {seq_len, NUM_EXPERTS});
check_tensor(routing_bias, "routing_bias", torch::kBFloat16, {NUM_EXPERTS});
check_tensor(hidden_states, "hidden_states", torch::kFloat8_e4m3fn, {seq_len, HIDDEN_SIZE});
check_tensor(hidden_states_scale, "hidden_states_scale", torch::kFloat32, {NUM_HIDDEN_BLOCKS, seq_len});
check_tensor(gemm1_weights, "gemm1_weights", torch::kFloat8_e4m3fn, {NUM_LOCAL_EXPERTS, GEMM1_OUT_SIZE, HIDDEN_SIZE});
check_tensor(gemm1_weights_scale, "gemm1_weights_scale", torch::kFloat32, {NUM_LOCAL_EXPERTS, NUM_GEMM1_OUT_BLOCKS, NUM_HIDDEN_BLOCKS});
check_tensor(gemm2_weights, "gemm2_weights", torch::kFloat8_e4m3fn, {NUM_LOCAL_EXPERTS, HIDDEN_SIZE, INTERMEDIATE_SIZE});
check_tensor(gemm2_weights_scale, "gemm2_weights_scale", torch::kFloat32, {NUM_LOCAL_EXPERTS, NUM_HIDDEN_BLOCKS, NUM_INTERMEDIATE_BLOCKS});
// --- Output and Workspace Allocation ---
auto options_bf16 = torch::TensorOptions().device(torch::kCUDA).dtype(torch::kBFloat16);
auto options_f32 = torch::TensorOptions().device(torch::kCUDA).dtype(torch::kFloat32);
auto options_i32 = torch::TensorOptions().device(torch::kCUDA).dtype(torch::kInt32);
auto output = torch::zeros({seq_len, HIDDEN_SIZE}, options_bf16);
auto topk_indices = torch::empty({seq_len, TOP_K}, options_i32);
auto topk_weights = torch::empty({seq_len, TOP_K}, options_f32);
auto expert_token_counts = torch::empty({NUM_LOCAL_EXPERTS}, options_i32);
auto expert_token_offsets = torch::empty({NUM_LOCAL_EXPERTS}, options_i32);
int max_dispatched_tokens = seq_len * TOP_K;
auto sorted_token_indices = torch::full({max_dispatched_tokens}, -1, options_i32);
auto token_expert_mapping = torch::full({max_dispatched_tokens}, -1, options_i32);
// FIX: Allocate as 1D, since it's accessed as 1D in the kernel.
auto token_storage_map = torch::full({(long)seq_len * TOP_K}, -1, options_i32);
auto temp_expert_output = torch::empty({(long)max_dispatched_tokens, HIDDEN_SIZE}, options_f32);
// --- Prepare Kernel Parameters ---
MoeKernelParams params;
params.seq_len = seq_len;
params.max_dispatched_tokens = max_dispatched_tokens;
// Inputs
params.routing_logits = routing_logits.data_ptr<float>();
params.routing_bias = reinterpret_cast<const __nv_bfloat16*>(routing_bias.data_ptr());
params.hidden_states = reinterpret_cast<const __nv_fp8_e4m3*>(hidden_states.data_ptr());
params.hidden_states_scale = hidden_states_scale.data_ptr<float>();
params.gemm1_weights = reinterpret_cast<const __nv_fp8_e4m3*>(gemm1_weights.data_ptr());
params.gemm1_weights_scale = gemm1_weights_scale.data_ptr<float>();
params.gemm2_weights = reinterpret_cast<const __nv_fp8_e4m3*>(gemm2_weights.data_ptr());
params.gemm2_weights_scale = gemm2_weights_scale.data_ptr<float>();
params.local_expert_offset = local_expert_offset;
params.routed_scaling_factor = routed_scaling_factor;
// Output
params.output = reinterpret_cast<__nv_bfloat16*>(output.data_ptr());
// Workspace
params.topk_indices = topk_indices.data_ptr<int>();
params.topk_weights = topk_weights.data_ptr<float>();
params.expert_token_counts = expert_token_counts.data_ptr<int>();
params.expert_token_offsets = expert_token_offsets.data_ptr<int>();
params.sorted_token_indices = sorted_token_indices.data_ptr<int>();
params.token_expert_mapping = token_expert_mapping.data_ptr<int>();
params.token_storage_map = token_storage_map.data_ptr<int>();
params.temp_expert_output = temp_expert_output.data_ptr<float>();
// --- Launch Kernels ---
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
run_moe_kernels(params, stream);
return output;
}
// --- Pybind11 Module Definition ---
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("run", &moe_fp8_block_scale_ds_routing, "MoE FP8 block-scale with DeepSeek Routing (CUDA)");
}scrolls · 104 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON