Skip to content
Open
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
93 changes: 93 additions & 0 deletions examples/models/muse-glimmer/runtime/engine/sampling_cuda.cu
Original file line number Diff line number Diff line change
Expand Up @@ -268,6 +268,33 @@ __global__ void accept_with_probability_kernel(
uniform_from_uint32(curand(&local_rng)) < probabilities[index];
}

__global__ void exclude_tokens_kernel(
float* probabilities,
int64_t row_count,
int64_t row_size,
const uint64_t* excluded_tokens) {
const int64_t row = static_cast<int64_t>(blockIdx.x) * blockDim.x +
threadIdx.x;
if (row < row_count && excluded_tokens[row] < row_size) {
probabilities[row * row_size + excluded_tokens[row]] = 0.0f;
}
}

__global__ void normalize_probability_rows_kernel(
float* probabilities,
const double* cumulative,
int64_t total_size,
int64_t row_size) {
const int64_t offset = static_cast<int64_t>(blockIdx.x) * blockDim.x +
threadIdx.x;
if (offset < total_size) {
const int64_t row = offset / row_size;
const float denominator = static_cast<float>(
cumulative[row * row_size + row_size - 1]);
probabilities[offset] /= denominator;
}
}

} // namespace

struct SamplingWorkspace::Impl {
Expand Down Expand Up @@ -626,4 +653,70 @@ cudaError_t accept_with_probability(
return cudaGetLastError();
}

cudaError_t sample_excluding_token_in_place(
float* probabilities,
int64_t row_count,
int64_t row_size,
const uint64_t* excluded_tokens,
uint64_t* sampled_tokens,
SamplingWorkspace& workspace,
cudaStream_t stream) {
if (probabilities == nullptr || excluded_tokens == nullptr ||
sampled_tokens == nullptr || row_count <= 0 || row_size <= 0 ||
row_count > std::numeric_limits<int>::max() ||
row_size > std::numeric_limits<int>::max() / row_count) {
return cudaErrorInvalidValue;
}
cudaError_t error = workspace.reserve(row_count, row_size, stream);
if (error != cudaSuccess) {
return error;
}
auto& state = *workspace.impl_;
const int row_blocks = static_cast<int>(
(row_count + kSamplingThreads - 1) / kSamplingThreads);
exclude_tokens_kernel<<<row_blocks, kSamplingThreads, 0, stream>>>(
probabilities, row_count, row_size, excluded_tokens);
error = cudaGetLastError();
if (error != cudaSuccess) {
return error;
}

const int64_t total_size = state.total_size;
const int item_blocks = static_cast<int>(
(total_size + kSamplingThreads - 1) / kSamplingThreads);
probabilities_to_double_kernel<<<
item_blocks, kSamplingThreads, 0, stream>>>(
probabilities, total_size, state.weights);
error = cudaGetLastError();
if (error != cudaSuccess) {
return error;
}
for (int64_t row = 0; row < row_count; ++row) {
error = cub::DeviceScan::InclusiveSum(
state.temporary_storage,
state.temporary_storage_bytes,
state.weights + row * row_size,
state.cumulative + row * row_size,
static_cast<int>(row_size),
stream);
if (error != cudaSuccess) {
return error;
}
}
normalize_probability_rows_kernel<<<
item_blocks, kSamplingThreads, 0, stream>>>(
probabilities, state.cumulative, total_size, row_size);
error = cudaGetLastError();
if (error != cudaSuccess) {
return error;
}
return categorical_sample(
probabilities,
row_count,
row_size,
sampled_tokens,
workspace,
stream);
}

} // namespace muse_glimmer::cuda
19 changes: 19 additions & 0 deletions examples/models/muse-glimmer/runtime/engine/sampling_cuda.h
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
/*
* Copyright (c) Meta Platforms, Inc. and affiliates.
* All rights reserved.
Expand Down Expand Up @@ -61,6 +61,14 @@
uint8_t*,
SamplingWorkspace&,
cudaStream_t);
friend cudaError_t sample_excluding_token_in_place(
float*,
int64_t,
int64_t,
const uint64_t*,
uint64_t*,
SamplingWorkspace&,
cudaStream_t);
};

// Computes one argmax per contiguous row of `values`.
Expand Down Expand Up @@ -109,4 +117,15 @@
SamplingWorkspace& workspace,
cudaStream_t stream);

// CUDA counterpart of muse_glimmer::sample_excluding_token_in_place for batched rows.
// Mutates each probability row by excluding and renormalizing before sampling.
cudaError_t sample_excluding_token_in_place(
float* probabilities,
int64_t row_count,
int64_t row_size,
const uint64_t* excluded_tokens,
uint64_t* sampled_tokens,
SamplingWorkspace& workspace,
cudaStream_t stream);

} // namespace muse_glimmer::cuda
71 changes: 71 additions & 0 deletions examples/models/muse-glimmer/tests/sampling_cuda_test.cpp
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
/*
* Copyright (c) Meta Platforms, Inc. and affiliates.
* All rights reserved.
Expand Down Expand Up @@ -318,4 +318,75 @@
ASSERT_CUDA_SUCCESS(cudaFree(device_probabilities));
}

TEST(CudaSamplingTest, ExcludingTokenMatchesHostSemantics) {
constexpr int64_t kRows = 12000;
constexpr int64_t kRowSize = 3;
const std::array<float, kRowSize> distribution = {0.2f, 0.5f, 0.3f};
std::vector<float> probabilities(kRows * kRowSize);
for (int64_t row = 0; row < kRows; ++row) {
std::copy(
distribution.begin(),
distribution.end(),
probabilities.begin() + row * kRowSize);
}
std::vector<uint64_t> excluded(kRows, 1);

float* device_probabilities = nullptr;
uint64_t* device_excluded = nullptr;
uint64_t* device_tokens = nullptr;
ASSERT_CUDA_SUCCESS(cudaMalloc(
&device_probabilities, probabilities.size() * sizeof(float)));
ASSERT_CUDA_SUCCESS(
cudaMalloc(&device_excluded, excluded.size() * sizeof(uint64_t)));
ASSERT_CUDA_SUCCESS(
cudaMalloc(&device_tokens, kRows * sizeof(uint64_t)));
ASSERT_CUDA_SUCCESS(cudaMemcpy(
device_probabilities,
probabilities.data(),
probabilities.size() * sizeof(float),
cudaMemcpyHostToDevice));
ASSERT_CUDA_SUCCESS(cudaMemcpy(
device_excluded,
excluded.data(),
excluded.size() * sizeof(uint64_t),
cudaMemcpyHostToDevice));

muse_glimmer::cuda::SamplingWorkspace workspace;
ASSERT_CUDA_SUCCESS(workspace.reserve(kRows, kRowSize, nullptr));
ASSERT_CUDA_SUCCESS(workspace.set_seed(9012, nullptr));
ASSERT_CUDA_SUCCESS(muse_glimmer::cuda::sample_excluding_token_in_place(
device_probabilities,
kRows,
kRowSize,
device_excluded,
device_tokens,
workspace,
nullptr));
std::vector<uint64_t> tokens(kRows);
ASSERT_CUDA_SUCCESS(cudaMemcpy(
tokens.data(),
device_tokens,
tokens.size() * sizeof(uint64_t),
cudaMemcpyDeviceToHost));
ASSERT_CUDA_SUCCESS(cudaMemcpy(
probabilities.data(),
device_probabilities,
probabilities.size() * sizeof(float),
cudaMemcpyDeviceToHost));

int64_t token_zero_count = 0;
for (int64_t row = 0; row < kRows; ++row) {
EXPECT_NEAR(probabilities[row * kRowSize], 0.4f, 1e-6f);
EXPECT_EQ(probabilities[row * kRowSize + 1], 0.0f);
EXPECT_NEAR(probabilities[row * kRowSize + 2], 0.6f, 1e-6f);
ASSERT_NE(tokens[row], 1);
token_zero_count += tokens[row] == 0;
}
EXPECT_NEAR(static_cast<double>(token_zero_count) / kRows, 0.4, 0.02);

ASSERT_CUDA_SUCCESS(cudaFree(device_tokens));
ASSERT_CUDA_SUCCESS(cudaFree(device_excluded));
ASSERT_CUDA_SUCCESS(cudaFree(device_probabilities));
}

} // namespace
Loading