Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -1312,6 +1312,7 @@ if(VLLM_CPP_HIP)
src/vt/rocm/rocm_gemma4_expert_geglu.hip
src/vt/rocm/rocm_fp8_channel_gemv.hip
src/vt/rocm/rocm_moe_router.hip
src/vt/rocm/rocm_sample.hip
src/vt/rocm/rocm_ops.hip)
if(VLLM_CPP_HIP_ARCHITECTURES)
set_source_files_properties(
Expand All @@ -1326,6 +1327,7 @@ if(VLLM_CPP_HIP)
src/vt/rocm/rocm_gemma4_expert_geglu.hip
src/vt/rocm/rocm_fp8_channel_gemv.hip
src/vt/rocm/rocm_moe_router.hip
src/vt/rocm/rocm_sample.hip
src/vt/rocm/rocm_ops.hip
PROPERTIES HIP_ARCHITECTURES "${VLLM_CPP_HIP_ARCHITECTURES}")
endif()
Expand Down
1 change: 1 addition & 0 deletions docs/FEATURES.md
Original file line number Diff line number Diff line change
Expand Up @@ -302,6 +302,7 @@ CPU elementwise GEMM (f32/f16/bf16) runs AVX2 and AVX-512 tiers on x86 where the
| Custom logits processors on CUDA | Open, not root-caused | Segfaults in a CUDA build, 232/232 green on CPU |
| Memory budgeting (`ROAD-V1-MEM`, #83) | M1+M2 landed (absolute bytes) | `--kv-cache-memory` sizes the KV pool from an absolute byte budget (ABI v16, group-aware divisor); `--num-blocks` overrides; `--gpu-memory-utilization` needs the M3 profile run (dgx-gated). See `specs/kv-sizing.md` |
| Gemma4 MoE ROCm fused helpers (`vt::fused_ops`) | Partial | Portable ROCm seam. Public: `VT_GEMMA4_EXPERT_VRAM_MB` (positive-MiB LRU cap; unset/0 unlimited) + `VT_SERVER_MAX_{PROMPT_CHARS,NEW_TOKENS}` (200000/4096; 0 disables). Nine tuning vars internal; defaults unchanged |
| ROCm V1 sampler ops | Landed (this PR) | temperature/top-k/top-p/min-p/penalties/masks/logprobs/random; parallel gumbel-max sample |

## How to read this page

Expand Down
6 changes: 6 additions & 0 deletions docs/USAGE.md
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,12 @@ directory (issue #85).

### One ROCm-specific behaviour

ROCm builds register the full V1 sampler surface (temperature, top-k/top-p, min-p,
penalties, allowed-token masks, logprobs, random sample) so EngineCore does not
fatal with `no kernel for op` after prefill on AMD. Random sample uses a parallel
gumbel-max reduce (not a 1-thread full-vocab scan) so temp>0 decode stays in the
same class as greedy on large vocabs (e.g. Gemma-4).

Worth knowing before you read a hang as a bug in the tests: a build that sets no
`CMAKE_BUILD_TYPE` floors **HIP device code** at `-O1` and prints a configure
line saying so. At `-O0` the ROCm runtime starts a hostcall listener the kernels
Expand Down
39 changes: 39 additions & 0 deletions src/vt/rocm/rocm_dense_basic.hip
Original file line number Diff line number Diff line change
Expand Up @@ -767,4 +767,43 @@ void GeluErfKernelRocm(Queue& q, Tensor& out, const Tensor& x) {
Check(hipGetLastError(), "gelu_erf");
}

// --- sampling logit masks (Hermes allowed_token_ids / bad_words) ------------
// mask TRUE => exclude (set -inf). Mirrors CPU/CUDA ApplyAllowedTokenIds.
__global__ void ApplyAllowedTokenIdsK(float* logits, const int8_t* mask, int64_t total) {
const int64_t step = static_cast<int64_t>(gridDim.x) * blockDim.x;
for (int64_t idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x; idx < total;
idx += step)
if (mask[idx]) logits[idx] = -INFINITY;
}

void ApplyAllowedTokenIdsKernelRocm(Queue& q, Tensor& logits, const Tensor& mask) {
const int64_t n = logits.shape[0], v = logits.shape[1];
if (n == 0 || v == 0) return;
VT_CHECK(logits.dtype == DType::kF32, "rocm apply_allowed_token_ids: logits f32");
VT_CHECK(mask.dtype == DType::kI8, "rocm apply_allowed_token_ids: mask i8");
const int64_t total = n * v;
ApplyAllowedTokenIdsK<<<GridFor(total), kBlock, 0, AsStream(q)>>>(
logits.Ptr<float>(), mask.Ptr<int8_t>(), total);
Check(hipGetLastError(), "apply_allowed_token_ids");
}

// Sparse -inf scatter at (rows[k], cols[k]) — bad_words path.
__global__ void ApplyTokenMaskK(float* logits, const int32_t* rows, const int32_t* cols,
int64_t v, int64_t m) {
const int64_t step = static_cast<int64_t>(gridDim.x) * blockDim.x;
for (int64_t k = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x; k < m;
k += step)
logits[static_cast<int64_t>(rows[k]) * v + cols[k]] = -INFINITY;
}

void ApplyTokenMaskKernelRocm(Queue& q, Tensor& logits, const Tensor& rows,
const Tensor& cols) {
const int64_t v = logits.shape[1], m = rows.shape[0];
if (m == 0) return;
VT_CHECK(logits.dtype == DType::kF32, "rocm apply_token_mask: logits f32");
ApplyTokenMaskK<<<GridFor(m), kBlock, 0, AsStream(q)>>>(
logits.Ptr<float>(), rows.Ptr<int32_t>(), cols.Ptr<int32_t>(), v, m);
Check(hipGetLastError(), "apply_token_mask");
}

} // namespace vt::rocm
42 changes: 42 additions & 0 deletions src/vt/rocm/rocm_ops.hip
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,20 @@ void GeluTanhKernelRocm(Queue& q, Tensor& out, const Tensor& x);
void GeluErfKernelRocm(Queue& q, Tensor& out, const Tensor& x);
void MoeRouterTopKKernelRocm(Queue& q, Tensor& weights, Tensor& indices, const Tensor& logits,
const MoeRouterTopKArgs& args, const Tensor* bias);
void ApplyAllowedTokenIdsKernelRocm(Queue& q, Tensor& logits, const Tensor& mask);
void ApplyTokenMaskKernelRocm(Queue& q, Tensor& logits, const Tensor& rows, const Tensor& cols);
void ApplyTemperatureKernelRocm(Queue& q, Tensor& logits, const Tensor& temp, bool all_random);
void ApplyTopKTopPKernelRocm(Queue& q, Tensor& logits, const Tensor* k, const Tensor* p);
void ComputeProbsKernelRocm(Queue& q, Tensor& probs, const Tensor& logits);
void ComputeLogprobsKernelRocm(Queue& q, Tensor& logprobs, const Tensor& logits);
void RandomSampleKernelRocm(Queue& q, Tensor& token_ids, const Tensor& probs, const Tensor& seeds);
void ApplyPenaltiesKernelRocm(Queue& q, Tensor& logits, const Tensor& prompt_mask,
const Tensor& output_bin_counts, const Tensor& output_mask,
const Tensor& frequency_penalties, const Tensor& presence_penalties,
const Tensor& repetition_penalties);
void ApplyMinPKernelRocm(Queue& q, Tensor& logits, const Tensor& min_p);
void ApplyLogitBiasKernelRocm(Queue& q, Tensor& logits, const Tensor& rows, const Tensor& cols,
const Tensor& biases);

namespace {

Expand Down Expand Up @@ -95,6 +109,34 @@ struct Registrar {
RegisterOp(OpId::kMoeRouterTopK, DeviceType::kROCM,
reinterpret_cast<void*>(
static_cast<MoeRouterTopKFn>(&MoeRouterTopKKernelRocm)));
// Full V1 sampler surface for Hermes (temp/top-p/allowed ids/penalties/...).
RegisterOp(OpId::kApplyTemperature, DeviceType::kROCM,
reinterpret_cast<void*>(
static_cast<ApplyTemperatureFn>(&ApplyTemperatureKernelRocm)));
RegisterOp(OpId::kApplyTopKTopP, DeviceType::kROCM,
reinterpret_cast<void*>(
static_cast<ApplyTopKTopPFn>(&ApplyTopKTopPKernelRocm)));
RegisterOp(OpId::kComputeProbs, DeviceType::kROCM,
reinterpret_cast<void*>(static_cast<ComputeProbsFn>(&ComputeProbsKernelRocm)));
RegisterOp(OpId::kComputeLogprobs, DeviceType::kROCM,
reinterpret_cast<void*>(
static_cast<ComputeLogprobsFn>(&ComputeLogprobsKernelRocm)));
RegisterOp(OpId::kRandomSample, DeviceType::kROCM,
reinterpret_cast<void*>(static_cast<RandomSampleFn>(&RandomSampleKernelRocm)));
RegisterOp(OpId::kApplyPenalties, DeviceType::kROCM,
reinterpret_cast<void*>(
static_cast<ApplyPenaltiesFn>(&ApplyPenaltiesKernelRocm)));
RegisterOp(OpId::kApplyMinP, DeviceType::kROCM,
reinterpret_cast<void*>(static_cast<ApplyMinPFn>(&ApplyMinPKernelRocm)));
RegisterOp(OpId::kApplyLogitBias, DeviceType::kROCM,
reinterpret_cast<void*>(
static_cast<ApplyLogitBiasFn>(&ApplyLogitBiasKernelRocm)));
RegisterOp(OpId::kApplyAllowedTokenIds, DeviceType::kROCM,
reinterpret_cast<void*>(
static_cast<ApplyAllowedTokenIdsFn>(&ApplyAllowedTokenIdsKernelRocm)));
RegisterOp(OpId::kApplyTokenMask, DeviceType::kROCM,
reinterpret_cast<void*>(
static_cast<ApplyTokenMaskFn>(&ApplyTokenMaskKernelRocm)));
}
} registrar;

Expand Down
Loading
Loading