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
39 changes: 39 additions & 0 deletions examples/models/muse-glimmer/runtime/engine/sampling_cuda.cu
Original file line number Diff line number Diff line change
Expand Up @@ -252,6 +252,22 @@ __global__ void categorical_sample_kernel(
tokens[row] = static_cast<uint64_t>(low);
}

__global__ void accept_with_probability_kernel(
const float* probabilities,
int64_t count,
DeviceRngState* rng,
uint8_t* accepted) {
const int64_t index = static_cast<int64_t>(blockIdx.x) * blockDim.x +
threadIdx.x;
if (index >= count) {
return;
}
curandStatePhilox4_32_10_t local_rng;
curand_init(rng->seed, 0, rng->base + index, &local_rng);
accepted[index] =
uniform_from_uint32(curand(&local_rng)) < probabilities[index];
}

} // namespace

struct SamplingWorkspace::Impl {
Expand Down Expand Up @@ -587,4 +603,27 @@ cudaError_t categorical_sample(
return cudaGetLastError();
}

cudaError_t accept_with_probability(
const float* probabilities,
int64_t count,
uint8_t* accepted,
SamplingWorkspace& workspace,
cudaStream_t stream) {
if (probabilities == nullptr || accepted == nullptr || count <= 0 ||
workspace.impl_->rng == nullptr) {
return cudaErrorInvalidValue;
}
auto& state = *workspace.impl_;
advance_rng_kernel<<<1, 1, 0, stream>>>(state.rng, count);
cudaError_t error = cudaGetLastError();
if (error != cudaSuccess) {
return error;
}
const int blocks =
static_cast<int>((count + kSamplingThreads - 1) / kSamplingThreads);
accept_with_probability_kernel<<<blocks, kSamplingThreads, 0, stream>>>(
probabilities, count, state.rng, accepted);
return cudaGetLastError();
}

} // namespace muse_glimmer::cuda
15 changes: 15 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 @@ -55,6 +55,12 @@
uint64_t*,
SamplingWorkspace&,
cudaStream_t);
friend cudaError_t accept_with_probability(
const float*,
int64_t,
uint8_t*,
SamplingWorkspace&,
cudaStream_t);
};

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

// CUDA counterpart of muse_glimmer::accept_with_probability. Each probability gets an
// independent Philox draw and produces a byte-valued device result.
cudaError_t accept_with_probability(
const float* probabilities,
int64_t count,
uint8_t* accepted,
SamplingWorkspace& workspace,
cudaStream_t stream);

} // namespace muse_glimmer::cuda
55 changes: 55 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 @@ -263,4 +263,59 @@
ASSERT_CUDA_SUCCESS(cudaFree(device_probabilities));
}

TEST(CudaSamplingTest, AcceptanceMatchesHostSemantics) {
constexpr int64_t kSamplesPerProbability = 10000;
constexpr std::array<float, 4> kProbabilities = {0.0f, 0.25f, 0.75f, 1.0f};
std::vector<float> probabilities;
probabilities.reserve(kSamplesPerProbability * kProbabilities.size());
for (const float probability : kProbabilities) {
probabilities.insert(
probabilities.end(), kSamplesPerProbability, probability);
}

float* device_probabilities = nullptr;
uint8_t* device_accepted = nullptr;
ASSERT_CUDA_SUCCESS(cudaMalloc(
&device_probabilities, probabilities.size() * sizeof(float)));
ASSERT_CUDA_SUCCESS(
cudaMalloc(&device_accepted, probabilities.size() * sizeof(uint8_t)));
ASSERT_CUDA_SUCCESS(cudaMemcpy(
device_probabilities,
probabilities.data(),
probabilities.size() * sizeof(float),
cudaMemcpyHostToDevice));

muse_glimmer::cuda::SamplingWorkspace workspace;
ASSERT_CUDA_SUCCESS(workspace.reserve(1, 1, nullptr));
ASSERT_CUDA_SUCCESS(workspace.set_seed(5678, nullptr));
ASSERT_CUDA_SUCCESS(muse_glimmer::cuda::accept_with_probability(
device_probabilities,
probabilities.size(),
device_accepted,
workspace,
nullptr));
std::vector<uint8_t> accepted(probabilities.size());
ASSERT_CUDA_SUCCESS(cudaMemcpy(
accepted.data(),
device_accepted,
accepted.size() * sizeof(uint8_t),
cudaMemcpyDeviceToHost));

for (size_t probability_index = 0;
probability_index < kProbabilities.size();
++probability_index) {
int64_t accepted_count = 0;
const size_t begin = probability_index * kSamplesPerProbability;
for (size_t index = begin; index < begin + kSamplesPerProbability; ++index) {
accepted_count += accepted[index];
}
const double frequency =
static_cast<double>(accepted_count) / kSamplesPerProbability;
EXPECT_NEAR(frequency, kProbabilities[probability_index], 0.02);
}

ASSERT_CUDA_SUCCESS(cudaFree(device_accepted));
ASSERT_CUDA_SUCCESS(cudaFree(device_probabilities));
}

} // namespace
Loading