Skip to content
KernelIndex
Search⌘K

gpt-o3 / cuda2ad247

gpt-o3_cuda_2ad247 · gpt-o3 · cuda · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 81 lines, Apache-2.0, pinned at da91508.

main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-cuda-2ad247?include=source"
interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16

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:adba3de3d00c84b1fb829662a8f93980d9d4caa707cc5a20dcbf072edadcb4e4
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20

Kernel source

main.cpp81 lines
#include "kernel.h"

/* PyTorch / CUDA headers */
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>

namespace py = pybind11;

/* --------------------------------------------------------------------- *
 *  Public entry point visible from Python                                *
 * --------------------------------------------------------------------- */
torch::Tensor run(torch::Tensor A,
                  torch::Tensor B,
                  py::args   /*unused*/ = {},
                  py::kwargs /*unused*/ = {})
{
    /* -------- Accept inputs on either CPU or GPU --------------------- */
    bool inputs_were_cuda = A.is_cuda() && B.is_cuda();

    TORCH_CHECK(A.scalar_type() == torch::kFloat16 &&
                B.scalar_type() == torch::kFloat16,
                "All tensors must be torch.float16.");

    /* If tensors are on CPU, move them to GPU 0 (default device)        */
    torch::Tensor A_d = inputs_were_cuda ? A  : A.to(torch::kCUDA);
    torch::Tensor B_d = inputs_were_cuda ? B  : B.to(torch::kCUDA);

    /* -------- Shape checks ------------------------------------------- */
    TORCH_CHECK(A_d.dim() == 2 && B_d.dim() == 2,
                "All tensors must be 2-D.");
    TORCH_CHECK(A_d.size(1) == GEMM_K,
                "A must have shape [M, ", GEMM_K, "]; got [",
                A_d.size(0), ", ", A_d.size(1), "].");
    TORCH_CHECK(B_d.size(0) == GEMM_N && B_d.size(1) == GEMM_K,
                "B must have shape [", GEMM_N, ", ", GEMM_K, "]; got [",
                B_d.size(0), ", ", B_d.size(1), "].");

    /* -------- Prepare output tensor ---------------------------------- */
    const int64_t M = A_d.size(0);
    torch::Tensor C_d = torch::empty({M, GEMM_N},
                                     A_d.options().dtype(torch::kFloat16));

    /* -------- Invoke GEMM launcher ----------------------------------- */
    cudaStream_t stream =
        at::cuda::getCurrentCUDAStream(A_d.device().index()).stream();
    launch_gemm_n_4096_k_14336(A_d, B_d, C_d, stream);

    /* -------- Move result back to original device if necessary ------- */
    torch::Tensor C = inputs_were_cuda ? C_d : C_d.cpu();
    return C;
}

/* --------------------------------------------------------------------- *
 *  pybind11 bindings                                                     *
 * --------------------------------------------------------------------- */
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
{
    m.doc() = R"pbdoc(
Optimised half-precision GEMM specialised for

    A : [M, 14336]
    B : [4096, 14336]

Computes

    C = A · Bᵀ   →  C ∈ ℝ^{M×4096}
)pbdoc";

    m.def("run",
          &run,
          py::arg("A"),
          py::arg("B"),
          py::arg("args")   = py::args(),
          py::arg("kwargs") = py::kwargs(),
          R"pbdoc(
Launch the fixed-shape GEMM on the current CUDA stream.  If the
inputs live on the CPU, they are transparently copied to the GPU
and the output is copied back before returning.
)pbdoc");
}
scrolls · 81 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON