Skip to content
KernelIndex
Search⌘K

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