gpt-o3 / cuda85371e
gpt-o3_cuda_85371e · gpt-o3 · cuda · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 88 lines, Apache-2.0, pinned at da91508.
main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-cuda-85371e?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:be52bb4ad122998dad3fdc14b1b2e51b9224e8c09715e2dc698359c35ba9959e
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20
Kernel source
main.cpp88 lines
/*
* main.cpp – PyTorch binding for
*
* C[M,6144] = A[M,4096] · B[6144,4096]ᵀ
*
* Both inputs & output are FP16, accumulation is FP32 for speed +
* sufficient accuracy. The kernel itself lives in kernel.cu.
*/
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include "kernel.h"
namespace py = pybind11;
/* ------------------------------------------------------------------ */
/* Light-weight input checker */
/* ------------------------------------------------------------------ */
static void check_inputs(const torch::Tensor& A,
const torch::Tensor& B)
{
TORCH_CHECK(A.is_cuda() && B.is_cuda(),
"A and B must be CUDA tensors");
TORCH_CHECK(A.dtype() == torch::kHalf && B.dtype() == torch::kHalf,
"A and B must be torch.float16 (half) tensors");
TORCH_CHECK(A.dim() == 2 && B.dim() == 2,
"A and B must be 2-D matrices");
TORCH_CHECK(A.size(1) == 4096,
"A has wrong shape (expected [M,4096], got ",
A.sizes(), ')');
TORCH_CHECK(B.size(0) == 6144 && B.size(1) == 4096,
"B has wrong shape (expected [6144,4096], got ",
B.sizes(), ')');
}
/* ------------------------------------------------------------------ */
/* Python-visible entry point */
/* ------------------------------------------------------------------ */
torch::Tensor run(torch::Tensor A, torch::Tensor B)
{
check_inputs(A, B);
A = A.contiguous();
B = B.contiguous();
const std::int64_t M = A.size(0);
/* allocate output on same device */
auto C = torch::empty({M, 6144},
torch::TensorOptions()
.dtype(torch::kHalf)
.device(A.device()));
cudaStream_t stream =
at::cuda::getCurrentCUDAStream(A.device().index()).stream();
gemm_n_6144_k_4096_launcher(
reinterpret_cast<const __half*>(A.data_ptr<at::Half>()),
reinterpret_cast<const __half*>(B.data_ptr<at::Half>()),
reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
M,
stream);
/* PyTorch takes care of stream-ordering, no explicit sync required */
return C;
}
/* ------------------------------------------------------------------ */
/* pybind11 module definition */
/* ------------------------------------------------------------------ */
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
{
m.doc() = "Specialised GEMM (N=6144, K=4096, FP16 I/O, FP32 acc)";
m.def("run", &run,
py::arg("A"),
py::arg("B"),
R"pbdoc(
Compute **C = A @ B.T**
A : torch.HalfTensor [M, 4096] (CUDA, contiguous)
B : torch.HalfTensor [6144, 4096] (CUDA, contiguous)
Returns **C** in FP16 with shape `[M, 6144]`.
)pbdoc");
}scrolls · 88 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON