diff --git a/examples/models/muse-glimmer/runtime/engine/sampling_cuda.cu b/examples/models/muse-glimmer/runtime/engine/sampling_cuda.cu index adaf9825cdf..e776200ddaf 100644 --- a/examples/models/muse-glimmer/runtime/engine/sampling_cuda.cu +++ b/examples/models/muse-glimmer/runtime/engine/sampling_cuda.cu @@ -252,6 +252,22 @@ __global__ void categorical_sample_kernel( tokens[row] = static_cast(low); } +__global__ void accept_with_probability_kernel( + const float* probabilities, + int64_t count, + DeviceRngState* rng, + uint8_t* accepted) { + const int64_t index = static_cast(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 { @@ -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((count + kSamplingThreads - 1) / kSamplingThreads); + accept_with_probability_kernel<<>>( + probabilities, count, state.rng, accepted); + return cudaGetLastError(); +} + } // namespace muse_glimmer::cuda diff --git a/examples/models/muse-glimmer/runtime/engine/sampling_cuda.h b/examples/models/muse-glimmer/runtime/engine/sampling_cuda.h index 0406a4bfc4a..6f61327bb88 100644 --- a/examples/models/muse-glimmer/runtime/engine/sampling_cuda.h +++ b/examples/models/muse-glimmer/runtime/engine/sampling_cuda.h @@ -55,6 +55,12 @@ class SamplingWorkspace { 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`. @@ -94,4 +100,13 @@ cudaError_t categorical_sample( 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 diff --git a/examples/models/muse-glimmer/tests/sampling_cuda_test.cpp b/examples/models/muse-glimmer/tests/sampling_cuda_test.cpp index 5994092d024..bcdd2d51887 100644 --- a/examples/models/muse-glimmer/tests/sampling_cuda_test.cpp +++ b/examples/models/muse-glimmer/tests/sampling_cuda_test.cpp @@ -263,4 +263,59 @@ TEST(CudaSamplingTest, CategoricalSampleMatchesHostDistribution) { ASSERT_CUDA_SUCCESS(cudaFree(device_probabilities)); } +TEST(CudaSamplingTest, AcceptanceMatchesHostSemantics) { + constexpr int64_t kSamplesPerProbability = 10000; + constexpr std::array kProbabilities = {0.0f, 0.25f, 0.75f, 1.0f}; + std::vector 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 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(accepted_count) / kSamplesPerProbability; + EXPECT_NEAR(frequency, kProbabilities[probability_index], 0.02); + } + + ASSERT_CUDA_SUCCESS(cudaFree(device_accepted)); + ASSERT_CUDA_SUCCESS(cudaFree(device_probabilities)); +} + } // namespace