Skip to content
KernelIndex
Search⌘K

gpt-5_cuda_927510

gpt-5-2025-08-07 · cuda · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-cuda-927510?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:06a9af868d43f80e54229601d6c10ff470ec07f4682d9c2049d0b0641dfc87ec
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20

Kernel source

main.cpp94 lines
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime.h>
#include "kernel.h"

static inline void check_inputs(const torch::Tensor& hidden_states, const torch::Tensor& weight) {
  TORCH_CHECK(hidden_states.dim() == 2, "hidden_states must be 2D [batch_size, 4096]");
  TORCH_CHECK(hidden_states.size(1) == RMSNORM_H4096, "hidden_states second dim must be 4096");
  TORCH_CHECK(weight.dim() == 1 && weight.size(0) == RMSNORM_H4096, "weight must be 1D of size 4096");
  TORCH_CHECK(hidden_states.dtype() == torch::kBFloat16, "hidden_states dtype must be bfloat16");
  TORCH_CHECK(weight.dtype() == torch::kBFloat16, "weight dtype must be bfloat16");
}

// PyTorch-facing launcher using current CUDA stream
void rmsnorm_h4096_cuda_launcher(const torch::Tensor& hidden_states,
                                 const torch::Tensor& weight,
                                 torch::Tensor& output) {
  TORCH_CHECK(hidden_states.is_cuda(), "hidden_states must be a CUDA tensor");
  TORCH_CHECK(weight.is_cuda(), "weight must be a CUDA tensor");
  TORCH_CHECK(output.is_cuda(), "output must be a CUDA tensor");
  TORCH_CHECK(hidden_states.is_contiguous(), "hidden_states must be contiguous");
  TORCH_CHECK(weight.is_contiguous(), "weight must be contiguous");
  TORCH_CHECK(output.is_contiguous(), "output must be contiguous");
  TORCH_CHECK(hidden_states.size(1) == RMSNORM_H4096, "hidden_size must be 4096");
  TORCH_CHECK(weight.dim() == 1 && weight.size(0) == RMSNORM_H4096, "weight must be [4096]");

  int64_t batch_size = hidden_states.size(0);
  if (batch_size == 0) return;

  cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();

  // Treat BF16 storage as 16-bit lanes
  const uint16_t* hidden_ptr = reinterpret_cast<const uint16_t*>(hidden_states.data_ptr<c10::BFloat16>());
  const uint16_t* weight_ptr = reinterpret_cast<const uint16_t*>(weight.data_ptr<c10::BFloat16>());
  uint16_t* out_ptr = reinterpret_cast<uint16_t*>(output.data_ptr<c10::BFloat16>());

  rmsnorm_h4096_launch(hidden_ptr, weight_ptr, out_ptr, static_cast<int>(batch_size), stream);

  // Post-launch error check
  cudaError_t err = cudaGetLastError();
  TORCH_CHECK(err == cudaSuccess, "rmsnorm_h4096_kernel launch failed: ", cudaGetErrorString(err));
}

// Entry point callable from Python
torch::Tensor run(const torch::Tensor& hidden_states_in, const torch::Tensor& weight_in) {
  check_inputs(hidden_states_in, weight_in);

  // Select device
  int device_index = -1;
  if (hidden_states_in.is_cuda()) {
    device_index = hidden_states_in.get_device();
  } else if (weight_in.is_cuda()) {
    device_index = weight_in.get_device();
  } else {
    device_index = c10::cuda::current_device();
  }

  c10::cuda::CUDAGuard device_guard(device_index);

  // Ensure tensors are on the right device, correct dtype, and contiguous
  torch::Tensor hidden_states = hidden_states_in;
  torch::Tensor weight = weight_in;

  if (!hidden_states.is_cuda() || hidden_states.get_device() != device_index) {
    hidden_states = hidden_states.to(torch::device(torch::kCUDA, device_index));
  }
  if (!weight.is_cuda() || weight.get_device() != device_index) {
    weight = weight.to(torch::device(torch::kCUDA, device_index));
  }

  if (hidden_states.dtype() != torch::kBFloat16) hidden_states = hidden_states.to(torch::kBFloat16);
  if (!hidden_states.is_contiguous()) hidden_states = hidden_states.contiguous();

  if (weight.dtype() != torch::kBFloat16) weight = weight.to(torch::kBFloat16);
  if (!weight.is_contiguous()) weight = weight.contiguous();

  // Allocate output
  torch::Tensor output = torch::empty_like(hidden_states, hidden_states.options().dtype(torch::kBFloat16));

  // Launch CUDA path
  rmsnorm_h4096_cuda_launcher(hidden_states, weight, output);

  // If original inputs were on CPU, move result back to CPU
  if (!hidden_states_in.is_cuda()) {
    return output.to(hidden_states_in.device());
  }
  return output;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("run", &run, "rmsnorm_h4096 CUDA kernel (BF16, B200-optimized)",
        py::arg("hidden_states_in"), py::arg("weight_in"));
}
scrolls · 94 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON