Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / cuda5fc7e3

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.cpp
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-cuda-5fc7e3?include=source"
interfacecuda
revisionda915083d4c7
symbolrun
pathmain.cpp
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32

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:ffacdb9d70668fc151cce7587ed7470f231c9b477d5d35a97c0867ecf74a31e0
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20

Kernel source

main.cpp137 lines
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cuda_runtime.h>
#include <vector>
#include <chrono>
#include <stdexcept>
#include "kernel.h"

namespace py = pybind11;

static inline void cuda_check_last(const char* msg) {
  cudaError_t err = cudaGetLastError();
  TORCH_CHECK(err == cudaSuccess, "CUDA error after ", msg, ": ", cudaGetErrorString(err));
}

static inline void cuda_check(cudaError_t err, const char* msg) {
  TORCH_CHECK(err == cudaSuccess, "CUDA error in ", msg, ": ", cudaGetErrorString(err));
}

torch::Tensor run(torch::Tensor probs, torch::Tensor top_p) {
  TORCH_CHECK(probs.is_cuda(), "probs must be a CUDA tensor");
  TORCH_CHECK(top_p.is_cuda(), "top_p must be a CUDA tensor");
  TORCH_CHECK(probs.dtype() == torch::kFloat32, "probs must be float32");
  TORCH_CHECK(top_p.dtype() == torch::kFloat32, "top_p must be float32");
  TORCH_CHECK(probs.dim() == 2, "probs must be 2D [batch, vocab]");
  TORCH_CHECK(top_p.dim() == 1, "top_p must be 1D [batch]");
  const int64_t B = probs.size(0);
  const int64_t V = probs.size(1);
  TORCH_CHECK(V == (int64_t)V_CONST, "vocab_size must be 128256");
  TORCH_CHECK(top_p.size(0) == B, "top_p shape mismatch with batch size");

  c10::cuda::CUDAGuard device_guard(probs.get_device());

  probs = probs.contiguous();
  top_p = top_p.contiguous();

  auto options = torch::TensorOptions().dtype(torch::kInt64).device(probs.device());
  auto samples = torch::empty({B}, options);

  auto cu_stream = at::cuda::getCurrentCUDAStream();

  // Compute status per row on device: 0 (argmax), 1 (top-p), 2 (full)
  auto status = torch::empty({B}, probs.options().dtype(torch::kInt32));
  compute_row_status_launcher(top_p.data_ptr<float>(), status.data_ptr<int32_t>(), B, cu_stream.stream());
  cuda_check_last("compute_row_status");
  cuda_check(cudaStreamSynchronize(cu_stream.stream()), "stream sync after compute_row_status");

  // Copy status and top_p to host to build list of rows needing top-p
  auto status_h = status.cpu();
  auto top_p_h = top_p.cpu();
  const int32_t* status_hp = status_h.data_ptr<int32_t>();
  const float* top_p_hp = top_p_h.data_ptr<float>();

  std::vector<int64_t> rows_top_p;
  rows_top_p.reserve(static_cast<size_t>(B));
  for (int64_t i = 0; i < B; ++i) {
    if (status_hp[i] == 1) rows_top_p.push_back(i);
  }

  // Step 1: p <= 0: argmax
  argmax_kernel_launcher(
      probs.data_ptr<float>(),
      status.data_ptr<int32_t>(),
      B, V,
      samples.data_ptr<int64_t>(),
      cu_stream.stream());
  cuda_check_last("argmax_kernel");

  // Step 2: p >= 1: sample from full distribution
  uint64_t base_seed =
      static_cast<uint64_t>(
          std::chrono::high_resolution_clock::now().time_since_epoch().count());
  sample_full_kernel_launcher(
      probs.data_ptr<float>(),
      status.data_ptr<int32_t>(),
      B, V,
      base_seed,
      samples.data_ptr<int64_t>(),
      cu_stream.stream());
  cuda_check_last("sample_full_kernel");

  // Step 3: 0 < p < 1: top-p via per-row CUB radix sort + sampling
  if (!rows_top_p.empty()) {
    float* d_sorted_keys = nullptr;
    int32_t* d_sorted_idx = nullptr;
    int32_t* d_base_idx = nullptr;
    void* d_cub_tmp = nullptr;

    cuda_check(cudaMalloc(&d_sorted_keys, sizeof(float) * (size_t)V), "cudaMalloc d_sorted_keys");
    cuda_check(cudaMalloc(&d_sorted_idx, sizeof(int32_t) * (size_t)V), "cudaMalloc d_sorted_idx");
    cuda_check(cudaMalloc(&d_base_idx, sizeof(int32_t) * (size_t)V), "cudaMalloc d_base_idx");
    arange_launcher(d_base_idx, V, cu_stream.stream());
    cuda_check_last("arange_launcher");

    size_t tmp_bytes = cub_sort_pairs_temp_bytes(V, cu_stream.stream());
    cuda_check(cudaMalloc(&d_cub_tmp, tmp_bytes), "cudaMalloc d_cub_tmp");

    const float* probs_ptr = probs.data_ptr<float>();
    int64_t* out_ptr = samples.data_ptr<int64_t>();

    for (size_t k = 0; k < rows_top_p.size(); ++k) {
      int64_t row = rows_top_p[k];
      const float* row_ptr = probs_ptr + row * V;

      // Sort this row descending
      sort_row_desc(row_ptr, d_sorted_keys, d_base_idx, d_sorted_idx, V, d_cub_tmp, tmp_bytes, cu_stream.stream());
      cuda_check_last("cub::SortPairsDescending");

      // Sample from top-p subset
      float pval = top_p_hp[row];
      uint64_t seed = base_seed + static_cast<uint64_t>(row * 1315423911ULL);
      top_p_sample_row_kernel_launcher(
          d_sorted_keys, d_sorted_idx,
          pval, V, seed,
          out_ptr + row,
          cu_stream.stream());
      cuda_check_last("top_p_sample_row_kernel");
    }

    cuda_check(cudaFree(d_cub_tmp), "cudaFree d_cub_tmp");
    cuda_check(cudaFree(d_base_idx), "cudaFree d_base_idx");
    cuda_check(cudaFree(d_sorted_idx), "cudaFree d_sorted_idx");
    cuda_check(cudaFree(d_sorted_keys), "cudaFree d_sorted_keys");
  }

  cuda_check_last("final");

  return samples;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("run", &run,
        "top_p_sampling_from_probs_v128256 (B200-optimized)",
        py::arg("probs"),
        py::arg("top_p"));
}
scrolls · 137 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON