gpt-5 / cudab963ec
gpt-5_cuda_b963ec · gpt-5-2025-08-07 · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 87 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-cuda-b963ec?include=source"interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16
Benchmark evidence
7 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reproduction-ready · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:b2f1aa7c94a09bb37500ca69c93601f1e976c98354d26d84f1209b6b745f8d91
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20
Kernel source
main.cpp87 lines
#include "kernel.h"
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <ATen/cuda/CUDAContext.h>
#include <vector>
#include <stdexcept>
namespace py = pybind11;
// Helper to ensure a tensor is contiguous, BF16, and on a target device
static inline torch::Tensor to_contig_bf16_on_device(const torch::Tensor& t, c10::Device device) {
auto tt = t;
if (tt.dtype() != torch::kBFloat16) {
tt = tt.to(torch::kBFloat16);
}
if (tt.device() != device) {
// Move to target device while preserving dtype
tt = tt.to(device, tt.scalar_type(), /*non_blocking=*/false, /*copy=*/true);
}
if (!tt.is_contiguous()) {
tt = tt.contiguous();
}
return tt;
}
torch::Tensor run(torch::Tensor hidden_states,
torch::Tensor residual,
torch::Tensor weight)
{
TORCH_CHECK(hidden_states.dim() == 2, "hidden_states must be 2D [batch_size, 2048]");
TORCH_CHECK(residual.dim() == 2, "residual must be 2D [batch_size, 2048]");
TORCH_CHECK(weight.dim() == 1, "weight must be 1D [2048]");
const int64_t batch_size = hidden_states.size(0);
const int64_t hidden_size = hidden_states.size(1);
TORCH_CHECK(hidden_size == HIDDEN_SIZE_2048, "hidden_size must be 2048");
TORCH_CHECK(residual.size(0) == batch_size && residual.size(1) == hidden_size,
"residual shape must match hidden_states");
TORCH_CHECK(weight.size(0) == hidden_size, "weight must have shape [2048]");
// Select target device: If any input is CUDA, use that device; otherwise default to current CUDA device
bool any_cuda = hidden_states.is_cuda() || residual.is_cuda() || weight.is_cuda();
c10::Device compute_device = any_cuda
? (hidden_states.is_cuda() ? hidden_states.device()
: (residual.is_cuda() ? residual.device() : weight.device()))
: c10::Device(c10::kCUDA, at::cuda::current_device());
// Prepare tensors on compute device in BF16 and contiguous layout
auto hs = to_contig_bf16_on_device(hidden_states, compute_device);
auto rs = to_contig_bf16_on_device(residual, compute_device);
auto wt = to_contig_bf16_on_device(weight, compute_device);
// Allocate output on compute device
auto out = torch::empty({batch_size, hidden_size}, hs.options().dtype(torch::kBFloat16).device(compute_device));
// Launch kernel on the current CUDA stream
cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
const __nv_bfloat16* hs_ptr = reinterpret_cast<const __nv_bfloat16*>(hs.data_ptr<at::BFloat16>());
const __nv_bfloat16* rs_ptr = reinterpret_cast<const __nv_bfloat16*>(rs.data_ptr<at::BFloat16>());
const __nv_bfloat16* wt_ptr = reinterpret_cast<const __nv_bfloat16*>(wt.data_ptr<at::BFloat16>());
__nv_bfloat16* out_ptr = reinterpret_cast<__nv_bfloat16*>(out.data_ptr<at::BFloat16>());
launch_fused_add_rmsnorm_h2048(
hs_ptr, rs_ptr, wt_ptr, out_ptr,
static_cast<int>(batch_size),
RMS_EPS,
stream
);
// If all original inputs were on CPU, bring result back to CPU
if (!hidden_states.is_cuda() && !residual.is_cuda() && !weight.is_cuda()) {
return out.to(hidden_states.options().device(torch::kCPU));
}
// Otherwise, return tensor on compute device
return out;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Fused Add + RMSNorm (hidden_size=2048, BF16) optimized for B200";
m.def("run", &run,
py::arg("hidden_states"),
py::arg("residual"),
py::arg("weight"),
"Fused add + RMSNorm with BF16 inputs and output, hidden_size fixed at 2048");
}scrolls · 87 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reproducible
JSON