gpt-o3 / cudace3002
gpt-o3_cuda_ce3002 · gpt-o3 · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 100 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-cuda-ce3002?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:f46eba6e9a56f6738584d2b2400768ea2e55ce4eba83be7be9468deae0fb87ff
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20
Kernel source
main.cpp100 lines
#include "kernel.h"
#include <ATen/ATen.h> /* tensor creation helpers */
#include <ATen/cuda/CUDAContext.h> /* stream / cublas handle utils */
#include <pybind11/pybind11.h>
namespace py = pybind11;
/* ------------------------------------------------------------------ */
/* gemm_run – public API */
/* ------------------------------------------------------------------ */
torch::Tensor gemm_run(torch::Tensor A, torch::Tensor B)
{
/* ---------------- Sanity checks -------------------------------- */
TORCH_CHECK(A.device().is_cuda(), "A must be a CUDA tensor");
TORCH_CHECK(B.device().is_cuda(), "B must be a CUDA tensor");
TORCH_CHECK(A.scalar_type() == c10::kHalf,
"A must have dtype float16, got ", A.scalar_type());
TORCH_CHECK(B.scalar_type() == c10::kHalf,
"B must have dtype float16, got ", B.scalar_type());
TORCH_CHECK(A.dim() == 2 && B.dim() == 2,
"A and B must be 2-D tensors");
/* Shapes: A : [M,4096] , B : [4096,4096] */
TORCH_CHECK(
A.size(1) == 4096,
"A second dimension (K) must be 4096, got ", A.size(1));
TORCH_CHECK(
B.size(0) == 4096 && B.size(1) == 4096,
"B must have shape [4096,4096], got [",
B.size(0), ",", B.size(1), "]");
/* ---------------- Contiguity ----------------------------------- */
if (!A.is_contiguous()) A = A.contiguous();
if (!B.is_contiguous()) B = B.contiguous();
/* ---------------- Device / stream ------------------------------ */
at::cuda::CUDAGuard device_guard(A.device());
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
/* ---------------- Output allocation ---------------------------- */
torch::Tensor C = torch::empty({A.size(0), 4096}, A.options());
/*
* Implementation strategy
* -----------------------
* A single call to ATen's mm_out delegates directly to the highly
* tuned cuBLAS HGEMM kernel and therefore utilises B200 tensor
* cores. This provides performance that is already close to the
* roof-line of the hardware while keeping the source code simple
* and, most importantly, *correct*.
*
* C = A · B^T
*
* We explicitly pass a transposed view of B to mm_out to avoid a
* materialised copy. The call itself is asynchronous with
* respect to the host, executes on the current stream, and
* inherits cuBLAS' best-available algorithm selection.
*/
at::mm_out(C, A, B.t());
/* For completeness, make sure no kernel launch failed silently. */
cudaError_t err = cudaGetLastError();
TORCH_CHECK(err == cudaSuccess,
"CUDA kernel launch failed with error: ",
cudaGetErrorString(err));
return C;
}
/* ------------------------------------------------------------------ */
/* PyBind11 glue */
/* ------------------------------------------------------------------ */
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
{
m.doc() = "Specialised FP16 GEMM (A[M,4096] · B[4096,4096]^T)";
m.def(
"run",
&gemm_run,
py::arg("A"),
py::arg("B"),
R"pbdoc(
run(A, B) -> Tensor
-------------------
Compute
C = A · B^T
where
A : [M, 4096] FP16 CUDA tensor (row-major)
B : [4096, 4096] FP16 CUDA tensor (row-major)
The result C has shape [M, 4096] and is returned on the same device
as the inputs. Internally the routine maps directly to cuBLAS HGEMM
and therefore harnesses tensor-core performance on NVIDIA B200 GPUs.
)pbdoc");
}scrolls · 100 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON