Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / cudaa6f41d

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Kernel source

main.cpp184 lines
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda.h>
#include <cuda_runtime.h>

#include <vector>
#include <random>
#include <chrono>
#include <stdexcept>
#include <cstring>
#include <limits>
#include <cstdint>

#include "kernel.h"

using namespace at::indexing;

// Helper to check tensor properties
static inline void check_inputs(const torch::Tensor& probs, const torch::Tensor& top_p) {
  TORCH_CHECK(probs.is_cuda(), "probs must be a CUDA tensor");
  TORCH_CHECK(top_p.is_cuda() || top_p.is_cpu(), "top_p must be on CPU or CUDA");
  TORCH_CHECK(probs.dim() == 2, "probs must be 2D [batch, vocab]");
  TORCH_CHECK(probs.scalar_type() == torch::kFloat32, "probs must be float32");
  TORCH_CHECK(probs.size(1) == VOCAB_SIZE_QWEN3, "vocab_size must be 151936");
  TORCH_CHECK(top_p.dim() == 1 && top_p.size(0) == probs.size(0), "top_p must be [batch]");
  TORCH_CHECK(top_p.scalar_type() == torch::kFloat32, "top_p must be float32");
}

// Build host vectors of row indices for each category based on top_p
static inline void partition_rows(const torch::Tensor& top_p, std::vector<int32_t>& rows_argmax,
                                  std::vector<int32_t>& rows_full, std::vector<int32_t>& rows_nucleus) {
  auto top_p_cpu = top_p.to(torch::kCPU, /*non_blocking=*/false);
  const float* tp = top_p_cpu.data_ptr<float>();
  int64_t B = top_p_cpu.size(0);
  rows_argmax.reserve(B);
  rows_full.reserve(B);
  rows_nucleus.reserve(B);

  for (int64_t i = 0; i < B; ++i) {
    float p = tp[i];
    if (!(p > 0.0f)) {
      rows_argmax.push_back(static_cast<int32_t>(i));
    } else if (p >= 1.0f) {
      rows_full.push_back(static_cast<int32_t>(i));
    } else {
      rows_nucleus.push_back(static_cast<int32_t>(i));
    }
  }
}

// Copy host vector to device tensor (int32)
static inline torch::Tensor vec_to_device_i32(const std::vector<int32_t>& v, c10::Device device, cudaStream_t stream) {
  auto opts = torch::TensorOptions().dtype(torch::kInt32).device(device);
  torch::Tensor t = torch::empty({static_cast<long long>(v.size())}, opts);
  if (!v.empty()) {
    CUDA_CHECK(cudaMemcpyAsync(t.data_ptr<int32_t>(), v.data(), v.size() * sizeof(int32_t),
                               cudaMemcpyHostToDevice, stream));
  }
  return t;
}

// Entry point
torch::Tensor run(torch::Tensor probs, torch::Tensor top_p) {
  check_inputs(probs, top_p);

  cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
  const int64_t B = probs.size(0);
  const int64_t V = probs.size(1);
  TORCH_CHECK(V == VOCAB_SIZE_QWEN3, "vocab size mismatch");

  // Ensure contiguous
  probs = probs.contiguous();

  // Make top_p available on device persistently
  torch::Tensor top_p_dev = top_p.is_cuda()
                                ? top_p.contiguous()
                                : top_p.to(probs.device(), /*non_blocking=*/false).contiguous();

  // Output on GPU
  auto out = torch::empty({B}, probs.options().dtype(torch::kInt64));

  // Partition rows
  std::vector<int32_t> rows_argmax, rows_full, rows_nucleus;
  partition_rows(top_p, rows_argmax, rows_full, rows_nucleus);

  // Device copies of row lists
  torch::Tensor d_rows_argmax = vec_to_device_i32(rows_argmax, probs.device(), stream);
  torch::Tensor d_rows_full   = vec_to_device_i32(rows_full,   probs.device(), stream);
  torch::Tensor d_rows_nucleus= vec_to_device_i32(rows_nucleus,probs.device(), stream);

  const float* d_probs = probs.data_ptr<float>();
  const float* d_top_p = top_p_dev.data_ptr<float>();
  int64_t* d_out = out.data_ptr<int64_t>();

  // Seed for RNG used by kernels (full sampling path)
  unsigned long long seed = static_cast<unsigned long long>(
      (std::random_device{}()) ^
      (static_cast<unsigned long long>(std::chrono::high_resolution_clock::now().time_since_epoch().count())));

  // 1) Argmax rows (top_p <= 0)
  if (!rows_argmax.empty()) {
    launch_argmax_kernel(d_probs,
                         d_rows_argmax.data_ptr<int32_t>(),
                         static_cast<int>(rows_argmax.size()),
                         static_cast<int>(V),
                         d_out,
                         stream);
  }

  // 2) Full multinomial rows (top_p >= 1)
  if (!rows_full.empty()) {
    launch_sample_full_kernel(d_probs,
                              d_rows_full.data_ptr<int32_t>(),
                              static_cast<int>(rows_full.size()),
                              static_cast<int>(V),
                              seed,
                              d_out,
                              stream);
  }

  // 3) Nucleus rows (0 < top_p < 1): Use exact PyTorch semantics on GPU to guarantee correctness
  if (!rows_nucleus.empty()) {
    // Gather per-row probabilities and top_p for nucleus rows
    auto sel_rows = d_rows_nucleus; // [R] int32 on device
    auto sel_probs = probs.index_select(0, sel_rows.to(torch::kLong)); // [R, V]
    auto sel_top_p = top_p_dev.index_select(0, sel_rows.to(torch::kLong)).unsqueeze(1); // [R,1]

    // Sort descending like reference
    auto sort_tuple = sel_probs.sort(1, /*descending=*/true);
    auto vals_sorted = std::get<0>(sort_tuple);         // [R, V], descending
    auto idx_sorted  = std::get<1>(sort_tuple).to(torch::kInt32); // [R, V], original indices

    // CDF and keep mask with "shift" semantics
    auto cdf = vals_sorted.cumsum(1);                   // [R, V]
    auto to_remove = cdf.gt(sel_top_p);                 // [R, V], bool

    // Shift mask right by 1, set first column to False
    auto to_remove_shifted = torch::zeros_like(to_remove, to_remove.options().dtype(torch::kBool));
    // to_remove_shifted[:, 1:] = to_remove[:, :-1]
    to_remove_shifted.index_put_({Slice(), Slice(1, None)}, to_remove.index({Slice(), Slice(None, -1)}));
    // to_remove_shifted[:, 0] already False

    auto keep = (~to_remove_shifted);                   // [R, V] bool

    // Build filtered distribution in sorted space and renormalize
    auto keep_f = keep.to(vals_sorted.scalar_type());   // float
    auto probs_kept_sorted = vals_sorted * keep_f;      // [R, V]
    auto sums = probs_kept_sorted.sum(1, true);         // [R, 1]

    // Handle degenerate rows defensively: if sum==0, fall back to picking the top token
    auto safe_sums = torch::where(sums > 0.0, sums, torch::ones_like(sums));
    auto normalized_sorted = probs_kept_sorted / safe_sums;

    // For rows with sum==0, normalized will be zeros; set first position to 1
    // Find rows where sums==0 and fix their normalized distributions
    auto zero_rows_mask = (sums <= 0.0).squeeze(1); // [R] bool
    if (zero_rows_mask.any().item<bool>()) {
      // Build an index tensor for those rows and set normalized_sorted[row, 0] = 1
      auto zero_rows = zero_rows_mask.nonzero().squeeze(1); // [Rz]
      if (zero_rows.numel() > 0) {
        normalized_sorted.index_put_({zero_rows, 0}, 1.0f);
      }
    }

    // Sample indices in sorted space using PyTorch multinomial (GPU)
    auto selected_sorted = torch::multinomial(normalized_sorted, /*num_samples=*/1, /*replacement=*/true); // [R,1]

    // Map back to original vocabulary indices
    auto final_idx = idx_sorted.gather(1, selected_sorted).squeeze(1).to(torch::kInt64); // [R]

    // Scatter into output at the appropriate batch rows
    out.index_copy_(0, sel_rows.to(torch::kLong), final_idx);
  }

  CUDA_CHECK(cudaGetLastError());
  CUDA_CHECK(cudaStreamSynchronize(stream));

  // Return results to CPU
  return out.cpu();
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("run", &run, "top_p_sampling_from_probs_v151936 (CUDA)");
}
scrolls · 184 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON