Skip to content
Merged
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
4 changes: 3 additions & 1 deletion src/ops/topk_gpu.cu
Original file line number Diff line number Diff line change
Expand Up @@ -259,7 +259,9 @@ namespace fastertransformer {

if (tid == 0) {
topks[ite] = total;
s_val[total.p] = cub::FpLimits<T>::Lowest();
// An exhausted reduction has no value to invalidate.
if (total.p != NOT_FOUND)
s_val[total.p] = cub::FpLimits<T>::Lowest();
}
__syncthreads();
}
Expand Down
66 changes: 66 additions & 0 deletions tests/ops_test.cc
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
#include <algorithm>
#include <cmath>
#include <limits>
#include "test_utils.h"
#include "ctranslate2/layers/attention.h"
#include "ctranslate2/ops/ops.h"
Expand Down Expand Up @@ -723,6 +725,70 @@ TEST_P(OpDeviceFPTest, TopK) {
expect_storage_eq(indices, expected_indices);
}

TEST_P(OpDeviceFPTest, TopKWithExhaustedCandidates) {
const Device device = GetParam().device;
const DataType dtype = GetParam().dtype;
const dim_t rows = 4;
const dim_t depth = 1025;
const float neg_inf = -std::numeric_limits<float>::infinity();
const float lowest = dtype == DataType::FLOAT16 ? -65504.f
: dtype == DataType::BFLOAT16 ? -std::ldexp(255.f, 120)
: std::numeric_limits<float>::lowest();

// Cover each specialized launcher and the generic k=3 branch. Exhaustion
// must also be safe for non-first rows and candidates outside block lane 0.
for (const int k : {1, 2, 3, 4, 6, 8, 10, 16, 32, 64}) {
for (const float masked : {neg_inf, lowest}) {
SCOPED_TRACE(::testing::Message() << "k=" << k << ", masked=" << masked);
std::vector<float> data(rows * depth, masked);
data[depth + 17] = 4.f;
data[2 * depth + 137] = 4.f;
data[2 * depth + 513] = 2.f;
for (int i = 0; i < k; ++i)
data[3 * depth + 100 + 3 * i] = float(k - i);

const StorageView expected_input({rows, depth}, data);
StorageView input = expected_input.to(dtype).to(device);
StorageView values(dtype, device);
StorageView indices(DataType::INT32, device);

const ops::TopK topk(k);
// Reuse the input: stage 1 temporarily mutates it and must restore it.
for (int repeat = 0; repeat < 2; ++repeat) {
topk(input, values, indices);
const StorageView result_values = values.to_float32().to(Device::CPU);
const StorageView result_indices = indices.to(Device::CPU);
assert_vector_eq(result_values.shape(), {rows, dim_t(k)});
assert_vector_eq(result_indices.shape(), {rows, dim_t(k)});
expect_storage_eq(input.to_float32().to(Device::CPU), expected_input);

for (dim_t row = 0; row < rows; ++row) {
SCOPED_TRACE(::testing::Message() << "row=" << row << ", repeat=" << repeat);
const int valid = row == 0 ? 0 : row == 1 ? 1 : row == 2 ? 2 : k;
for (int i = 0; i < k; ++i) {
const dim_t offset = row * k + i;
const int32_t index = result_indices.at<int32_t>(offset);
EXPECT_GE(index, 0);
EXPECT_LT(index, depth);
if (i < valid) {
const int32_t expected_index = row == 1 ? 17
: row == 2 ? (i == 0 ? 137 : 513) : 100 + 3 * i;
EXPECT_EQ(index, expected_index);
EXPECT_FLOAT_EQ(result_values.at<float>(offset),
data[row * depth + expected_index]);
} else {
// Preserve the existing CUDA sentinel value behavior without
// prescribing a tied/sentinel index as a public TopK contract.
EXPECT_FLOAT_EQ(result_values.at<float>(offset),
device == Device::CUDA ? lowest : masked);
}
}
}
}
}
}
}

TEST_P(OpDeviceTest, TopKVariableDepth) {
Device device = GetParam();
const int k = 3;
Expand Down
Loading