gpt-o3 / cuda3c881e
gpt-o3_cuda_3c881e · gpt-o3 · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 95 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-cuda-3c881e?include=source"interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16
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:295b435669d95971a352c40aad3f0a37bce5e740dff0c7e44dc3eb124e43bedd
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20
Kernel source
main.cpp95 lines
/*
* PyTorch C++ / CUDA extension – entry point for rmsnorm_h2048
*/
#include "kernel.h"
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <stdexcept>
#include <string>
namespace py = pybind11;
/* -------------------------------------------------------------------- */
/* Light-weight CUDA error checker */
/* -------------------------------------------------------------------- */
inline void check_cuda(cudaError_t err, const char* where)
{
if (err != cudaSuccess)
throw std::runtime_error(std::string("CUDA error @ ") + where + ": "
+ cudaGetErrorString(err));
}
/* -------------------------------------------------------------------- */
/* Python-visible wrapper */
/* -------------------------------------------------------------------- */
torch::Tensor run(torch::Tensor hidden_states,
torch::Tensor weight,
double eps = 1e-6)
{
/* Sanity checks – stay strict, crash early */
TORCH_CHECK(hidden_states.is_cuda(), "hidden_states must be CUDA");
TORCH_CHECK(weight.is_cuda(), "weight must be CUDA");
TORCH_CHECK(hidden_states.scalar_type() == torch::kBFloat16,
"hidden_states must be bfloat16");
TORCH_CHECK(weight.scalar_type() == torch::kBFloat16,
"weight must be bfloat16");
TORCH_CHECK(hidden_states.size(-1) == HIDDEN_SIZE,
"last dimension must be 2048");
TORCH_CHECK(weight.numel() == HIDDEN_SIZE,
"weight length must be 2048");
TORCH_CHECK(hidden_states.device() == weight.device(),
"hidden_states and weight must be on same device");
/* Ensure contiguous layout (kernel expects it) */
auto x = hidden_states.contiguous();
auto wt = weight.contiguous();
auto out = torch::empty_like(x);
const int batch = static_cast<int>(x.size(0));
const int device_idx = x.device().index();
cudaStream_t stream = at::cuda::getCurrentCUDAStream(device_idx);
/* Copy weight into GPU constant memory (async) */
load_rmsnorm_weight(reinterpret_cast<const __nv_bfloat16*>(wt.data_ptr()),
stream);
/* Launch CUDA kernel */
launch_rmsnorm_h2048(
reinterpret_cast<const __nv_bfloat16*>(x.data_ptr()),
reinterpret_cast< __nv_bfloat16*>(out.data_ptr()),
batch,
static_cast<float>(eps),
stream);
check_cuda(cudaGetLastError(), "rmsnorm_kernel launch");
return out;
}
/* -------------------------------------------------------------------- */
/* PyBind11 module definition */
/* -------------------------------------------------------------------- */
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
{
m.def("run", &run,
py::arg("hidden_states"),
py::arg("weight"),
py::arg("eps") = 1e-6,
R"pbdoc(
Fixed-size (hidden = 2048) BF16 RMSNorm kernel optimised for NVIDIA B200.
Parameters
----------
hidden_states : torch.Tensor (batch, 2048) – bf16, CUDA
weight : torch.Tensor (2048) – bf16, CUDA
eps : float (default 1e-6)
Returns
-------
torch.Tensor (same shape and dtype as hidden_states)
)pbdoc");
}scrolls · 95 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON