Skip to content
KernelIndex
Search⌘K

gpt-5 / cuda0d807b

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Kernel source

main.cpp99 lines
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <ATen/cuda/CUDAGuard.h>
#include <vector>
#include <stdexcept>
#include <string>

#include "kernel.h"

using torch::Tensor;
namespace py = pybind11;

static inline void check_inputs(const Tensor& hidden_states,
                                const Tensor& residual,
                                const Tensor& weight) {
  TORCH_CHECK(hidden_states.dim() == 2, "hidden_states must be 2D [batch_size, 7168]");
  TORCH_CHECK(residual.dim() == 2,      "residual must be 2D [batch_size, 7168]");
  TORCH_CHECK(weight.dim() == 1,        "weight must be 1D [7168]");

  TORCH_CHECK(hidden_states.size(1) == fused_add_rmsnorm_h7168_cfg::HIDDEN_SIZE,
              "hidden_size must be 7168 for this kernel");
  TORCH_CHECK(residual.size(1) == fused_add_rmsnorm_h7168_cfg::HIDDEN_SIZE,
              "hidden_size must be 7168 for this kernel");
  TORCH_CHECK(weight.size(0) == fused_add_rmsnorm_h7168_cfg::HIDDEN_SIZE,
              "weight size must be 7168 for this kernel");

  TORCH_CHECK(hidden_states.scalar_type() == at::kBFloat16,
              "hidden_states must be torch.bfloat16");
  TORCH_CHECK(residual.scalar_type() == at::kBFloat16,
              "residual must be torch.bfloat16");
  TORCH_CHECK(weight.scalar_type() == at::kBFloat16,
              "weight must be torch.bfloat16");
}

// Ensure tensor is on a target device (CUDA), dtype bfloat16, and contiguous.
static inline Tensor to_device_bf16_contig(const Tensor& t, c10::Device device) {
  auto opts = t.options().dtype(at::kBFloat16).device(device);
  if (!(t.device() == device) || t.scalar_type() != at::kBFloat16 || !t.is_contiguous()) {
    return t.to(opts, /*non_blocking=*/true).contiguous();
  }
  return t.contiguous();
}

// Core run implementation (positional and keyword arguments supported via pybind names).
Tensor run(Tensor hidden_states, Tensor residual, Tensor weight) {
  check_inputs(hidden_states, residual, weight);

  const bool inputs_on_cpu =
      hidden_states.is_cpu() && residual.is_cpu() && weight.is_cpu();

  // Select device: if any input is CUDA, use its device; otherwise default to current CUDA device 0.
  c10::Device target_device = c10::Device(c10::kCUDA, 0);
  if (hidden_states.is_cuda()) target_device = hidden_states.device();
  else if (residual.is_cuda()) target_device = residual.device();
  else if (weight.is_cuda()) target_device = weight.device();

  // Move inputs to target CUDA device as needed.
  at::cuda::CUDAGuard device_guard(target_device);
  Tensor h = to_device_bf16_contig(hidden_states, target_device);
  Tensor r = to_device_bf16_contig(residual,       target_device);
  Tensor w = to_device_bf16_contig(weight,         target_device);

  const int64_t batch_size = h.size(0);
  auto out = torch::empty({batch_size, (int64_t)fused_add_rmsnorm_h7168_cfg::HIDDEN_SIZE},
                          h.options().dtype(at::kBFloat16).device(target_device));

  // Current stream for the selected device (Driver API stream type).
  CUstream cu_stream = at::cuda::getCurrentCUDAStream().stream();

  // Convert to CUDA BF16 pointers (reinterpret from PyTorch storage).
  const __nv_bfloat16* h_ptr = reinterpret_cast<const __nv_bfloat16*>(
      h.data_ptr<c10::BFloat16>());
  const __nv_bfloat16* r_ptr = reinterpret_cast<const __nv_bfloat16*>(
      r.data_ptr<c10::BFloat16>());
  const __nv_bfloat16* w_ptr = reinterpret_cast<const __nv_bfloat16*>(
      w.data_ptr<c10::BFloat16>());
  __nv_bfloat16* out_ptr = reinterpret_cast<__nv_bfloat16*>(
      out.data_ptr<c10::BFloat16>());

  // Launch kernel
  launch_fused_add_rmsnorm_h7168(
      h_ptr, r_ptr, w_ptr, out_ptr, static_cast<int>(batch_size),
      reinterpret_cast<CUstream_st*>(cu_stream));

  // If original inputs were on CPU, move result back to CPU
  if (inputs_on_cpu) {
    return out.to(torch::kCPU, /*non_blocking=*/false);
  }
  return out;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.doc() = "Fused Add + RMSNorm kernel (hidden_size=7168, bf16) optimized for B200";
  m.def("run", &run,
        py::arg("hidden_states"),
        py::arg("residual"),
        py::arg("weight"),
        "Run fused_add_rmsnorm_h7168");
}
scrolls · 99 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON