#include "torch_utils.h"
#include "dispatch_utils.h"

#include "../cuda_utils.h"
#include "../cuda_compat.h"

#include "quantization/vectorization_utils.cuh"
#include "concat_mla_q.cuh"

#if !defined(USE_ROCM) && defined(ENABLE_NVFP4_SM100) && ENABLE_NVFP4_SM100
  #include "nvfp4_ds_mla_cache.h"
#endif

#ifdef USE_ROCM
  #include "../quantization/w8a8/fp8/amd/quant_utils.cuh"
#else
  #include "../quantization/w8a8/fp8/nvidia/quant_utils.cuh"
#endif

#include <algorithm>
#include <cassert>
#include <cstdlib>

#ifdef USE_ROCM
  #include <hip/hip_bf16.h>
typedef __hip_bfloat16 __nv_bfloat16;
#else
  #include <cuda.h>
#endif

#if defined(__gfx942__)
constexpr float kFp8ScaleDivisor = 224.f;
#else
constexpr float kFp8ScaleDivisor = 448.f;
#endif

void swap_blocks(torch::stable::Tensor& src, torch::stable::Tensor& dst,
                 int64_t block_size_in_bytes,
                 const torch::stable::Tensor& block_mapping) {
  torch::stable::Device src_device = src.device();
  torch::stable::Device dst_device = dst.device();
  cudaMemcpyKind memcpy_type;
  if (src_device.is_cuda() && dst_device.is_cuda()) {
    STD_TORCH_CHECK(src_device.index() == dst_device.index(),
                    "src and dst must be on the same GPU");
    memcpy_type = cudaMemcpyDeviceToDevice;
  } else if (src_device.is_cuda() && dst_device.is_cpu()) {
    memcpy_type = cudaMemcpyDeviceToHost;
  } else if (src_device.is_cpu() && dst_device.is_cuda()) {
    memcpy_type = cudaMemcpyHostToDevice;
  } else {
    STD_TORCH_CHECK(false, "Invalid device combination");
  }

  // NOTE(youkaichao): keep in mind that `block_mapping` should be
  // a cpu tensor, otherwise every `item` call will require a gpu-cpu
  // synchronization.
  STD_TORCH_CHECK(block_mapping.device().is_cpu(),
                  "block_mapping must be on CPU");

  char* src_ptr = static_cast<char*>(src.data_ptr());
  char* dst_ptr = static_cast<char*>(dst.data_ptr());

  auto guard_device = src_device.is_cuda() ? src_device : dst_device;
  const torch::stable::accelerator::DeviceGuard device_guard(
      guard_device.index());
  const cudaStream_t stream = get_current_cuda_stream();
  // NOTE(woosuk): This can be slow if the number of blocks is large.
  const int64_t num_blocks = block_mapping.size(0);
  const int64_t* bm_ptr = block_mapping.const_data_ptr<int64_t>();
  const int64_t bm_stride0 = block_mapping.stride(0);
  const int64_t bm_stride1 = block_mapping.stride(1);
  for (size_t i = 0; i < num_blocks; i++) {
    int64_t src_block_number = bm_ptr[i * bm_stride0];
    int64_t dst_block_number = bm_ptr[i * bm_stride0 + bm_stride1];
    int64_t src_offset = src_block_number * block_size_in_bytes;
    int64_t dst_offset = dst_block_number * block_size_in_bytes;
    cudaMemcpyAsync(dst_ptr + dst_offset, src_ptr + src_offset,
                    block_size_in_bytes, memcpy_type, stream);
  }
}

namespace {
// ROCm hipMemcpyBatchAsync faults for count > 8192 (MI350X/gfx950), so chunk
// at that ceiling.
#if defined(USE_ROCM)
constexpr int64_t kRocmDefaultMaxBatchDescriptors = 8192;
#endif

constexpr const char* kMaxBatchDescriptorsEnv =
    "VLLM_KV_OFFLOAD_MAX_BATCH_DESCRIPTORS";

// Max copy descriptors per batch-memcpy call (0 = unlimited). ROCm caps and
// chunks; CUDA is uncapped. The env var (>0) overrides on any platform.
int64_t resolve_max_batch_descriptors() {
  static const int64_t cached = []() -> int64_t {
    const char* val = std::getenv(kMaxBatchDescriptorsEnv);
    const int64_t override_val = val ? std::atoll(val) : 0;
    if (override_val > 0) return override_val;
#if defined(USE_ROCM)
    return kRocmDefaultMaxBatchDescriptors;
#else
    return 0;
#endif
  }();
  return cached;
}
}  // namespace

void swap_blocks_batch(const torch::stable::Tensor& src_ptrs,
                       const torch::stable::Tensor& dst_ptrs,
                       const torch::stable::Tensor& sizes,
                       bool is_src_access_order_any) {
  STD_TORCH_CHECK(src_ptrs.device().is_cpu(), "src_ptrs must be on CPU");
  STD_TORCH_CHECK(dst_ptrs.device().is_cpu(), "dst_ptrs must be on CPU");
  STD_TORCH_CHECK(sizes.device().is_cpu(), "sizes must be on CPU");
  STD_TORCH_CHECK(src_ptrs.scalar_type() == torch::headeronly::ScalarType::Long,
                  "src_ptrs must be int64");
  STD_TORCH_CHECK(dst_ptrs.scalar_type() == torch::headeronly::ScalarType::Long,
                  "dst_ptrs must be int64");
  STD_TORCH_CHECK(sizes.scalar_type() == torch::headeronly::ScalarType::Long,
                  "sizes must be int64");

  const int64_t n = src_ptrs.size(0);
  STD_TORCH_CHECK(dst_ptrs.size(0) == n, "dst_ptrs length must match src_ptrs");
  STD_TORCH_CHECK(sizes.size(0) == n, "sizes length must match src_ptrs");

  if (n == 0) return;

  int64_t* src_data = src_ptrs.mutable_data_ptr<int64_t>();
  int64_t* dst_data = dst_ptrs.mutable_data_ptr<int64_t>();
  int64_t* size_data = sizes.mutable_data_ptr<int64_t>();

  const cudaStream_t stream = get_current_cuda_stream();

  // Use cuMemcpyBatchAsync / hipMemcpyBatchAsync to submit all copies in a
  // single driver call, amortizing per-copy submission overhead. int64_t
  // and CUdeviceptr/void*/size_t are all 8 bytes on 64-bit platforms, so we
  // reinterpret_cast the tensor data directly to avoid copies.
  static_assert(sizeof(size_t) == sizeof(int64_t));
#if !defined(USE_ROCM) && defined(CUDA_VERSION) && CUDA_VERSION >= 12080
  static_assert(sizeof(CUdeviceptr) == sizeof(int64_t));
  // Resolve cuMemcpyBatchAsync at runtime via cuGetProcAddress so that
  // binaries compiled with CUDA 12.8+ still work on older drivers, and
  // we avoid the CUDA 13.0 header remapping (#define to _v2 signature).
  // The function pointer is cached after the first call.
  using BatchFn =
      CUresult (*)(CUdeviceptr*, CUdeviceptr*, size_t*, size_t,
                   CUmemcpyAttributes*, size_t*, size_t, size_t*, CUstream);
  static BatchFn batch_fn = []() -> BatchFn {
    CUdriverProcAddressQueryResult sym_status;
    void* fn_ptr = nullptr;
    CUresult res = cuGetProcAddress("cuMemcpyBatchAsync", &fn_ptr, 12080,
                                    CU_GET_PROC_ADDRESS_DEFAULT, &sym_status);
    if (res != CUDA_SUCCESS || fn_ptr == nullptr) {
      return nullptr;
    }
    return reinterpret_cast<BatchFn>(fn_ptr);
  }();

  // cuMemcpyBatchAsync rejects the legacy default stream (handle 0 /
  // cudaStreamLegacy) with CUDA_ERROR_INVALID_VALUE; route it to the per-copy
  // fallback below, which is correct on any stream. Real and per-thread-default
  // streams take the batch fast path.
  const bool usable_stream = stream != nullptr && stream != cudaStreamLegacy;
  if (batch_fn != nullptr && usable_stream) {
    CUmemcpyAttributes attr = {};
    // ANY lets the DMA engine prefetch source bytes out of stream order,
    // which is only safe when no GPU stream is concurrently writing the
    // source.
    attr.srcAccessOrder = is_src_access_order_any
                              ? CU_MEMCPY_SRC_ACCESS_ORDER_ANY
                              : CU_MEMCPY_SRC_ACCESS_ORDER_STREAM;
    size_t attrs_idx = 0;
    size_t fail_idx = 0;
    // Uncapped on CUDA (max_desc == 0 -> single call) unless overridden by
    // VLLM_KV_OFFLOAD_MAX_BATCH_DESCRIPTORS; chunk to honor the override.
    const int64_t max_desc = resolve_max_batch_descriptors();
    const int64_t step = max_desc <= 0 ? n : max_desc;
    for (int64_t off = 0; off < n; off += step) {
      const int64_t cnt = std::min(step, n - off);
      CUresult result = batch_fn(reinterpret_cast<CUdeviceptr*>(dst_data + off),
                                 reinterpret_cast<CUdeviceptr*>(src_data + off),
                                 reinterpret_cast<size_t*>(size_data + off),
                                 static_cast<size_t>(cnt), &attr, &attrs_idx, 1,
                                 &fail_idx, static_cast<CUstream>(stream));
      STD_TORCH_CHECK(result == CUDA_SUCCESS,
                      "cuMemcpyBatchAsync failed at index ", fail_idx,
                      " with error ", result);
    }
    return;
  }
#elif defined(USE_ROCM) && defined(HIP_VERSION) && HIP_VERSION >= 70100000
  // ROCm 7.1+ exposes hipMemcpyBatchAsync. ROCm 7.2.1-7.2.3 early-return
  // hipErrorNotSupported whenever numAttrs > 0 (see ROCm/clr @ rocm-7.2.1
  // hipamd/src/hip_memory.cpp:2819-2822), so those releases must call with
  // numAttrs=0. ROCm 7.13+ accepts numAttrs > 0.
  // rocm-7.14+ has better performance.
  {
    hipMemcpyAttributes attr = {};
    size_t attrs_idx = 0;
    size_t fail_idx = 0;
    size_t num_attrs = 0;
  #if HIP_VERSION >= 71300000
    static const bool runtime_accepts_attrs = []() {
      int runtime_version = 0;
      return hipRuntimeGetVersion(&runtime_version) == hipSuccess &&
             runtime_version >= 71300000;
    }();
    if (runtime_accepts_attrs) {
      attr.srcAccessOrder = is_src_access_order_any
                                ? hipMemcpySrcAccessOrderAny
                                : hipMemcpySrcAccessOrderStream;
      num_attrs = 1;
    }
  #endif
    // hipMemcpyBatchAsync faults (GPU memory access fault) above 8192
    // descriptors/call on ROCm 7.15, so chunk the batch at the resolved cap.
    const int64_t max_desc = resolve_max_batch_descriptors();
    const int64_t step = max_desc <= 0 ? n : max_desc;
    for (int64_t off = 0; off < n; off += step) {
      const int64_t cnt = std::min(step, n - off);
      hipError_t result = hipMemcpyBatchAsync(
          reinterpret_cast<void**>(dst_data + off),
          reinterpret_cast<void**>(src_data + off),
          reinterpret_cast<size_t*>(size_data + off), static_cast<size_t>(cnt),
          &attr, &attrs_idx, num_attrs, &fail_idx,
          static_cast<hipStream_t>(stream));
      STD_TORCH_CHECK(result == hipSuccess,
                      "hipMemcpyBatchAsync failed at index ", fail_idx,
                      " with error ", result);
    }
    return;
  }
#endif
  {
    // Fallback for CUDA < 12.8, older CUDA drivers, and ROCm < 7.1:
    // individual async copies. cudaMemcpyDefault lets the driver infer
    // direction from pointer types.
    for (int64_t i = 0; i < n; i++) {
      cudaMemcpyAsync(reinterpret_cast<void*>(dst_data[i]),
                      reinterpret_cast<void*>(src_data[i]),
                      static_cast<size_t>(size_data[i]), cudaMemcpyDefault,
                      stream);
    }
  }
}

namespace vllm {

// Used to copy/convert one element
template <typename OutT, typename InT, Fp8KVCacheDataType kv_dt>
struct CopyWithScaleOp {
  float scale;

  __device__ __forceinline__ void operator()(OutT& dst, const InT src) const {
    if constexpr (kv_dt == Fp8KVCacheDataType::kAuto) {
      dst = static_cast<OutT>(src);
    } else {
      dst = fp8::scaled_convert<OutT, InT, kv_dt>(src, scale);
    }
  }
};

template <typename scalar_t, typename cache_t, Fp8KVCacheDataType kv_dt>
__global__ void reshape_and_cache_kernel(
    const scalar_t* __restrict__ key,    // [num_tokens, num_heads, head_size]
    const scalar_t* __restrict__ value,  // [num_tokens, num_heads, head_size]
    cache_t* __restrict__ key_cache,     // [num_blocks, num_heads, head_size/x,
                                         // block_size, x]
    cache_t* __restrict__ value_cache,   // [num_blocks, num_heads, head_size,
                                         // block_size]
    const int64_t* __restrict__ slot_mapping,  // [num_tokens]
    const int key_stride, const int value_stride, const int num_heads,
    const int head_size, const int block_size, const int x,
    const float* k_scale, const float* v_scale) {
  const int64_t token_idx = blockIdx.x;
  const int64_t slot_idx = slot_mapping[token_idx];
  if (slot_idx < 0) {
    return;
  }

  const int64_t block_idx = slot_idx / block_size;
  const int64_t block_offset = slot_idx % block_size;
  const int h_block_count = head_size / x;  // head_size//x

  const int h_block_idx = threadIdx.x;
  if (h_block_idx >= num_heads * h_block_count) {
    return;
  }

  const int head_idx = h_block_idx / h_block_count;
  const int h_block = h_block_idx % h_block_count;

  const scalar_t* __restrict__ key_src =
      key + token_idx * key_stride + head_idx * head_size + h_block * x;
  const int64_t src_value_start =
      token_idx * value_stride + head_idx * head_size + h_block * x;

  cache_t* __restrict__ key_dst =
      key_cache + block_idx * num_heads * h_block_count * block_size * x +
      head_idx * h_block_count * block_size * x + h_block * block_size * x +
      block_offset * x;
  const int64_t tgt_value_start =
      block_idx * num_heads * h_block_count * x * block_size +
      head_idx * h_block_count * x * block_size + h_block * x * block_size +
      block_offset;

  constexpr int VEC_SIZE = (sizeof(scalar_t) == 2) ? 8 : 4;
  float k_scale_val = (kv_dt == Fp8KVCacheDataType::kAuto) ? 0.f : *k_scale;
  CopyWithScaleOp<cache_t, scalar_t, kv_dt> k_op{k_scale_val};
  float v_scale_val = (kv_dt == Fp8KVCacheDataType::kAuto) ? 0.f : *v_scale;
  CopyWithScaleOp<cache_t, scalar_t, kv_dt> v_op{v_scale_val};

  vectorize_with_alignment<VEC_SIZE>(key_src, key_dst, x, 0, 1, k_op);

  const scalar_t* __restrict__ value_src = value + src_value_start;
  cache_t* __restrict__ value_dst = value_cache + tgt_value_start;
#pragma unroll
  for (int i = 0; i < x; i++) {
    v_op(value_dst[i * block_size], value_src[i]);
  }
}

template <typename scalar_t, typename cache_t, Fp8KVCacheDataType kv_dt>
__global__ void reshape_and_cache_flash_kernel(
    const scalar_t* __restrict__ key,    // [num_tokens, num_heads, head_size]
    const scalar_t* __restrict__ value,  // [num_tokens, num_heads, head_size]
    cache_t* __restrict__ key_cache,     // NHD or HND, shape see comments below
    cache_t* __restrict__ value_cache,   // same above
    const int64_t* __restrict__ slot_mapping,  // [num_tokens]
    const int64_t block_stride, const int64_t page_stride,
    const int64_t head_stride, const int64_t key_stride,
    const int64_t value_stride, const int num_heads, const int head_size,
    const int block_size, const float* k_scale, const float* v_scale,
    const int kv_scale_stride) {
  const int64_t token_idx = blockIdx.x;
  const int64_t slot_idx = slot_mapping[token_idx];
  // NOTE: slot_idx can be -1 if the token is padded
  if (slot_idx < 0) {
    return;
  }
  const int64_t block_idx = slot_idx / block_size;
  const int64_t block_offset = slot_idx % block_size;
  const int n_elems = num_heads * head_size;

  // pointers to the beginning of the source row for this token.
  const scalar_t* __restrict__ key_src = key + token_idx * key_stride;
  const scalar_t* __restrict__ value_src = value + token_idx * value_stride;

  // find the start position inside the kv-cache for this token.
  cache_t* __restrict__ key_dst =
      key_cache + block_idx * block_stride + block_offset * page_stride;
  cache_t* __restrict__ value_dst =
      value_cache + block_idx * block_stride + block_offset * page_stride;

  // this is true for the NHD layout where `head_stride == head_size`
  const bool is_contiguous_heads = (head_stride == head_size);

  constexpr int VEC_SIZE = (sizeof(scalar_t) == 2) ? 8 : 4;

  if (is_contiguous_heads && kv_scale_stride == 0) {
    // NHD layout and k/v_scales are [1] (i.e. single scale for all heads)
    // kv cache: [num_blocks, block_size, num_heads, head_size]
    float k_scale_val = (kv_dt == Fp8KVCacheDataType::kAuto) ? 0.f : *k_scale;
    float v_scale_val = (kv_dt == Fp8KVCacheDataType::kAuto) ? 0.f : *v_scale;

    CopyWithScaleOp<cache_t, scalar_t, kv_dt> k_op{k_scale_val};
    CopyWithScaleOp<cache_t, scalar_t, kv_dt> v_op{v_scale_val};

    vectorize_with_alignment<VEC_SIZE>(key_src, key_dst, n_elems, threadIdx.x,
                                       blockDim.x, k_op);
    vectorize_with_alignment<VEC_SIZE>(value_src, value_dst, n_elems,
                                       threadIdx.x, blockDim.x, v_op);
  } else {
    // HND layout OR k/v_scales are [num_heads] (i.e. per-attn-head)
    // HND layout: heads are strided, but each head_size segment is contiguous
    // kv cache: [num_blocks, num_heads, block_size, head_size]
    const int lane = threadIdx.x & 31;     // 0..31 within warp
    const int warp_id = threadIdx.x >> 5;  // warp index within block
    const int warps_per_block = blockDim.x >> 5;

    for (int head = warp_id; head < num_heads; head += warps_per_block) {
      const scalar_t* __restrict__ k_src_h = key_src + head * head_size;
      const scalar_t* __restrict__ v_src_h = value_src + head * head_size;

      cache_t* __restrict__ k_dst_h =
          key_dst + static_cast<int64_t>(head) * head_stride;
      cache_t* __restrict__ v_dst_h =
          value_dst + static_cast<int64_t>(head) * head_stride;

      float k_scale_val = (kv_dt == Fp8KVCacheDataType::kAuto)
                              ? 0.f
                              : k_scale[head * kv_scale_stride];
      float v_scale_val = (kv_dt == Fp8KVCacheDataType::kAuto)
                              ? 0.f
                              : v_scale[head * kv_scale_stride];

      CopyWithScaleOp<cache_t, scalar_t, kv_dt> k_op{k_scale_val};
      CopyWithScaleOp<cache_t, scalar_t, kv_dt> v_op{v_scale_val};

      // within each head, let the 32 threads of the warp perform the vector
      // copy
      vectorize_with_alignment<VEC_SIZE>(k_src_h, k_dst_h, head_size, lane, 32,
                                         k_op);

      vectorize_with_alignment<VEC_SIZE>(v_src_h, v_dst_h, head_size, lane, 32,
                                         v_op);
    }
  }
}

template <typename scalar_t, typename cache_t, Fp8KVCacheDataType kv_dt>
__global__ void concat_and_cache_mla_kernel(
    const scalar_t* __restrict__ kv_c,  // [num_tokens, kv_lora_rank]
    const scalar_t* __restrict__ k_pe,  // [num_tokens, pe_dim]
    cache_t* __restrict__ kv_cache,  // [num_blocks, block_size, (kv_lora_rank
                                     // + pe_dim)]
    const int64_t* __restrict__ slot_mapping,  // [num_tokens]
    const int block_stride,                    //
    const int entry_stride,                    //
    const int kv_c_stride,                     //
    const int k_pe_stride,                     //
    const int kv_lora_rank,                    //
    const int pe_dim,                          //
    const int block_size,                      //
    const float* scale                         //
) {
  const int64_t token_idx = blockIdx.x;
  const int64_t slot_idx = slot_mapping[token_idx];
  // NOTE: slot_idx can be -1 if the token is padded
  if (slot_idx < 0) {
    return;
  }
  const int64_t block_idx = slot_idx / block_size;
  const int64_t block_offset = slot_idx % block_size;

  auto copy = [&](const scalar_t* __restrict__ src, cache_t* __restrict__ dst,
                  int src_stride, int dst_stride, int size, int offset) {
    for (int i = threadIdx.x; i < size; i += blockDim.x) {
      const int64_t src_idx = token_idx * src_stride + i;
      const int64_t dst_idx =
          block_idx * block_stride + block_offset * entry_stride + i + offset;
      if constexpr (kv_dt == Fp8KVCacheDataType::kAuto) {
        dst[dst_idx] = src[src_idx];
      } else {
        dst[dst_idx] =
            fp8::scaled_convert<cache_t, scalar_t, kv_dt>(src[src_idx], *scale);
      }
    }
  };

  copy(kv_c, kv_cache, kv_c_stride, block_stride, kv_lora_rank, 0);
  copy(k_pe, kv_cache, k_pe_stride, block_stride, pe_dim, kv_lora_rank);
}

// Grouped variant of concat_and_cache_mla: inserts the context K/V for every
// draft layer in a single launch. Grid is (num_tokens, num_layers); each layer
// reads its own cache base pointer from kv_cache_ptrs (same pointer-array
// pattern as copy_blocks_kernel).
template <typename scalar_t, typename cache_t, Fp8KVCacheDataType kv_dt>
__global__ void concat_and_cache_mla_grouped_kernel(
    const scalar_t* __restrict__ kv_c,  // [num_layers, num_tokens,
                                        // kv_lora_rank]
    const scalar_t* __restrict__ k_pe,  // [num_layers, num_tokens, pe_dim]
    const int64_t* __restrict__ kv_cache_ptrs,  // [num_layers]
    const float* __restrict__ kv_scales,        // [num_layers] or nullptr
    const int64_t* __restrict__ slot_mapping,   // [num_layers, num_tokens]
    const int64_t kv_c_layer_stride, const int64_t kv_c_token_stride,
    const int64_t k_pe_layer_stride, const int64_t k_pe_token_stride,
    const int64_t slot_layer_stride, const int64_t block_stride,
    const int64_t entry_stride, const int kv_lora_rank, const int pe_dim,
    const int block_size) {
  const int64_t token_idx = blockIdx.x;
  const int64_t layer_idx = blockIdx.y;
  const int64_t slot_idx =
      slot_mapping[layer_idx * slot_layer_stride + token_idx];
  // NOTE: slot_idx can be -1 if the token is padded
  if (slot_idx < 0) {
    return;
  }
  const int64_t block_idx = slot_idx / block_size;
  const int64_t block_offset = slot_idx % block_size;

  cache_t* __restrict__ kv_cache =
      reinterpret_cast<cache_t*>(kv_cache_ptrs[layer_idx]);
  const scalar_t* __restrict__ kv_c_layer =
      kv_c + layer_idx * kv_c_layer_stride;
  const scalar_t* __restrict__ k_pe_layer =
      k_pe + layer_idx * k_pe_layer_stride;
  float scale = 0.0f;
  if constexpr (kv_dt != Fp8KVCacheDataType::kAuto) {
    scale = kv_scales[layer_idx];
  }

  auto copy = [&](const scalar_t* __restrict__ src, int64_t src_token_stride,
                  int size, int offset) {
    for (int i = threadIdx.x; i < size; i += blockDim.x) {
      const int64_t src_idx = token_idx * src_token_stride + i;
      const int64_t dst_idx =
          block_idx * block_stride + block_offset * entry_stride + i + offset;
      if constexpr (kv_dt == Fp8KVCacheDataType::kAuto) {
        kv_cache[dst_idx] = src[src_idx];
      } else {
        kv_cache[dst_idx] =
            fp8::scaled_convert<cache_t, scalar_t, kv_dt>(src[src_idx], scale);
      }
    }
  };

  copy(kv_c_layer, kv_c_token_stride, kv_lora_rank, 0);
  copy(k_pe_layer, k_pe_token_stride, pe_dim, kv_lora_rank);
}

template <typename scalar_t, typename cache_t, Fp8KVCacheDataType kv_dt>
__global__ void concat_and_cache_ds_mla_kernel(
    const scalar_t* __restrict__ kv_c,  // [num_tokens, kv_lora_rank]
    const scalar_t* __restrict__ k_pe,  // [num_tokens, pe_dim]
    cache_t* __restrict__ kv_cache,  // [num_blocks, block_size, (kv_lora_rank
                                     // + pe_dim)]
    const int64_t* __restrict__ slot_mapping,  // [num_tokens]
    const int block_stride,                    //
    const int entry_stride,                    //
    const int kv_c_stride,                     //
    const int k_pe_stride,                     //
    const int kv_lora_rank,                    //
    const int pe_dim,                          //
    const int block_size,                      //
    const float* scale                         //
) {
  const int64_t token_idx = blockIdx.x;
  const int64_t slot_idx = slot_mapping[token_idx];
  // NOTE: slot_idx can be -1 if the token is padded
  if (slot_idx < 0) {
    return;
  }
  const int64_t block_idx = slot_idx / block_size;
  const int64_t block_offset = slot_idx % block_size;
  const int64_t dst_idx_start =
      block_idx * block_stride + block_offset * entry_stride;

  // For the NoPE part, each tile of 128 elements is handled by half of one warp
  // (16 threads). There are 4 total tiles, so 2 warps (64 threads).
  // Lanes 0 and 16 of each warp write the scale values for that warp's tiles.
  // The RoPE part (last 64 elements) is handled by another 1 warp (32 threads).
  // So in total, we use 3 warps (96 threads) per block.

  // Cast kv_cache to 16_bit for RoPE values
  scalar_t* kv_cache_16bit =
      reinterpret_cast<scalar_t*>(&kv_cache[dst_idx_start]);

  // Zero the reserved RoPE tail for NoPE rows.
  if (threadIdx.x >= 64) {
    // Each thread handles two elements of RoPE
    const int8_t pe_idx_start = (threadIdx.x - 64) * 2;
    // RoPE values start after the packed 8-bit NoPE values and the
    // 32-bit scales
    const int64_t dst_idx = kv_lora_rank / 2 + 8 + pe_idx_start;
    if (pe_dim == 0) {
      *reinterpret_cast<int32_t*>(&kv_cache_16bit[dst_idx]) = 0;
      return;
    }
    const int64_t src_idx = token_idx * k_pe_stride + pe_idx_start;
    // Vectorized load of two 16-bit values, performed as one 32-bit load
    const int32_t vals = *reinterpret_cast<const int32_t*>(&k_pe[src_idx]);
    // Vectorized store of two 16-bit values, performed as one 32-bit store
    *reinterpret_cast<int32_t*>(&kv_cache_16bit[dst_idx]) = vals;
    return;
  }

  // The first two warps handle the NoPE part
  const int8_t warp_idx = threadIdx.x >> 5;
  const int8_t lane_idx = threadIdx.x & 31;
  const int8_t tile_idx = warp_idx * 2 + (lane_idx >> 4);

  // Each thread handles 8 elements of NoPE
  // Load the NoPE elements for this thread into registers
  const int64_t src_idx_start = token_idx * kv_c_stride + (threadIdx.x * 8);
  // Vectorized load of eight 16-bit values, performed as an int4 load
  const int4 vals_i4 = *reinterpret_cast<const int4*>(&kv_c[src_idx_start]);
  const scalar_t* vals = reinterpret_cast<const scalar_t*>(&vals_i4);

  // Max absolute value of this thread's elements
  float max_abs = fmaxf(fmaxf(fmaxf(fabsf(vals[0]), fabsf(vals[1])),
                              fmaxf(fabsf(vals[2]), fabsf(vals[3]))),
                        fmaxf(fmaxf(fabsf(vals[4]), fabsf(vals[5])),
                              fmaxf(fabsf(vals[6]), fabsf(vals[7]))));

  // Warp-level reduction to find the max absolute value in each half-warp
#pragma unroll
  for (int offset = 8; offset > 0; offset /= 2) {
    max_abs = fmaxf(max_abs, VLLM_SHFL_XOR_SYNC_WIDTH(max_abs, offset, 16));
  }

  // Both SM90 and SM100 readers preserve power-of-two fp32 scales exactly.
  float tile_scale = fmaxf(max_abs / kFp8ScaleDivisor, 1e-4f);
  tile_scale = exp2f(ceilf(log2f(tile_scale)));

  // The first lane of each half-warp writes the scale to kv_cache
  if ((lane_idx == 0) || (lane_idx == 16)) {
    float* kv_cache_32bit = reinterpret_cast<float*>(&kv_cache[dst_idx_start]);
    const uint64_t dst_idx = kv_lora_rank / 4 + tile_idx;
    kv_cache_32bit[dst_idx] = tile_scale;
  }

  // Now all threads in the block scale and write their elements
  // NoPE data occupies the first kv_lora_rank bytes (512 bytes).
  const int64_t dst_idx_base = dst_idx_start + (threadIdx.x * 8);

  uint8_t result[8];
#pragma unroll
  for (int i = 0; i < 8; i++) {
    result[i] =
        fp8::scaled_convert<uint8_t, scalar_t, Fp8KVCacheDataType::kFp8E4M3>(
            vals[i], tile_scale);
  }

  // Store as aligned 64-bit writes
  *reinterpret_cast<uint64_t*>(&kv_cache[dst_idx_base]) =
      *reinterpret_cast<const uint64_t*>(result);
}

template <typename scalar_t, typename cache_t, Fp8KVCacheDataType kv_dt>
__global__ void indexer_k_quant_and_cache_kernel(
    const scalar_t* __restrict__ k,  // [num_tokens, head_dim]
    cache_t* __restrict__ kv_cache,  // [num_blocks, block_size, cache_stride]
    const int64_t* __restrict__ slot_mapping,  // [num_tokens]
    const int head_dim,                        // dimension of each head
    const int quant_block_size,                // quantization block size
    const int cache_block_size,                // cache block size
    const int64_t cache_block_stride,  // stride for each block in kv_cache

    const bool use_ue8m0  // use ue8m0 scale format
) {
  constexpr int VEC_SIZE = 4;
  const int64_t token_idx = blockIdx.x;
  const int64_t head_dim_idx = (blockIdx.y * blockDim.y * blockDim.x +
                                threadIdx.y * blockDim.x + threadIdx.x) *
                               VEC_SIZE;
  const int64_t slot_idx = slot_mapping[token_idx];
  const int64_t block_idx = slot_idx / cache_block_size;
  const int64_t block_offset = slot_idx % cache_block_size;

  // NOTE: slot_idx can be -1 if the token is padded
  if (slot_idx < 0 || (head_dim_idx >= head_dim)) {
    return;
  }

  float2 k_val = (reinterpret_cast<const float2*>(
      k))[(token_idx * head_dim + head_dim_idx) / VEC_SIZE];
  scalar_t* k_val_ptr = reinterpret_cast<scalar_t*>(&k_val);
  float amax = 0.0f;
  for (int i = 0; i < VEC_SIZE; i++) {
    amax = fmaxf(amax, fabsf(float(k_val_ptr[i])));
  }

  // Reduced amax
  for (int mask = 16; mask > 0; mask /= 2) {
#ifdef USE_ROCM
    amax = fmaxf(amax, __shfl_xor_sync(uint64_t(-1), amax, mask));
#else
    amax = fmaxf(amax, __shfl_xor_sync(unsigned(-1), amax, mask));
#endif
  }

  float scale = fmaxf(amax, 1e-4) / kFp8ScaleDivisor;

  if (use_ue8m0) {
    scale = exp2f(ceilf(log2f(scale)));
  }

  const int64_t dst_offset =
      block_idx * cache_block_stride + block_offset * head_dim + head_dim_idx;
  for (int i = 0; i < VEC_SIZE; i++) {
    kv_cache[dst_offset + i] =
        fp8::scaled_convert<cache_t, scalar_t, kv_dt>(k_val_ptr[i], scale);
  }
  if (threadIdx.x == 0) {
    const int64_t dst_scale_idx =
        block_idx * cache_block_stride + cache_block_size * head_dim +
        (block_offset * head_dim + head_dim_idx) * 4 / quant_block_size;
    reinterpret_cast<float*>(kv_cache)[dst_scale_idx / 4] = scale;
  }
}

template <int BLOCK_Y_SIZE>
__global__ void cp_gather_indexer_k_quant_cache_kernel(
    const char* __restrict__ kv_cache,  // [num_blocks, block_size,
                                        // cache_stride]
    char* __restrict__ dst_k,           // [num_tokens, head_dim]
    char* __restrict__ dst_scale,  // [num_tokens, head_dim / quant_block_size *
                                   // 4]
    const int* __restrict__ block_table,  // [batch_size, num_blocks]
    const int* __restrict__ cu_seq_lens,  // [batch_size + 1]
    const int batch_size,                 // batch size
    const int64_t token_stride,           // stride for each token in dst_k
    const int64_t head_dim,               // dimension of each head
    const int64_t block_stride,           // stride for each block in kv_cache
    const int64_t cache_token_stride,     // stride for each token in kv_cache
    const int64_t cache_block_size,  // num_tokens for each block in kv_cache
    const int num_blocks,            // number of blocks
    const int num_tokens,            // number of tokens
    const int quant_block_size       // quantization block size
) {
  constexpr int VEC_SIZE = sizeof(float4) / sizeof(char);
  const int token_idx = blockIdx.x * blockDim.y + threadIdx.y;
  const int head_idx = (blockIdx.y * blockDim.x + threadIdx.x) * VEC_SIZE;
  // Find batch index within a block
  __shared__ int batch_idx[BLOCK_Y_SIZE];
  if (threadIdx.x == 0) {
    batch_idx[threadIdx.y] = -1;
  }
  __syncthreads();

  for (int iter = 0; iter < cuda_utils::ceil_div(batch_size, int(blockDim.x));
       iter++) {
    int tid = iter * blockDim.x + threadIdx.x;
    if (tid < batch_size) {
      const int seq_start = cu_seq_lens[tid];
      const int seq_end = cu_seq_lens[tid + 1];
      if (token_idx >= seq_start && token_idx < seq_end) {
        batch_idx[threadIdx.y] = tid;
      }
    }
  }

  __syncthreads();

  // num_tokens may be an allocation upper bound when Python avoids a D2H sync.
  // Only tokens covered by the exact device-side cu_seq_lens are valid to
  // gather.
  const int batch = batch_idx[threadIdx.y];
  if (head_idx >= head_dim || token_idx >= num_tokens || batch < 0) {
    return;
  }
  const int inbatch_seq_idx = token_idx - cu_seq_lens[batch];
  const int block_idx =
      block_table[batch * num_blocks + inbatch_seq_idx / cache_block_size];
  const int64_t src_block_offset = block_idx * block_stride;
  const int64_t cache_inblock_offset =
      (inbatch_seq_idx % cache_block_size) * head_dim + head_idx;
  const int64_t src_inblock_offset = src_block_offset + cache_inblock_offset;
  const int64_t dst_inblock_offset = token_idx * token_stride + head_idx;

  reinterpret_cast<float4*>(dst_k)[dst_inblock_offset / VEC_SIZE] =
      reinterpret_cast<const float4*>(kv_cache)[src_inblock_offset / VEC_SIZE];
  ;
  if (threadIdx.x == 0) {
    const int64_t src_scale_offset =
        src_block_offset + cache_block_size * head_dim +
        cache_inblock_offset * 4 / quant_block_size;
    reinterpret_cast<float*>(dst_scale)[dst_inblock_offset / quant_block_size] =
        reinterpret_cast<const float*>(kv_cache)[src_scale_offset / 4];
  }
}

}  // namespace vllm

// KV_T is the data type of key and value tensors.
// CACHE_T is the stored data type of kv-cache.
// KV_DTYPE is the real data type of kv-cache.
#define CALL_RESHAPE_AND_CACHE(KV_T, CACHE_T, KV_DTYPE)                     \
  vllm::reshape_and_cache_kernel<KV_T, CACHE_T, KV_DTYPE>                   \
      <<<grid, block, 0, stream>>>(                                         \
          reinterpret_cast<KV_T*>(key.data_ptr()),                          \
          reinterpret_cast<KV_T*>(value.data_ptr()),                        \
          reinterpret_cast<CACHE_T*>(key_cache.data_ptr()),                 \
          reinterpret_cast<CACHE_T*>(value_cache.data_ptr()),               \
          slot_mapping.const_data_ptr<int64_t>(), key_stride, value_stride, \
          num_heads, head_size, block_size, x,                              \
          reinterpret_cast<const float*>(k_scale.data_ptr()),               \
          reinterpret_cast<const float*>(v_scale.data_ptr()));

void reshape_and_cache(
    torch::stable::Tensor& key,    // [num_tokens, num_heads, head_size]
    torch::stable::Tensor& value,  // [num_tokens, num_heads, head_size]
    torch::stable::Tensor&
        key_cache,  // [num_blocks, num_heads, head_size/x, block_size, x]
    torch::stable::Tensor&
        value_cache,  // [num_blocks, num_heads, head_size, block_size]
    torch::stable::Tensor& slot_mapping,  // [num_tokens]
    const std::string& kv_cache_dtype, torch::stable::Tensor& k_scale,
    torch::stable::Tensor& v_scale) {
  int num_tokens = slot_mapping.size(0);
  int num_heads = key.size(1);
  int head_size = key.size(2);
  int block_size = key_cache.size(3);
  int x = key_cache.size(4);

  int key_stride = key.stride(0);
  int value_stride = value.stride(0);
  int head_div_x = head_size / x;

  dim3 grid(num_tokens);
  dim3 block(std::min(num_heads * head_div_x, 512));
  const torch::stable::accelerator::DeviceGuard device_guard(
      key.get_device_index());
  const cudaStream_t stream = get_current_cuda_stream();

  DISPATCH_BY_KV_CACHE_DTYPE(key.scalar_type(), kv_cache_dtype,
                             CALL_RESHAPE_AND_CACHE);
}

// KV_T is the data type of key and value tensors.
// CACHE_T is the stored data type of kv-cache.
// KV_DTYPE is the real data type of kv-cache.
#define CALL_RESHAPE_AND_CACHE_FLASH(KV_T, CACHE_T, KV_DTYPE)                \
  vllm::reshape_and_cache_flash_kernel<KV_T, CACHE_T, KV_DTYPE>              \
      <<<grid, block, 0, stream>>>(                                          \
          reinterpret_cast<KV_T*>(key.data_ptr()),                           \
          reinterpret_cast<KV_T*>(value.data_ptr()),                         \
          reinterpret_cast<CACHE_T*>(key_cache.data_ptr()),                  \
          reinterpret_cast<CACHE_T*>(value_cache.data_ptr()),                \
          slot_mapping.const_data_ptr<int64_t>(), block_stride, page_stride, \
          head_stride, key_stride, value_stride, num_heads, head_size,       \
          block_size, reinterpret_cast<const float*>(k_scale.data_ptr()),    \
          reinterpret_cast<const float*>(v_scale.data_ptr()),                \
          kv_scale_stride);

void reshape_and_cache_flash(
    torch::stable::Tensor& key,    // [num_tokens, num_heads, head_size]
    torch::stable::Tensor& value,  // [num_tokens, num_heads, head_size]
    torch::stable::Tensor&
        key_cache,  // [num_blocks, block_size, num_heads, head_size]
    torch::stable::Tensor&
        value_cache,  // [num_blocks, block_size, num_heads, head_size]
    torch::stable::Tensor& slot_mapping,  // [num_tokens] or [num_actual_tokens]
    const std::string& kv_cache_dtype,
    torch::stable::Tensor& k_scale,    // [1] or [num_heads]
    torch::stable::Tensor& v_scale) {  // [1] or [num_heads]
  // NOTE(woosuk): In vLLM V1, key.size(0) can be different from
  // slot_mapping.size(0) because of padding for CUDA graphs.
  // In vLLM V0, key.size(0) is always equal to slot_mapping.size(0) because
  // both include padding.
  // In vLLM V1, however, key.size(0) can be larger than slot_mapping.size(0)
  // since key includes padding for CUDA graphs, while slot_mapping does not.
  // In this case, slot_mapping.size(0) represents the actual number of tokens
  // before padding.
  // For compatibility with both cases, we use slot_mapping.size(0) as the
  // number of tokens.
  int num_tokens = slot_mapping.size(0);
  int num_heads = key.size(1);
  int head_size = key.size(2);

  const torch::stable::accelerator::DeviceGuard device_guard(
      key.get_device_index());
  const cudaStream_t stream = get_current_cuda_stream();

  if (kv_cache_dtype == "nvfp4" || kv_cache_dtype == "nvfp4_4over6") {
#if defined(ENABLE_NVFP4_SM100) || defined(ENABLE_NVFP4_SM120)
    // NVFP4 dispatch is compiled separately for SM100+.
    extern void reshape_and_cache_nvfp4_dispatch(
        torch::stable::Tensor & key, torch::stable::Tensor & value,
        torch::stable::Tensor & key_cache, torch::stable::Tensor & value_cache,
        torch::stable::Tensor & slot_mapping, torch::stable::Tensor & k_scale,
        torch::stable::Tensor & v_scale, const std::string& kv_cache_dtype);
    reshape_and_cache_nvfp4_dispatch(key, value, key_cache, value_cache,
                                     slot_mapping, k_scale, v_scale,
                                     kv_cache_dtype);
    return;
#else
    STD_TORCH_CHECK(
        false,
        "NVFP4 KV cache requires SM100+ (Blackwell). "
        "Please rebuild vllm with a Blackwell-compatible CUDA target.");
#endif
  }

  // Original FP8/auto path.
  int block_size = key_cache.size(1);

  int64_t key_stride = key.stride(0);
  int64_t value_stride = value.stride(0);
  int64_t block_stride = key_cache.stride(0);
  int64_t page_stride = key_cache.stride(1);
  int64_t head_stride = key_cache.stride(2);
  STD_TORCH_CHECK(key_cache.stride(0) == value_cache.stride(0));

  STD_TORCH_CHECK(k_scale.sizes().equals(v_scale.sizes()),
                  "k_scale and v_scale must have the same shape");
  STD_TORCH_CHECK(k_scale.numel() == 1 || k_scale.numel() == num_heads,
                  "k_scale and v_scale must be of shape [1] or [num_heads]");
  int kv_scale_stride = (k_scale.numel() > 1) ? 1 : 0;

  dim3 grid(num_tokens);
  dim3 block(std::min(num_heads * head_size, 512));

  DISPATCH_BY_KV_CACHE_DTYPE(key.scalar_type(), kv_cache_dtype,
                             CALL_RESHAPE_AND_CACHE_FLASH);
}

// KV_T is the data type of key and value tensors.
// CACHE_T is the stored data type of kv-cache.
// KV_DTYPE is the real data type of kv-cache.
#define CALL_CONCAT_AND_CACHE_MLA(KV_T, CACHE_T, KV_DTYPE)                    \
  vllm::concat_and_cache_mla_kernel<KV_T, CACHE_T, KV_DTYPE>                  \
      <<<grid, block, 0, stream>>>(                                           \
          reinterpret_cast<KV_T*>(kv_c.data_ptr()),                           \
          reinterpret_cast<KV_T*>(k_pe.data_ptr()),                           \
          reinterpret_cast<CACHE_T*>(kv_cache.data_ptr()),                    \
          slot_mapping.const_data_ptr<int64_t>(), block_stride, entry_stride, \
          kv_c_stride, k_pe_stride, kv_lora_rank, pe_dim, block_size,         \
          reinterpret_cast<const float*>(scale.data_ptr()));

// KV_T is the data type of key and value tensors.
// CACHE_T is the stored data type of kv-cache.
#define CALL_CONCAT_AND_CACHE_DS_MLA(KV_T, CACHE_T, KV_DTYPE)                 \
  vllm::concat_and_cache_ds_mla_kernel<KV_T, CACHE_T, KV_DTYPE>               \
      <<<grid, block, 0, stream>>>(                                           \
          reinterpret_cast<KV_T*>(kv_c.data_ptr()),                           \
          reinterpret_cast<KV_T*>(k_pe.data_ptr()),                           \
          reinterpret_cast<CACHE_T*>(kv_cache.data_ptr()),                    \
          slot_mapping.const_data_ptr<int64_t>(), block_stride, entry_stride, \
          kv_c_stride, k_pe_stride, kv_lora_rank, pe_dim, block_size,         \
          reinterpret_cast<const float*>(scale.data_ptr()));

void concat_and_cache_mla(
    torch::stable::Tensor& kv_c,      // [num_tokens, kv_lora_rank]
    torch::stable::Tensor& k_pe,      // [num_tokens, pe_dim]
    torch::stable::Tensor& kv_cache,  // [num_blocks, block_size, (kv_lora_rank
                                      // + pe_dim)]
    torch::stable::Tensor& slot_mapping,  // [num_tokens] or [num_actual_tokens]
    const std::string& kv_cache_dtype, torch::stable::Tensor& scale) {
  // NOTE(woosuk): In vLLM V1, key.size(0) can be different from
  // slot_mapping.size(0) because of padding for CUDA graphs.
  // In vLLM V0, key.size(0) is always equal to slot_mapping.size(0) because
  // both include padding.
  // In vLLM V1, however, key.size(0) can be larger than slot_mapping.size(0)
  // since key includes padding for CUDA graphs, while slot_mapping does not.
  // In this case, slot_mapping.size(0) represents the actual number of tokens
  // before padding.
  // For compatibility with both cases, we use slot_mapping.size(0) as the
  // number of tokens.
  int num_tokens = slot_mapping.size(0);
  int kv_lora_rank = kv_c.size(1);
  int pe_dim = k_pe.size(1);
  int block_size = kv_cache.size(1);

  const bool is_nvfp4_ds_mla = kv_cache_dtype == "nvfp4_ds_mla";
  if (kv_cache_dtype == "fp8_ds_mla") {
    STD_TORCH_CHECK(kv_lora_rank == 512,
                    "kv_lora_rank must be 512 for fp8_ds_mla");
    STD_TORCH_CHECK(pe_dim == 64 || pe_dim == 0,
                    "pe_dim must be 64 or 0 for fp8_ds_mla");
    STD_TORCH_CHECK(kv_cache.size(2) == 656 / kv_cache.element_size(),
                    "kv_cache.size(2) must be 656 bytes for fp8_ds_mla");
    STD_TORCH_CHECK(kv_c.element_size() == 2,
                    "kv_c.element_size() must be 2 for fp8_ds_mla");
    STD_TORCH_CHECK(k_pe.element_size() == 2,
                    "k_pe.element_size() must be 2 for fp8_ds_mla");
  } else if (is_nvfp4_ds_mla) {
    constexpr int bytes_per_token = 352;
    STD_TORCH_CHECK(kv_lora_rank == 512, "kv_lora_rank must be 512 for ",
                    kv_cache_dtype);
    STD_TORCH_CHECK(pe_dim == 64, "pe_dim must be 64 for ", kv_cache_dtype);
    STD_TORCH_CHECK(
        kv_cache.size(2) == bytes_per_token / kv_cache.element_size(),
        "kv_cache.size(2) must be ", bytes_per_token, " bytes for ",
        kv_cache_dtype);
    STD_TORCH_CHECK(
        kv_c.scalar_type() == torch::headeronly::ScalarType::BFloat16,
        "kv_c must be bfloat16 for ", kv_cache_dtype);
    STD_TORCH_CHECK(
        k_pe.scalar_type() == torch::headeronly::ScalarType::BFloat16,
        "k_pe must be bfloat16 for ", kv_cache_dtype);
  } else {
    STD_TORCH_CHECK(kv_cache.size(2) == kv_lora_rank + pe_dim);
  }

  int kv_c_stride = kv_c.stride(0);
  int k_pe_stride = k_pe.stride(0);
  int block_stride = kv_cache.stride(0);
  int entry_stride = kv_cache.stride(1);

  const torch::stable::accelerator::DeviceGuard device_guard(
      kv_c.get_device_index());
  const cudaStream_t stream = get_current_cuda_stream();

  if (kv_cache_dtype == "fp8_ds_mla") {
    dim3 grid(num_tokens);
    // For the NoPE part, each tile of 128 elements is handled by half of one
    // warp (16 threads). There are 4 total tiles, so 2 warps (64 threads).
    // Lanes 0 and 16 of each warp write the scale values for that warp's tiles.
    // The RoPE part (last 64 elements) is handled by another 1 warp (32
    // threads). So in total, we use 3 warps (96 threads) per block.
    dim3 block(96);
    DISPATCH_BY_KV_CACHE_DTYPE(kv_c.scalar_type(), kv_cache_dtype,
                               CALL_CONCAT_AND_CACHE_DS_MLA);
  } else if (is_nvfp4_ds_mla) {
    // The kernels live in nvfp4_ds_mla_cache_kernels.cu, which is only built
    // when the target list contains an SM100 arch.
#if !defined(USE_ROCM) && defined(ENABLE_NVFP4_SM100) && ENABLE_NVFP4_SM100
    const cudaDeviceProp* props = get_device_prop();
    STD_TORCH_CHECK(props->major == 10,
                    "nvfp4_ds_mla KV cache requires SM100 (Blackwell); "
                    "got SM",
                    props->major, props->minor);
    vllm::launch_concat_and_cache_nvfp4_ds_mla(
        kv_c.data_ptr(), k_pe.data_ptr(), kv_cache.data_ptr(),
        slot_mapping.const_data_ptr<int64_t>(), block_stride, entry_stride,
        kv_c_stride, k_pe_stride, block_size, num_tokens, stream);
#else
    STD_TORCH_CHECK(false,
                    "the nvfp4 ds_mla kv-cache format requires a CUDA build "
                    "targeting SM100 (Blackwell)");
#endif
  } else {
    dim3 grid(num_tokens);
    dim3 block(std::min(kv_lora_rank, 512));
    DISPATCH_BY_KV_CACHE_DTYPE(kv_c.scalar_type(), kv_cache_dtype,
                               CALL_CONCAT_AND_CACHE_MLA);
  }
}

void concat_and_cache_mla_grouped(
    torch::stable::Tensor& kv_c,  // [num_layers, num_tokens, kv_lora_rank]
    torch::stable::Tensor& k_pe,  // [num_layers, num_tokens, pe_dim]
    torch::stable::Tensor& kv_cache_ptrs,  // [num_layers] int64, on device
    torch::stable::Tensor& slot_mapping,   // [num_layers, num_tokens] int64
    int64_t block_size, int64_t block_stride, int64_t entry_stride,
    std::optional<torch::stable::Tensor> kv_scales,  // [num_layers] or None
    const std::string& kv_cache_dtype) {
  const bool use_fp8 = kv_cache_dtype == "fp8" ||
                       kv_cache_dtype == "fp8_e4m3" ||
                       kv_cache_dtype == "fp8_e5m2";
#ifdef USE_ROCM
  STD_TORCH_CHECK(kv_cache_dtype != "fp8_e5m2",
                  "concat_and_cache_mla_grouped does not support fp8_e5m2 "
                  "KV cache on ROCm");
#endif
  STD_TORCH_CHECK(
      use_fp8 || kv_cache_dtype == "auto" || kv_cache_dtype == "bfloat16",
      "concat_and_cache_mla_grouped only supports BF16 and plain "
      "FP8 KV cache; got ",
      kv_cache_dtype);

  STD_TORCH_CHECK(
      kv_c.scalar_type() == torch::headeronly::ScalarType::BFloat16 &&
          k_pe.scalar_type() == torch::headeronly::ScalarType::BFloat16,
      "concat_and_cache_mla_grouped requires BF16 inputs; got kv_c=",
      kv_c.scalar_type(), ", k_pe=", k_pe.scalar_type());
  STD_TORCH_CHECK(
      kv_cache_ptrs.scalar_type() == torch::headeronly::ScalarType::Long &&
          slot_mapping.scalar_type() == torch::headeronly::ScalarType::Long,
      "cache pointers and slot mapping must be int64");

  const int num_layers = kv_c.size(0);
  const int num_tokens = kv_c.size(1);
  const int kv_lora_rank = kv_c.size(2);
  const int pe_dim = k_pe.size(2);
  const float* kv_scales_ptr = nullptr;
  if (use_fp8) {
    STD_TORCH_CHECK(kv_scales.has_value(),
                    "FP8 grouped cache insert requires kv_scales");
    STD_TORCH_CHECK(
        kv_scales->scalar_type() == torch::headeronly::ScalarType::Float,
        "kv_scales must be float32");
    STD_TORCH_CHECK(kv_scales->numel() == num_layers,
                    "kv_scales must contain one scale per layer");
    STD_TORCH_CHECK(kv_scales->is_cuda() && kv_scales->get_device_index() ==
                                                kv_c.get_device_index(),
                    "kv_scales must be on the same CUDA device as kv_c");
    STD_TORCH_CHECK(kv_scales->is_contiguous(), "kv_scales must be contiguous");
    kv_scales_ptr = kv_scales->const_data_ptr<float>();
  } else {
    STD_TORCH_CHECK(!kv_scales.has_value(),
                    "BF16 grouped cache insert does not use kv_scales");
  }

  if (num_tokens == 0 || num_layers == 0) {
    return;
  }

  const int64_t kv_c_layer_stride = kv_c.stride(0);
  const int64_t kv_c_token_stride = kv_c.stride(1);
  const int64_t k_pe_layer_stride = k_pe.stride(0);
  const int64_t k_pe_token_stride = k_pe.stride(1);
  const int64_t slot_layer_stride = slot_mapping.stride(0);

  const torch::stable::accelerator::DeviceGuard device_guard(
      kv_c.get_device_index());
  const cudaStream_t stream = get_current_cuda_stream();

  const dim3 grid(num_tokens, num_layers);
  const dim3 block(std::min(kv_lora_rank, 512));

  if (!use_fp8) {
    vllm::concat_and_cache_mla_grouped_kernel<uint16_t, uint16_t,
                                              vllm::Fp8KVCacheDataType::kAuto>
        <<<grid, block, 0, stream>>>(
            reinterpret_cast<const uint16_t*>(kv_c.data_ptr()),
            reinterpret_cast<const uint16_t*>(k_pe.data_ptr()),
            kv_cache_ptrs.const_data_ptr<int64_t>(), nullptr,
            slot_mapping.const_data_ptr<int64_t>(), kv_c_layer_stride,
            kv_c_token_stride, k_pe_layer_stride, k_pe_token_stride,
            slot_layer_stride, block_stride, entry_stride, kv_lora_rank, pe_dim,
            block_size);
    return;
  }

#define LAUNCH_GROUPED_FP8(KV_DTYPE)                                           \
  vllm::concat_and_cache_mla_grouped_kernel<__nv_bfloat16, uint8_t, KV_DTYPE>  \
      <<<grid, block, 0, stream>>>(                                            \
          reinterpret_cast<const __nv_bfloat16*>(kv_c.data_ptr()),             \
          reinterpret_cast<const __nv_bfloat16*>(k_pe.data_ptr()),             \
          kv_cache_ptrs.const_data_ptr<int64_t>(), kv_scales_ptr,              \
          slot_mapping.const_data_ptr<int64_t>(), kv_c_layer_stride,           \
          kv_c_token_stride, k_pe_layer_stride, k_pe_token_stride,             \
          slot_layer_stride, block_stride, entry_stride, kv_lora_rank, pe_dim, \
          block_size)

  if (kv_cache_dtype == "fp8_e5m2") {
    LAUNCH_GROUPED_FP8(vllm::Fp8KVCacheDataType::kFp8E5M2);
  } else {
    LAUNCH_GROUPED_FP8(vllm::Fp8KVCacheDataType::kFp8E4M3);
  }
#undef LAUNCH_GROUPED_FP8
}

namespace vllm {

template <typename Tout, typename Tin, Fp8KVCacheDataType kv_dt>
__global__ void convert_fp8_kernel(const Tin* __restrict__ src_cache,
                                   Tout* __restrict__ dst_cache,
                                   const float scale,
                                   const int64_t block_stride) {
  const int64_t block_idx = blockIdx.x;
  for (int i = threadIdx.x; i < block_stride; i += blockDim.x) {
    int64_t idx = block_idx * block_stride + i;
    dst_cache[idx] =
        fp8::scaled_convert<Tout, Tin, kv_dt>(src_cache[idx], scale);
  }
}

}  // namespace vllm

#define CALL_CONVERT_FP8(Tout, Tin, KV_DTYPE)                                \
  vllm::convert_fp8_kernel<Tout, Tin, KV_DTYPE><<<grid, block, 0, stream>>>( \
      reinterpret_cast<Tin*>(src_cache.data_ptr()),                          \
      reinterpret_cast<Tout*>(dst_cache.data_ptr()), scale, block_stride);

// Only for testing.
void convert_fp8(torch::stable::Tensor& dst_cache,
                 torch::stable::Tensor& src_cache, const double scale,
                 const std::string& kv_cache_dtype) {
  torch::stable::Device src_device = src_cache.device();
  torch::stable::Device dst_device = dst_cache.device();
  STD_TORCH_CHECK(src_device.is_cuda(), "src must be on a GPU")
  STD_TORCH_CHECK(dst_device.is_cuda(), "dst must be on a GPU")
  STD_TORCH_CHECK(src_device.index() == dst_device.index(),
                  "src and dst must be on the same GPU");
  torch::stable::accelerator::DeviceGuard device_guard(src_device.index());

  int64_t num_blocks = src_cache.size(0);
  int64_t block_stride = src_cache.stride(0);

  dim3 grid(num_blocks);
  dim3 block(std::min(block_stride, int64_t(512)));
  const cudaStream_t stream = get_current_cuda_stream();

  if (kv_cache_dtype == "auto") {
    if (src_cache.scalar_type() == torch::headeronly::ScalarType::Float) {
      CALL_CONVERT_FP8(uint8_t, float, vllm::Fp8KVCacheDataType::kAuto);
    } else if (src_cache.scalar_type() == torch::headeronly::ScalarType::Half) {
      CALL_CONVERT_FP8(uint8_t, uint16_t, vllm::Fp8KVCacheDataType::kAuto);
    } else if (src_cache.scalar_type() ==
               torch::headeronly::ScalarType::BFloat16) {
      CALL_CONVERT_FP8(uint8_t, __nv_bfloat16, vllm::Fp8KVCacheDataType::kAuto);
    } else if (dst_cache.scalar_type() ==
               torch::headeronly::ScalarType::Float) {
      CALL_CONVERT_FP8(float, uint8_t, vllm::Fp8KVCacheDataType::kAuto);
    } else if (dst_cache.scalar_type() == torch::headeronly::ScalarType::Half) {
      CALL_CONVERT_FP8(uint16_t, uint8_t, vllm::Fp8KVCacheDataType::kAuto);
    } else if (dst_cache.scalar_type() ==
               torch::headeronly::ScalarType::BFloat16) {
      CALL_CONVERT_FP8(__nv_bfloat16, uint8_t, vllm::Fp8KVCacheDataType::kAuto);
    }
  } else if (kv_cache_dtype == "fp8" || kv_cache_dtype == "fp8_e4m3") {
    if (src_cache.scalar_type() == torch::headeronly::ScalarType::Float) {
      CALL_CONVERT_FP8(uint8_t, float, vllm::Fp8KVCacheDataType::kFp8E4M3);
    } else if (src_cache.scalar_type() == torch::headeronly::ScalarType::Half) {
      CALL_CONVERT_FP8(uint8_t, uint16_t, vllm::Fp8KVCacheDataType::kFp8E4M3);
    } else if (src_cache.scalar_type() ==
               torch::headeronly::ScalarType::BFloat16) {
      CALL_CONVERT_FP8(uint8_t, __nv_bfloat16,
                       vllm::Fp8KVCacheDataType::kFp8E4M3);
    } else if (dst_cache.scalar_type() ==
               torch::headeronly::ScalarType::Float) {
      CALL_CONVERT_FP8(float, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3);
    } else if (dst_cache.scalar_type() == torch::headeronly::ScalarType::Half) {
      CALL_CONVERT_FP8(uint16_t, uint8_t, vllm::Fp8KVCacheDataType::kFp8E4M3);
    } else if (dst_cache.scalar_type() ==
               torch::headeronly::ScalarType::BFloat16) {
      CALL_CONVERT_FP8(__nv_bfloat16, uint8_t,
                       vllm::Fp8KVCacheDataType::kFp8E4M3);
    }
  } else {
    STD_TORCH_CHECK(false, "Unsupported data type: ", kv_cache_dtype);
  }
}

namespace vllm {

struct GatherPageTask {
  int32_t req_id;
  int32_t logical_block;
  int32_t page_token_begin;
  int32_t page_token_end;
  int32_t output_token_begin;
};

template <bool has_terminal_start>
__device__ __forceinline__ bool map_gather_page_task(
    int32_t task, const int32_t* __restrict__ output_starts, int32_t num_reqs,
    int32_t total_tokens, int32_t block_size,
    const int32_t* __restrict__ seq_starts, GatherPageTask& page) {
  int32_t relative_page = task;
  for (int32_t req_id = 0; req_id < num_reqs; ++req_id) {
    const int32_t output_begin = min(output_starts[req_id], total_tokens);
    int32_t output_end;
    if constexpr (has_terminal_start) {
      output_end = min(output_starts[req_id + 1], total_tokens);
    } else {
      output_end =
          min(req_id + 1 < num_reqs ? output_starts[req_id + 1] : total_tokens,
              total_tokens);
    }
    const int32_t seq_len = max(output_end - output_begin, 0);
    const int32_t source_begin = seq_starts == nullptr ? 0 : seq_starts[req_id];
    const int32_t first_block = source_begin / block_size;
    const int32_t num_pages =
        cuda_utils::ceil_div(source_begin + seq_len, block_size) - first_block;
    if (relative_page < num_pages) {
      page.req_id = req_id;
      page.logical_block = first_block + relative_page;
      const int32_t page_begin = page.logical_block * block_size;
      const int32_t copy_begin = max(source_begin, page_begin);
      const int32_t copy_end =
          min(source_begin + seq_len, page_begin + block_size);
      page.page_token_begin = copy_begin - page_begin;
      page.page_token_end = copy_end - page_begin;
      page.output_token_begin = output_begin + copy_begin - source_begin;
      return true;
    }
    relative_page -= num_pages;
  }
  return false;
}

template <typename scalar_t, typename cache_t, Fp8KVCacheDataType kv_dt,
          int ENTRY_SIZE>
__global__ void gather_and_maybe_dequant_cache_page(
    const cache_t* __restrict__ src_cache, scalar_t* __restrict__ dst,
    const int32_t* __restrict__ block_table,
    const int32_t* __restrict__ cu_seq_lens, const int32_t num_reqs,
    const int32_t num_tokens, const int32_t block_size,
    const int64_t block_table_stride, const int64_t cache_block_stride,
    const int64_t cache_entry_stride, const int64_t dst_entry_stride,
    const float* __restrict__ scale, const int32_t* __restrict__ seq_starts) {
  constexpr int32_t vec_size = sizeof(float4) / sizeof(scalar_t);
  constexpr int32_t vec_iter_cnt = ENTRY_SIZE / vec_size;
  static_assert(ENTRY_SIZE % vec_size == 0);
  using ltype = vllm::vec_n_t<cache_t, vec_size>;
  using stype = vllm::vec_n_t<scalar_t, vec_size>;

  __shared__ GatherPageTask page;
  __shared__ int32_t physical_block;
  __shared__ bool has_task;
  __shared__ bool copy_task;
  __shared__ float scale_value;

  if (threadIdx.x == 0) {
    if constexpr (kv_dt != Fp8KVCacheDataType::kAuto) {
      scale_value = *scale;
    }
  }
  __syncthreads();

  if (threadIdx.x == 0) {
    has_task =
        map_gather_page_task<true>(blockIdx.x, cu_seq_lens, num_reqs,
                                   num_tokens, block_size, seq_starts, page);
    copy_task = has_task && page.logical_block < block_table_stride;
    if (copy_task) {
      physical_block =
          block_table[page.req_id * block_table_stride + page.logical_block];
    }
  }
  __syncthreads();
  if (!has_task) {
    return;
  }

  if (copy_task) {
    const int32_t page_vectors =
        (page.page_token_end - page.page_token_begin) * vec_iter_cnt;
    for (int32_t flat_idx = threadIdx.x; flat_idx < page_vectors;
         flat_idx += blockDim.x) {
      const int32_t token_offset = flat_idx / vec_iter_cnt;
      const int32_t idx = flat_idx - token_offset * vec_iter_cnt;
      const int32_t page_token = page.page_token_begin + token_offset;
      const int32_t output_token = page.output_token_begin + token_offset;
      const cache_t* src = src_cache + physical_block * cache_block_stride +
                           page_token * cache_entry_stride;
      scalar_t* output = dst + output_token * dst_entry_stride;

      if constexpr (kv_dt == Fp8KVCacheDataType::kAuto) {
        reinterpret_cast<stype*>(output)[idx] =
            static_cast<stype>(reinterpret_cast<const ltype*>(src)[idx]);
      } else {
        const ltype loaded = reinterpret_cast<const ltype*>(src)[idx];
        stype converted;
#pragma unroll
        for (int32_t j = 0; j < vec_size; ++j) {
          converted.val[j] = fp8::scaled_convert<scalar_t, cache_t, kv_dt>(
              loaded.val[j], scale_value);
        }
        reinterpret_cast<stype*>(output)[idx] = converted;
      }
    }
  }
}

}  // namespace vllm

// Macro to dispatch the kernel based on the data type.
// SCALAR_T is the data type of the destination tensor.
// CACHE_T is the stored data type of kv-cache.
// KV_DTYPE is the real data type of kv-cache.
#define CALL_GATHER_CACHE(SCALAR_T, CACHE_T, KV_DTYPE, ENTRY_SZ)            \
  vllm::gather_and_maybe_dequant_cache_page<SCALAR_T, CACHE_T, KV_DTYPE,    \
                                            ENTRY_SZ>                       \
      <<<grid, block, 0, stream>>>(                                         \
          reinterpret_cast<CACHE_T*>(src_cache.data_ptr()),                 \
          reinterpret_cast<SCALAR_T*>(dst.data_ptr()),                      \
          block_table.const_data_ptr<int32_t>(),                            \
          cu_seq_lens.const_data_ptr<int32_t>(), num_reqs,                  \
          static_cast<int32_t>(num_tokens), block_size, block_table_stride, \
          cache_block_stride, cache_entry_stride, dst_entry_stride,         \
          reinterpret_cast<const float*>(scale.data_ptr()), seq_starts_ptr)

#define CALL_GATHER_CACHE_576(SCALAR_T, CACHE_T, KV_DTYPE) \
  CALL_GATHER_CACHE(SCALAR_T, CACHE_T, KV_DTYPE, 576)

#define CALL_GATHER_CACHE_512(SCALAR_T, CACHE_T, KV_DTYPE) \
  CALL_GATHER_CACHE(SCALAR_T, CACHE_T, KV_DTYPE, 512)

#define CALL_GATHER_CACHE_320(SCALAR_T, CACHE_T, KV_DTYPE) \
  CALL_GATHER_CACHE(SCALAR_T, CACHE_T, KV_DTYPE, 320)

// Gather sequences from the cache into the destination tensor.
//  - cu_seq_lens contains the cumulative sequence lengths for each batch
//  - block_table contains the cache block indices for each sequence
//  - Optionally, seq_starts (if provided) offsets the starting block index by
//  (seq_starts[bid] / page_size)
void gather_and_maybe_dequant_cache(
    torch::stable::Tensor const&
        src_cache,                     // [NUM_BLOCKS, BLOCK_SIZE, ENTRIES...]
    torch::stable::Tensor const& dst,  // [TOT_TOKENS, ENTRIES...]
    torch::stable::Tensor const& block_table,   // [BATCH, BLOCK_INDICES]
    torch::stable::Tensor const& cu_seq_lens,   // [BATCH+1]
    torch::stable::Tensor const& token_to_seq,  // [MAX_TOKEN_ACROSS_CHUNKS]
    int64_t num_tokens, const std::string& kv_cache_dtype,
    torch::stable::Tensor const& scale,
    std::optional<torch::stable::Tensor> seq_starts = std::nullopt) {
  torch::stable::accelerator::DeviceGuard device_guard(
      src_cache.get_device_index());
  const cudaStream_t stream = get_current_cuda_stream();
  (void)token_to_seq;

  int32_t block_size = src_cache.size(1);
  int32_t head_dim = dst.size(-1);

  STD_TORCH_CHECK(
      block_table.scalar_type() == torch::headeronly::ScalarType::Int,
      "block_table must be int32");
  STD_TORCH_CHECK(
      cu_seq_lens.scalar_type() == torch::headeronly::ScalarType::Int,
      "cu_seq_lens must be int32");
  if (seq_starts.has_value()) {
    STD_TORCH_CHECK(
        seq_starts.value().scalar_type() == torch::headeronly::ScalarType::Int,
        "seq_starts must be int32");
  }
  STD_TORCH_CHECK(head_dim == 320 || head_dim == 512 || head_dim == 576,
                  "gather_and_maybe_dequant_cache only support the head_dim to "
                  "320 or 512 or 576 for better performance")

  STD_TORCH_CHECK(src_cache.device() == dst.device(),
                  "src_cache and dst must be on the same device");
  STD_TORCH_CHECK(src_cache.device() == block_table.device(),
                  "src_cache and block_table must be on the same device");
  STD_TORCH_CHECK(src_cache.device() == cu_seq_lens.device(),
                  "src_cache and cu_seq_lens must be on the same device");
  if (seq_starts.has_value()) {
    STD_TORCH_CHECK(src_cache.device() == seq_starts.value().device(),
                    "src_cache and seq_starts must be on the same device");
  }

  if (num_tokens == 0) {
    return;
  }

  const int32_t num_reqs = cu_seq_lens.size(0) - 1;
  const int64_t block_table_stride = block_table.stride(0);
  const int64_t cache_block_stride = src_cache.stride(0);
  const int64_t cache_entry_stride = src_cache.stride(1);
  const int64_t dst_entry_stride = dst.stride(0);
  const int32_t page_threads = num_tokens >= (1 << 20) ? 128 : 256;
  const int32_t required_blocks =
      cuda_utils::ceil_div(static_cast<int32_t>(num_tokens), block_size) +
      2 * num_reqs;
  const dim3 grid(required_blocks);
  const dim3 block(page_threads);

  const int32_t* seq_starts_ptr =
      seq_starts.has_value() ? seq_starts.value().const_data_ptr<int32_t>()
                             : nullptr;

  if (head_dim == 576) {
    DISPATCH_BY_KV_CACHE_DTYPE(dst.scalar_type(), kv_cache_dtype,
                               CALL_GATHER_CACHE_576);
  } else if (head_dim == 512) {
    DISPATCH_BY_KV_CACHE_DTYPE(dst.scalar_type(), kv_cache_dtype,
                               CALL_GATHER_CACHE_512);
  } else {
    DISPATCH_BY_KV_CACHE_DTYPE(dst.scalar_type(), kv_cache_dtype,
                               CALL_GATHER_CACHE_320);
  }
}

namespace vllm {

__device__ __forceinline__ void gather_and_upconvert_fp8_token(
    const uint8_t* __restrict__ token_ptr, __nv_bfloat16* __restrict__ dst_ptr,
    int32_t lane_id) {
  const uint2* fp8_src = reinterpret_cast<const uint2*>(token_ptr);
  const float* scales = reinterpret_cast<const float*>(token_ptr + 512);
  int4* nope_dst = reinterpret_cast<int4*>(dst_ptr);

#pragma unroll
  for (int32_t phase = 0; phase < 2; ++phase) {
    const int32_t chunk = phase * 32 + lane_id;
    const uint2 fp8_data = fp8_src[chunk];
    const float scale = scales[chunk >> 4];
#ifdef USE_ROCM
    const bf16_8_t bf16_data =
        fp8::scaled_vec_conversion<bf16_8_t, uint2>(fp8_data, scale);
#else
    const bf16_8_t bf16_data =
        fp8::scaled_vec_conversion<bf16_8_t, uint2>(fp8_data, scale, __NV_E4M3);
#endif
    nope_dst[chunk] = *reinterpret_cast<const int4*>(&bf16_data);
  }

  const int* rope_src = reinterpret_cast<const int*>(token_ptr + 528);
  int* rope_dst = reinterpret_cast<int*>(dst_ptr + 512);
  rope_dst[lane_id] = rope_src[lane_id];
}

__global__ void cp_gather_and_upconvert_fp8_kv_cache_page(
    const uint8_t* __restrict__ src_cache, __nv_bfloat16* __restrict__ dst,
    const int32_t* __restrict__ block_table,
    const int32_t* __restrict__ workspace_starts, const int32_t num_reqs,
    const int32_t block_size, const int32_t total_tokens,
    const int64_t block_table_stride, const int64_t cache_block_stride,
    const int64_t cache_entry_stride, const int64_t dst_entry_stride,
    const int32_t* __restrict__ seq_starts,
    const uint8_t* __restrict__ host_cache,
    const int32_t* __restrict__ host_row_ids,
    const int32_t* __restrict__ device_row_ids,
    const int64_t host_entry_stride) {
  constexpr int32_t warps_per_cta = 16;
  __shared__ GatherPageTask page;
  __shared__ int32_t physical_block;
  __shared__ bool has_task;

  if (threadIdx.x == 0) {
    has_task =
        map_gather_page_task<false>(blockIdx.x, workspace_starts, num_reqs,
                                    total_tokens, block_size, seq_starts, page);
    if (has_task) {
      physical_block =
          block_table[page.req_id * block_table_stride + page.logical_block];
    }
  }
  __syncthreads();
  if (!has_task) {
    return;
  }

  const int32_t warp_id = threadIdx.x >> 5;
  const int32_t lane_id = threadIdx.x & 31;
  for (int32_t page_token = page.page_token_begin + warp_id;
       page_token < page.page_token_end; page_token += warps_per_cta) {
    const int32_t output_token =
        page.output_token_begin + page_token - page.page_token_begin;
    const int32_t source_row = physical_block * block_size + page_token;
    const uint8_t* token_ptr;
    if (device_row_ids == nullptr) {
      token_ptr = src_cache + physical_block * cache_block_stride +
                  page_token * cache_entry_stride;
    } else {
      const int32_t device_row = device_row_ids[source_row];
      if (device_row >= 0) {
        token_ptr = src_cache + (device_row / block_size) * cache_block_stride +
                    (device_row % block_size) * cache_entry_stride;
      } else {
        token_ptr =
            host_cache +
            static_cast<int64_t>(host_row_ids[source_row]) * host_entry_stride;
      }
    }
    __nv_bfloat16* dst_ptr = dst + output_token * dst_entry_stride;
    gather_and_upconvert_fp8_token(token_ptr, dst_ptr, lane_id);
  }
}

template <bool contiguous_entries, bool vectorized>
__global__ void cp_gather_cache_page(
    const uint8_t* __restrict__ src_cache,    // [NUM_BLOCKS, BLOCK_SIZE,
                                              // ENTRY_SIZE_BYTES]
    uint8_t* __restrict__ dst,                // [TOT_TOKENS, ENTRY_SIZE_BYTES]
    const int32_t* __restrict__ block_table,  // [BATCH, BLOCK_INDICES]
    const int32_t* __restrict__ cu_seq_lens,  // [BATCH+1]
    const int32_t num_reqs, const int32_t block_size,
    const int32_t entry_size_bytes, const int32_t total_tokens,
    const int64_t block_table_stride, const int64_t cache_block_stride,
    const int64_t cache_entry_stride, const int64_t dst_entry_stride,
    const int32_t* __restrict__ seq_starts) {
  constexpr int32_t threads_per_token = 32;
  constexpr int32_t tokens_per_cta = 256 / threads_per_token;
  __shared__ GatherPageTask page;
  __shared__ int32_t physical_block;
  __shared__ bool has_task;

  if (threadIdx.x == 0) {
    has_task =
        map_gather_page_task<true>(blockIdx.x, cu_seq_lens, num_reqs,
                                   total_tokens, block_size, seq_starts, page);
    if (has_task) {
      physical_block =
          block_table[page.req_id * block_table_stride + page.logical_block];
    }
  }
  __syncthreads();
  if (!has_task) {
    return;
  }

  if constexpr (contiguous_entries) {
    const int64_t src_offset = physical_block * cache_block_stride +
                               page.page_token_begin * entry_size_bytes;
    const int64_t dst_offset = page.output_token_begin * entry_size_bytes;
    const int32_t copy_bytes =
        (page.page_token_end - page.page_token_begin) * entry_size_bytes;
    if constexpr (vectorized) {
      const int4* src = reinterpret_cast<const int4*>(src_cache + src_offset);
      int4* output = reinterpret_cast<int4*>(dst + dst_offset);
      const int32_t num_vectors = copy_bytes / sizeof(int4);
      int32_t idx = threadIdx.x;
      for (; idx + blockDim.x < num_vectors; idx += 2 * blockDim.x) {
        const int4 first = src[idx];
        const int4 second = src[idx + blockDim.x];
        output[idx] = first;
        output[idx + blockDim.x] = second;
      }
      for (; idx < num_vectors; idx += blockDim.x) {
        output[idx] = src[idx];
      }
    } else {
      const uint8_t* src = src_cache + src_offset;
      uint8_t* output = dst + dst_offset;
      for (int32_t idx = threadIdx.x; idx < copy_bytes; idx += blockDim.x) {
        output[idx] = src[idx];
      }
    }
  } else {
    const int32_t token_group = threadIdx.x / threads_per_token;
    const int32_t lane_id = threadIdx.x % threads_per_token;
    for (int32_t page_token = page.page_token_begin + token_group;
         page_token < page.page_token_end; page_token += tokens_per_cta) {
      const int32_t output_token =
          page.output_token_begin + page_token - page.page_token_begin;
      const uint8_t* src = src_cache + physical_block * cache_block_stride +
                           page_token * cache_entry_stride;
      uint8_t* output = dst + output_token * dst_entry_stride;
      if constexpr (vectorized) {
        const int4* src_vec = reinterpret_cast<const int4*>(src);
        int4* dst_vec = reinterpret_cast<int4*>(output);
        const int32_t num_vectors = entry_size_bytes / sizeof(int4);
        for (int32_t idx = lane_id; idx < num_vectors;
             idx += threads_per_token) {
          dst_vec[idx] = src_vec[idx];
        }
      } else {
        for (int32_t idx = lane_id; idx < entry_size_bytes;
             idx += threads_per_token) {
          output[idx] = src[idx];
        }
      }
    }
  }
}
}  // namespace vllm

// Gather sequences from the cache into the destination tensor.
//  - cu_seq_lens contains the cumulative sequence lengths for each batch
//  - block_table contains the cache block indices for each sequence
//  - Optionally, seq_starts (if provided) offsets the starting slot index by
//  seq_starts[bid]
void cp_gather_cache(
    torch::stable::Tensor const&
        src_cache,                     // [NUM_BLOCKS, BLOCK_SIZE, ENTRIES...]
    torch::stable::Tensor const& dst,  // [TOT_TOKENS, ENTRIES...]
    torch::stable::Tensor const& block_table,  // [BATCH, BLOCK_INDICES]
    torch::stable::Tensor const& cu_seq_lens,  // [BATCH+1]
    int64_t batch_size,
    std::optional<torch::stable::Tensor> seq_starts = std::nullopt) {
  torch::stable::accelerator::DeviceGuard device_guard(
      src_cache.get_device_index());
  const cudaStream_t stream = get_current_cuda_stream();

  int32_t block_size = src_cache.size(1);
  int32_t entry_size = torch::stable::flatten(src_cache, 2, -1).size(2);

  STD_TORCH_CHECK(
      block_table.scalar_type() == torch::headeronly::ScalarType::Int,
      "block_table must be int32");
  STD_TORCH_CHECK(
      cu_seq_lens.scalar_type() == torch::headeronly::ScalarType::Int,
      "cu_seq_lens must be int32");
  if (seq_starts.has_value()) {
    STD_TORCH_CHECK(
        seq_starts.value().scalar_type() == torch::headeronly::ScalarType::Int,
        "seq_starts must be int32");
  }

  STD_TORCH_CHECK(src_cache.device() == dst.device(),
                  "src_cache and dst must be on the same device");
  STD_TORCH_CHECK(src_cache.device() == block_table.device(),
                  "src_cache and block_table must be on the same device");
  STD_TORCH_CHECK(src_cache.device() == cu_seq_lens.device(),
                  "src_cache and cu_seq_lens must be on the same device");
  if (seq_starts.has_value()) {
    STD_TORCH_CHECK(src_cache.device() == seq_starts.value().device(),
                    "src_cache and seq_starts must be on the same device");
  }

  STD_TORCH_CHECK(src_cache.scalar_type() == dst.scalar_type(),
                  "src_cache and dst must have the same dtype");
  const int32_t element_size = src_cache.element_size();
  STD_TORCH_CHECK(element_size == 1 || element_size == 2 || element_size == 4,
                  "Unsupported data type width: ", element_size * 8);
  STD_TORCH_CHECK(batch_size == cu_seq_lens.size(0) - 1,
                  "batch_size must match cu_seq_lens");

  const int32_t total_tokens = dst.size(0);
  if (total_tokens == 0) {
    return;
  }

  const int32_t entry_size_bytes = entry_size * element_size;
  const int64_t block_table_stride = block_table.stride(0);
  const int64_t cache_block_stride = src_cache.stride(0) * element_size;
  const int64_t cache_entry_stride = src_cache.stride(1) * element_size;
  const int64_t dst_entry_stride = dst.stride(0) * element_size;
  const int32_t* seq_starts_ptr =
      seq_starts.has_value() ? seq_starts.value().const_data_ptr<int32_t>()
                             : nullptr;

  constexpr std::uintptr_t vector_alignment = alignof(int4);
  const bool vectorized =
      reinterpret_cast<std::uintptr_t>(src_cache.data_ptr()) %
              vector_alignment ==
          0 &&
      reinterpret_cast<std::uintptr_t>(dst.data_ptr()) % vector_alignment ==
          0 &&
      entry_size_bytes % vector_alignment == 0 &&
      cache_block_stride % vector_alignment == 0 &&
      cache_entry_stride % vector_alignment == 0 &&
      dst_entry_stride % vector_alignment == 0;
  const bool contiguous_entries = cache_entry_stride == entry_size_bytes &&
                                  dst_entry_stride == entry_size_bytes;

  const int32_t num_reqs = static_cast<int32_t>(batch_size);
  const int32_t required_blocks =
      cuda_utils::ceil_div(total_tokens, block_size) + 2 * num_reqs;
  const dim3 grid(required_blocks);
  const dim3 block(256);

#define CALL_CP_GATHER_CACHE(CONTIGUOUS, VECTORIZED)                   \
  vllm::cp_gather_cache_page<CONTIGUOUS, VECTORIZED>                   \
      <<<grid, block, 0, stream>>>(                                    \
          reinterpret_cast<const uint8_t*>(src_cache.data_ptr()),      \
          reinterpret_cast<uint8_t*>(dst.data_ptr()),                  \
          block_table.const_data_ptr<int32_t>(),                       \
          cu_seq_lens.const_data_ptr<int32_t>(), num_reqs, block_size, \
          entry_size_bytes, total_tokens, block_table_stride,          \
          cache_block_stride, cache_entry_stride, dst_entry_stride,    \
          seq_starts_ptr)

  if (contiguous_entries && vectorized) {
    CALL_CP_GATHER_CACHE(true, true);
  } else if (contiguous_entries) {
    CALL_CP_GATHER_CACHE(true, false);
  } else if (vectorized) {
    CALL_CP_GATHER_CACHE(false, true);
  } else {
    CALL_CP_GATHER_CACHE(false, false);
  }

#undef CALL_CP_GATHER_CACHE
}

void cp_gather_and_upconvert_fp8_kv_cache(
    torch::stable::Tensor const& src_cache,    // [NUM_BLOCKS, BLOCK_SIZE, 656]
    torch::stable::Tensor const& dst,          // [TOT_TOKENS, 576]
    torch::stable::Tensor const& block_table,  // [BATCH, BLOCK_INDICES]
    torch::stable::Tensor const& workspace_starts,  // [BATCH]
    int64_t batch_size,
    std::optional<torch::stable::Tensor> seq_starts = std::nullopt,
    std::optional<torch::stable::Tensor> host_cache = std::nullopt,
    std::optional<torch::stable::Tensor> host_row_ids = std::nullopt,
    std::optional<torch::stable::Tensor> device_row_ids = std::nullopt) {
  torch::stable::accelerator::DeviceGuard device_guard(
      src_cache.get_device_index());
  const cudaStream_t stream = get_current_cuda_stream();

  int32_t block_size = src_cache.size(1);
  int32_t head_dim = dst.size(1);

  STD_TORCH_CHECK(
      block_table.scalar_type() == torch::headeronly::ScalarType::Int,
      "block_table must be int32");
  STD_TORCH_CHECK(
      workspace_starts.scalar_type() == torch::headeronly::ScalarType::Int,
      "workspace_starts must be int32");
  if (seq_starts.has_value()) {
    STD_TORCH_CHECK(
        seq_starts.value().scalar_type() == torch::headeronly::ScalarType::Int,
        "seq_starts must be int32");
  }

  const bool has_host_rows = host_cache.has_value();
  STD_TORCH_CHECK(
      has_host_rows == host_row_ids.has_value() &&
          has_host_rows == device_row_ids.has_value(),
      "host_cache, host_row_ids, and device_row_ids must be provided together");

  STD_TORCH_CHECK(src_cache.device() == dst.device(),
                  "src_cache and dst must be on the same device");
  STD_TORCH_CHECK(src_cache.device() == block_table.device(),
                  "src_cache and block_table must be on the same device");
  STD_TORCH_CHECK(src_cache.device() == workspace_starts.device(),
                  "src_cache and workspace_starts must be on the same device");
  if (seq_starts.has_value()) {
    STD_TORCH_CHECK(src_cache.device() == seq_starts.value().device(),
                    "src_cache and seq_starts must be on the same device");
  }
  const uint8_t* host_cache_ptr = nullptr;
  const int32_t* host_row_ids_ptr = nullptr;
  const int32_t* device_row_ids_ptr = nullptr;
  int64_t host_entry_stride = 0;
  if (has_host_rows) {
    auto const& host = host_cache.value();
    auto const& host_rows = host_row_ids.value();
    auto const& device_rows = device_row_ids.value();
    STD_TORCH_CHECK(
        host.device().is_cpu() &&
            host.scalar_type() == torch::headeronly::ScalarType::Byte &&
            host.dim() == 2 && host.is_contiguous() &&
            host.size(1) == src_cache.size(2),
        "host_cache must be a contiguous uint8 CPU row matrix matching "
        "src_cache");
    cudaPointerAttributes attributes{};
    const cudaError_t pointer_status =
        cudaPointerGetAttributes(&attributes, host.const_data_ptr());
    if (pointer_status != cudaSuccess) {
      cudaGetLastError();
    }
    STD_TORCH_CHECK(
        pointer_status == cudaSuccess && attributes.type == cudaMemoryTypeHost,
        "host_cache must be CUDA-accessible pinned memory");
    STD_TORCH_CHECK(
        host_rows.is_cuda() && device_rows.is_cuda() &&
            host_rows.scalar_type() == torch::headeronly::ScalarType::Int &&
            device_rows.scalar_type() == torch::headeronly::ScalarType::Int &&
            host_rows.is_contiguous() && device_rows.is_contiguous() &&
            host_rows.numel() == device_rows.numel() && host_rows.dim() <= 2 &&
            device_rows.dim() <= 2,
        "host/device row maps must be matching contiguous CUDA int32 tensors");
    host_cache_ptr = reinterpret_cast<const uint8_t*>(host.const_data_ptr());
    host_row_ids_ptr = host_rows.const_data_ptr<int32_t>();
    device_row_ids_ptr = device_rows.const_data_ptr<int32_t>();
    host_entry_stride = host.stride(0);
  }
  auto dtype = src_cache.scalar_type();
  STD_TORCH_CHECK(
      dtype == torch::headeronly::ScalarType::Byte ||               // uint8
          dtype == torch::headeronly::ScalarType::Float8_e4m3fn ||  // fp8 e4m3
          dtype == torch::headeronly::ScalarType::Float8_e5m2,      // fp8 e5m2
      "src_cache must be uint8, float8_e4m3fn, or float8_e5m2, but got ",
      src_cache.scalar_type());
  STD_TORCH_CHECK(dst.scalar_type() == torch::headeronly::ScalarType::BFloat16,
                  "dst must be bfloat16");
  STD_TORCH_CHECK(head_dim == 576, "head_dim must be 576 for MLA");

  int64_t block_table_stride = block_table.stride(0);
  int64_t cache_block_stride = src_cache.stride(0);
  int64_t cache_entry_stride = src_cache.stride(1);
  int64_t dst_entry_stride = dst.stride(0);

  const uint8_t* src_ptr = nullptr;
  if (dtype == torch::headeronly::ScalarType::Byte) {
    src_ptr = src_cache.const_data_ptr<uint8_t>();
  } else {
    // float8_e4m3fn or float8_e5m2
    src_ptr = reinterpret_cast<const uint8_t*>(src_cache.data_ptr());
  }

  const int total_tokens = dst.size(0);
  if (total_tokens == 0) {
    return;
  }
  constexpr int warps_per_block = 16;
  const int block_size_threads = warps_per_block * 32;
  const int32_t* seq_starts_ptr =
      seq_starts.has_value() ? seq_starts.value().const_data_ptr<int32_t>()
                             : nullptr;
  const int required_blocks = cuda_utils::ceil_div(total_tokens, block_size) +
                              2 * static_cast<int32_t>(batch_size);
  vllm::cp_gather_and_upconvert_fp8_kv_cache_page<<<
      required_blocks, block_size_threads, 0, stream>>>(
      src_ptr, reinterpret_cast<__nv_bfloat16*>(dst.data_ptr()),
      block_table.const_data_ptr<int32_t>(),
      workspace_starts.const_data_ptr<int32_t>(),
      static_cast<int32_t>(batch_size), block_size, total_tokens,
      block_table_stride, cache_block_stride, cache_entry_stride,
      dst_entry_stride, seq_starts_ptr, host_cache_ptr, host_row_ids_ptr,
      device_row_ids_ptr, host_entry_stride);
}

void cp_gather_and_upconvert_nvfp4_kv_cache(
    torch::stable::Tensor const& src_cache,         // [NUM_BLOCKS, BLOCK_SIZE,
                                                    // 352]
    torch::stable::Tensor const& dst,               // [TOT_TOKENS, 576]
    torch::stable::Tensor const& block_table,       // [BATCH, BLOCK_INDICES]
    torch::stable::Tensor const& workspace_starts,  // [BATCH]
    int64_t batch_size) {
  // The kernel lives in nvfp4_ds_mla_cache_kernels.cu, which is only built when
  // the target list contains an SM100 arch.
#if !defined(USE_ROCM) && defined(ENABLE_NVFP4_SM100) && ENABLE_NVFP4_SM100
  torch::stable::accelerator::DeviceGuard device_guard(
      src_cache.get_device_index());
  const cudaStream_t stream = get_current_cuda_stream();

  const cudaDeviceProp* props = get_device_prop();
  STD_TORCH_CHECK(props->major == 10,
                  "nvfp4_ds_mla KV cache requires SM100 (Blackwell); "
                  "got SM",
                  props->major, props->minor);

  int32_t block_size = src_cache.size(1);
  int32_t head_dim = dst.size(1);

  STD_TORCH_CHECK(
      block_table.scalar_type() == torch::headeronly::ScalarType::Int,
      "block_table must be int32");
  STD_TORCH_CHECK(
      workspace_starts.scalar_type() == torch::headeronly::ScalarType::Int,
      "workspace_starts must be int32");
  STD_TORCH_CHECK(src_cache.device() == dst.device(),
                  "src_cache and dst must be on the same device");
  STD_TORCH_CHECK(src_cache.device() == block_table.device(),
                  "src_cache and block_table must be on the same device");
  STD_TORCH_CHECK(src_cache.device() == workspace_starts.device(),
                  "src_cache and workspace_starts must be on the same device");
  STD_TORCH_CHECK(
      src_cache.scalar_type() == torch::headeronly::ScalarType::Byte,
      "src_cache must be uint8 for the nvfp4 ds_mla format, but got ",
      src_cache.scalar_type());
  STD_TORCH_CHECK(dst.scalar_type() == torch::headeronly::ScalarType::BFloat16,
                  "dst must be bfloat16");
  STD_TORCH_CHECK(head_dim == 576, "head_dim must be 576 for MLA");
  STD_TORCH_CHECK(src_cache.size(2) == 352,
                  "src_cache.size(2) must be 352 for the nvfp4 ds_mla format");

  int64_t block_table_stride = block_table.stride(0);
  int64_t cache_block_stride = src_cache.stride(0);
  int64_t cache_entry_stride = src_cache.stride(1);
  int64_t dst_entry_stride = dst.stride(0);

  const int total_tokens = dst.size(0);

  vllm::launch_cp_gather_and_upconvert_nvfp4_kv_cache(
      src_cache.const_data_ptr<uint8_t>(), dst.data_ptr(),
      block_table.const_data_ptr<int32_t>(),
      workspace_starts.const_data_ptr<int32_t>(),
      static_cast<int32_t>(batch_size), block_size, total_tokens,
      block_table_stride, cache_block_stride, cache_entry_stride,
      dst_entry_stride, stream);
#else
  STD_TORCH_CHECK(false,
                  "the nvfp4 ds_mla kv-cache format requires a CUDA build "
                  "targeting SM100 (Blackwell)");
#endif
}

// Macro to dispatch the kernel based on the data type.
#define CALL_INDEXER_K_QUANT_AND_CACHE(KV_T, CACHE_T, KV_DTYPE)               \
  vllm::indexer_k_quant_and_cache_kernel<KV_T, CACHE_T, KV_DTYPE>             \
      <<<grid, block, 0, stream>>>(                                           \
          reinterpret_cast<KV_T*>(k.data_ptr()),                              \
          reinterpret_cast<CACHE_T*>(kv_cache.data_ptr()),                    \
          slot_mapping.const_data_ptr<int64_t>(), head_dim, quant_block_size, \
          cache_block_size, cache_block_stride, use_ue8m0);

void indexer_k_quant_and_cache(
    torch::stable::Tensor& k,         // [num_tokens, head_dim]
    torch::stable::Tensor& kv_cache,  // [num_blocks, block_size, cache_stride]
    torch::stable::Tensor& slot_mapping,  // [num_tokens]
    int64_t quant_block_size,             // quantization block size
    const std::string& scale_fmt) {
  int num_tokens = k.size(0);
  int head_dim = k.size(1);
  int cache_block_size = kv_cache.size(1);
  int64_t cache_block_stride = kv_cache.stride(0);
  bool use_ue8m0 = scale_fmt == "ue8m0";

  STD_TORCH_CHECK(k.device() == kv_cache.device(),
                  "k and kv_cache must be on the same device");
  STD_TORCH_CHECK(k.device() == slot_mapping.device(),
                  "k and slot_mapping must be on the same device");
  STD_TORCH_CHECK(head_dim % quant_block_size == 0,
                  "head_dim must be divisible by quant_block_size");

  constexpr int vec_size = 4;
  dim3 grid(num_tokens, (head_dim + quant_block_size * vec_size - 1) /
                            (quant_block_size * vec_size));
  dim3 block(32, vec_size);
  const torch::stable::accelerator::DeviceGuard device_guard(
      k.get_device_index());
  const cudaStream_t stream = get_current_cuda_stream();

  static const std::string kv_cache_dtype = "fp8_e4m3";
  DISPATCH_BY_KV_CACHE_DTYPE(k.scalar_type(), kv_cache_dtype,
                             CALL_INDEXER_K_QUANT_AND_CACHE);
}

// Macro to dispatch the kernel based on the data amount.
#define CALL_CP_GATHER_INDEXER_K_QUANT_CACHE(BLOCK_Y_SIZE)                    \
  vllm::cp_gather_indexer_k_quant_cache_kernel<BLOCK_Y_SIZE>                  \
      <<<dim3((num_tokens + BLOCK_Y_SIZE - 1) / BLOCK_Y_SIZE,                 \
              (head_dim + 8 * vec_size - 1) / (8 * vec_size)),                \
         dim3(8, BLOCK_Y_SIZE), 0, stream>>>(                                 \
          reinterpret_cast<char*>(kv_cache.data_ptr()),                       \
          reinterpret_cast<char*>(dst_k.data_ptr()),                          \
          reinterpret_cast<char*>(dst_scale.data_ptr()),                      \
          block_table.const_data_ptr<int32_t>(),                              \
          cu_seq_lens.const_data_ptr<int32_t>(), batch_size, dst_k.stride(0), \
          dst_k.size(1), kv_cache.stride(0), kv_cache.stride(1),              \
          kv_cache.size(1), block_table.size(1), num_tokens,                  \
          quant_block_size);

void cp_gather_indexer_k_quant_cache(
    const torch::stable::Tensor&
        kv_cache,                  // [num_blocks, block_size, cache_stride]
    torch::stable::Tensor& dst_k,  // [num_tokens, head_dim]
    torch::stable::Tensor&
        dst_scale,  // [num_tokens, head_dim / quant_block_size * 4]
    const torch::stable::Tensor& block_table,  // [batch_size, num_blocks]
    const torch::stable::Tensor& cu_seq_lens   // [batch_size + 1]
) {
  int batch_size = block_table.size(0);
  int num_tokens = dst_k.size(0);
  int head_dim = dst_k.size(1);
  int quant_block_size = head_dim * 4 / dst_scale.size(1);

  STD_TORCH_CHECK(kv_cache.device() == dst_k.device(),
                  "kv_cache and dst_k must be on the same device");
  STD_TORCH_CHECK(kv_cache.device() == dst_scale.device(),
                  "kv_cache and dst_scale must be on the same device");
  STD_TORCH_CHECK(kv_cache.device() == block_table.device(),
                  "kv_cache and block_table must be on the same device");
  STD_TORCH_CHECK(kv_cache.device() == cu_seq_lens.device(),
                  "kv_cache and cu_seq_lens must be on the same device");
  STD_TORCH_CHECK(head_dim % quant_block_size == 0,
                  "head_dim must be divisible by quant_block_size");

  constexpr int vec_size = 16;
  const torch::stable::accelerator::DeviceGuard device_guard(
      kv_cache.get_device_index());
  const cudaStream_t stream = get_current_cuda_stream();

  if (num_tokens < 32) {
    CALL_CP_GATHER_INDEXER_K_QUANT_CACHE(1);
  } else if (num_tokens < 64) {
    CALL_CP_GATHER_INDEXER_K_QUANT_CACHE(2);
  } else if (num_tokens < 128) {
    CALL_CP_GATHER_INDEXER_K_QUANT_CACHE(4);
  } else if (num_tokens < 256) {
    CALL_CP_GATHER_INDEXER_K_QUANT_CACHE(8);
  } else if (num_tokens < 512) {
    CALL_CP_GATHER_INDEXER_K_QUANT_CACHE(16);
  } else {
    CALL_CP_GATHER_INDEXER_K_QUANT_CACHE(32);
  }
}

// Concatenate ql_nope and q_pe into a contiguous q_out tensor for MLA/DSA.
// Replaces torch.cat((ql_nope, q_pe), dim=-1).
void concat_mla_q(
    torch::stable::Tensor& ql_nope,  // [num_tokens, num_heads, nope_dim]
    torch::stable::Tensor& q_pe,     // [num_tokens, num_heads, rope_dim]
    torch::stable::Tensor& q_out     // [num_tokens, num_heads, nope_dim +
                                     // rope_dim]
) {
  const int num_tokens = ql_nope.size(0);
  const int num_heads = ql_nope.size(1);
  const int nope_dim = ql_nope.size(2);
  const int rope_dim = q_pe.size(2);

  STD_TORCH_CHECK(nope_dim % 512 == 0,
                  "nope_dim must be a multiple of 512, got ", nope_dim);
  STD_TORCH_CHECK(rope_dim == 64, "rope_dim must be 64, got ", rope_dim);
  STD_TORCH_CHECK(q_out.size(2) == nope_dim + rope_dim);

  STD_TORCH_CHECK(ql_nope.stride(2) == 1,
                  "ql_nope must have stride 1 in dim 2");
  STD_TORCH_CHECK(q_pe.stride(2) == 1, "q_pe must have stride 1 in dim 2");
  STD_TORCH_CHECK(q_out.stride(2) == 1, "q_out must have stride 1 in dim 2");
  STD_TORCH_CHECK(
      ql_nope.scalar_type() == torch::headeronly::ScalarType::Half ||
          ql_nope.scalar_type() == torch::headeronly::ScalarType::BFloat16,
      "ql_nope must be float16 or bfloat16 dtype");

  if (num_tokens == 0) return;

  constexpr int warps_per_block = 8;
  const int total_warps = num_tokens * num_heads;
  const int grid_size = (total_warps + warps_per_block - 1) / warps_per_block;
  const int block_size = warps_per_block * 32;

  const torch::stable::accelerator::DeviceGuard device_guard(
      ql_nope.get_device_index());
  const cudaStream_t stream = get_current_cuda_stream();

  VLLM_STABLE_DISPATCH_HALF_TYPES(ql_nope.scalar_type(), "concat_mla_q", [&] {
    vllm::ConcatMLAQKernel<scalar_t, 512><<<grid_size, block_size, 0, stream>>>(
        q_out.mutable_data_ptr<scalar_t>(), ql_nope.const_data_ptr<scalar_t>(),
        q_pe.const_data_ptr<scalar_t>(), num_tokens, num_heads, q_out.stride(0),
        q_out.stride(1), ql_nope.stride(0), ql_nope.stride(1), q_pe.stride(0),
        q_pe.stride(1));
  });
}
