Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / cudaaec5f2

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

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

Kernel source

main.cpp158 lines
#include "kernel.h"

#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime_api.h>
#include <chrono>
#include <limits>
#include <pybind11/pybind11.h>

namespace {

inline void check_inputs(const torch::Tensor& probs,
                         const torch::Tensor& top_k,
                         const torch::Tensor& top_p) {
  TORCH_CHECK(probs.is_cuda(), "probs must be a CUDA tensor");
  TORCH_CHECK(top_k.is_cuda(), "top_k 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_k.dtype() == torch::kInt32, "top_k must be int32");
  TORCH_CHECK(top_p.dtype() == torch::kFloat32, "top_p must be float32");

  TORCH_CHECK(probs.dim() == 2, "probs must be 2D [batch_size, vocab_size]");
  TORCH_CHECK(probs.size(1) == VOCAB_SIZE_V128256,
              "vocab_size must be 128256, got ", probs.size(1));
  TORCH_CHECK(top_k.dim() == 1 && top_k.size(0) == probs.size(0),
              "top_k must be 1D with length equal to batch_size");
  TORCH_CHECK(top_p.dim() == 1 && top_p.size(0) == probs.size(0),
              "top_p must be 1D with length equal to batch_size");

  TORCH_CHECK(probs.device().index() == top_k.device().index() &&
              probs.device().index() == top_p.device().index(),
              "All inputs must be on the same CUDA device");
}

} // namespace

torch::Tensor top_k_top_p_sampling_from_probs_v128256(torch::Tensor probs,
                                                      torch::Tensor top_k,
                                                      torch::Tensor top_p,
                                                      c10::optional<int64_t> seed_opt) {
  check_inputs(probs, top_k, top_p);

  at::cuda::CUDAGuard device_guard(probs.device());
  probs = probs.contiguous();
  top_k = top_k.contiguous();
  top_p = top_p.contiguous();

  const int batch_size = static_cast<int>(probs.size(0));
  const int vocab_size = VOCAB_SIZE_V128256;
  if (batch_size == 0) {
    return torch::empty({0}, probs.options().dtype(torch::kInt64));
  }

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

  const int64_t total_items_64 = static_cast<int64_t>(batch_size) * vocab_size;
  TORCH_CHECK(total_items_64 <= static_cast<int64_t>(std::numeric_limits<int>::max()),
              "Total items exceed supported range");
  const int total_items = static_cast<int>(total_items_64);

  // Allocate global sort buffers
  auto keys_dtype = torch::kInt64; // store uint64_t keys in int64_t tensor
  auto idx_dtype  = torch::kInt32;
  auto keys_in  = torch::empty({total_items_64}, probs.options().dtype(keys_dtype));
  auto keys_out = torch::empty({total_items_64}, probs.options().dtype(keys_dtype));
  auto vals_in  = torch::empty({total_items_64}, probs.options().dtype(idx_dtype));
  auto vals_out = torch::empty({total_items_64}, probs.options().dtype(idx_dtype));

  // Build composite keys and linear indices
  build_composite_keys_and_linidx_launcher(
      probs.data_ptr<float>(),
      reinterpret_cast<uint64_t*>(keys_in.data_ptr<int64_t>()),
      vals_in.data_ptr<int32_t>(),
      total_items,
      vocab_size,
      stream);

  // CUB global radix sort on (composite_key, linidx) pairs (ascending)
  uint64_t* d_keys_in  = reinterpret_cast<uint64_t*>(keys_in.data_ptr<int64_t>());
  uint64_t* d_keys_out = reinterpret_cast<uint64_t*>(keys_out.data_ptr<int64_t>());
  int32_t*  d_vals_in  = vals_in.data_ptr<int32_t>();
  int32_t*  d_vals_out = vals_out.data_ptr<int32_t>();

  size_t temp_bytes = radix_sort_pairs_temp_bytes(
      d_keys_in, d_keys_out, d_vals_in, d_vals_out, total_items, stream);

  auto temp_storage = torch::empty({static_cast<int64_t>(temp_bytes)},
                                   probs.options().dtype(torch::kUInt8));
  void* temp_ptr = static_cast<void*>(temp_storage.data_ptr<uint8_t>());

  radix_sort_pairs_launcher(
      d_keys_in, d_keys_out, d_vals_in, d_vals_out,
      total_items, temp_ptr, temp_bytes, stream);

  // Gather sorted probs and token indices from linidx result
  auto sorted_probs   = torch::empty_like(probs); // [B, V]
  auto sorted_indices = torch::empty({total_items_64}, probs.options().dtype(idx_dtype));

  gather_sorted_from_linidx_launcher(
      probs.data_ptr<float>(),
      d_vals_out,
      sorted_probs.data_ptr<float>(),
      sorted_indices.data_ptr<int32_t>(),
      total_items,
      vocab_size,
      stream);

  // Output samples
  auto samples = torch::empty({batch_size}, probs.options().dtype(torch::kInt64));

  // Seed
  uint64_t seed;
  if (seed_opt.has_value() && seed_opt.value() != 0) {
    seed = static_cast<uint64_t>(seed_opt.value());
  } else {
    seed = static_cast<uint64_t>(
        std::chrono::high_resolution_clock::now().time_since_epoch().count());
  }

  // Launch sampling kernel: one block per row
  sample_from_sorted_launcher(
      sorted_probs.data_ptr<float>(),
      sorted_indices.data_ptr<int32_t>(),
      top_k.data_ptr<int32_t>(),
      top_p.data_ptr<float>(),
      samples.data_ptr<int64_t>(),
      batch_size,
      vocab_size,
      seed,
      stream);

  return samples;
}

// Python binding entry point
torch::Tensor run(torch::Tensor probs,
                  torch::Tensor top_k,
                  torch::Tensor top_p,
                  c10::optional<int64_t> seed) {
  return top_k_top_p_sampling_from_probs_v128256(std::move(probs),
                                                 std::move(top_k),
                                                 std::move(top_p),
                                                 seed);
}

namespace py = pybind11;

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.doc() = "Top-K + Top-P sampling from probability distributions (vocab=128256) optimized for B200";
  m.def("run",
        &run,
        "Run top-k top-p sampling from probs (v128256)",
        py::arg("probs"),
        py::arg("top_k"),
        py::arg("top_p"),
        py::arg("seed") = py::none());
}
scrolls · 158 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON