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
26 changes: 20 additions & 6 deletions onnxruntime/core/providers/cpu/llm/tensorscatter.cc
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,18 @@ Status TensorScatter::Compute(OpKernelContext* context) const {
write_indices = write_indices_tensor->Data<int64_t>();
}

if (write_indices != nullptr) {
for (int64_t batch_idx = 0; batch_idx < batch_size; ++batch_idx) {
const int64_t wi = write_indices[batch_idx];
ORT_ENFORCE(wi >= 0, "TensorScatter: write_indices[", batch_idx, "] = ", wi, " is negative");
if (!circular_) {
ORT_ENFORCE(wi <= max_sequence_length - sequence_length,
"TensorScatter linear mode: write_indices[", batch_idx, "] + sequence_length (",
wi, " + ", sequence_length, ") exceeds max_sequence_length (", max_sequence_length, ")");
}
}
}

// Allocate output with the same shape as past_cache.
Tensor* present_cache = context->Output(0, cache_shape);

Expand All @@ -83,12 +95,16 @@ Status TensorScatter::Compute(OpKernelContext* context) const {
const size_t total_bytes = SafeInt<size_t>(cache_shape.Size()) * element_size;
const auto* src_raw = past_cache->DataRaw();
auto* dst_raw = present_cache->MutableDataRaw();
if (dst_raw != src_raw) {
if (dst_raw != src_raw && total_bytes > 0) {
LOGS(context->Logger(), WARNING) << "TensorScatter: in-place optimization not activated, copying past_cache to present_cache ("
<< total_bytes << " bytes)";
memcpy(dst_raw, src_raw, total_bytes);
}

if (sequence_length == 0) {
return Status::OK();
}

// Step 2: Scatter the update into present_cache.
//
// Layout: (batch_size, D1, ..., D_{axis-1}, max_seq_len, D_{axis+1}, ..., D_{n-1})
Expand Down Expand Up @@ -126,25 +142,23 @@ Status TensorScatter::Compute(OpKernelContext* context) const {
for (int64_t p = 0; p < prefix_count; ++p) {
int64_t batch_idx = p / prefix_stride_for_batch;
int64_t wi = (write_indices != nullptr) ? write_indices[batch_idx] : 0;
ORT_ENFORCE(wi >= 0, "TensorScatter: write_indices[", batch_idx, "] = ", wi, " is negative");

ptrdiff_t update_offset = static_cast<ptrdiff_t>(SafeInt<size_t>(p) * update_axis_stride);
ptrdiff_t cache_offset = static_cast<ptrdiff_t>(SafeInt<size_t>(p) * cache_axis_stride);
const uint8_t* update_base = update_raw + update_offset;
uint8_t* cache_base = dst_bytes + cache_offset;

if (!circular_) {
ORT_ENFORCE(wi + sequence_length <= max_sequence_length,
"TensorScatter linear mode: write_indices[", batch_idx, "] + sequence_length (",
wi, " + ", sequence_length, ") exceeds max_sequence_length (", max_sequence_length, ")");
// Single contiguous memcpy for the whole slice.
ptrdiff_t wi_offset = static_cast<ptrdiff_t>(SafeInt<size_t>(wi) * suffix_bytes);
size_t copy_len = SafeInt<size_t>(sequence_length) * suffix_bytes;
memcpy(cache_base + wi_offset, update_base, copy_len);
} else {
// Circular: each sequence position wraps independently.
const int64_t wi_mod = wi % max_sequence_length;
const int64_t distance_to_end = max_sequence_length - wi_mod;
for (int64_t s = 0; s < sequence_length; ++s) {
int64_t cache_pos = (wi + s) % max_sequence_length;
const int64_t cache_pos = s >= distance_to_end ? s - distance_to_end : wi_mod + s;
ptrdiff_t dst_off = static_cast<ptrdiff_t>(SafeInt<size_t>(cache_pos) * suffix_bytes);
ptrdiff_t src_off = static_cast<ptrdiff_t>(SafeInt<size_t>(s) * suffix_bytes);
memcpy(cache_base + dst_off, update_base + src_off, suffix_bytes);
Expand Down
142 changes: 142 additions & 0 deletions onnxruntime/test/providers/cpu/llm/tensorscatter_op_test.cc
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.

#include <limits>

#include "gtest/gtest.h"
#include "test/providers/provider_test_utils.h"

Expand Down Expand Up @@ -335,6 +337,21 @@ TEST(TensorScatterTest, Linear_OutOfBoundsWriteIndex) {
{}, nullptr, &execution_providers);
}

TEST(TensorScatterTest, Linear_WriteIndexAdditionOverflow) {
OpTester test("TensorScatter", 24);
test.AddAttribute<std::string>("mode", "linear");

test.AddInput<float>("past_cache", {1, 4, 1}, {0, 0, 0, 0});
test.AddInput<float>("update", {1, 2, 1}, {1, 2});
test.AddInput<int64_t>("write_indices", {1}, {std::numeric_limits<int64_t>::max()});
test.AddOutput<float>("present_cache", {1, 4, 1}, {0, 0, 0, 0});

std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultCpuExecutionProvider());
test.Run(OpTester::ExpectResult::kExpectFailure, "exceeds max_sequence_length",
{}, nullptr, &execution_providers);
}

// Circular mode: negative write_indices should still fail.
// Run CPU-only: CUDA validates asynchronously via CUDA_KERNEL_ASSERT.
TEST(TensorScatterTest, Circular_NegativeWriteIndex) {
Expand All @@ -354,6 +371,131 @@ TEST(TensorScatterTest, Circular_NegativeWriteIndex) {
{}, nullptr, &execution_providers);
}

TEST(TensorScatterTest, Circular_LargeWriteIndexWrapsWithoutOverflow) {
OpTester test("TensorScatter", 24);
test.AddAttribute<std::string>("mode", "circular");

test.AddInput<float>("past_cache", {1, 4, 1}, {0, 0, 0, 0});
test.AddInput<float>("update", {1, 2, 1}, {1, 2});
test.AddInput<int64_t>("write_indices", {1}, {std::numeric_limits<int64_t>::max()});
test.AddOutput<float>("present_cache", {1, 4, 1}, {2, 0, 0, 1});

std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultCpuExecutionProvider());
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}

TEST(TensorScatterTest, Circular_ZeroSequenceLengthIsNoOp) {
OpTester test("TensorScatter", 24);
test.AddAttribute<std::string>("mode", "circular");

test.AddInput<float>("past_cache", {1, 0, 1}, {});
test.AddInput<float>("update", {1, 0, 1}, {});
test.AddInput<int64_t>("write_indices", {1}, {std::numeric_limits<int64_t>::max()});
test.AddOutput<float>("present_cache", {1, 0, 1}, {});

std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultCpuExecutionProvider());
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}

TEST(TensorScatterTest, ZeroSequenceLengthWithoutWriteIndicesSkipsLargeBatch) {
OpTester test("TensorScatter", 24);
test.AddAttribute<std::string>("mode", "circular");

constexpr int64_t large_batch = std::numeric_limits<int64_t>::max();
test.AddInput<float>("past_cache", {large_batch, 0, 1}, {});
test.AddInput<float>("update", {large_batch, 0, 1}, {});
test.AddOptionalInputEdge<int64_t>();
test.AddOutput<float>("present_cache", {large_batch, 0, 1}, {});

std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultCpuExecutionProvider());
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}

TEST(TensorScatterTest, ExplicitWriteIndicesValidatedWithZeroPrefixDimension) {
OpTester test("TensorScatter", 24);
test.AddAttribute<int64_t>("axis", 2);
test.AddAttribute<std::string>("mode", "circular");

test.AddInput<float>("past_cache", {2, 0, 4, 1}, {});
test.AddInput<float>("update", {2, 0, 1, 1}, {});
test.AddInput<int64_t>("write_indices", {2}, {-1, 0});
test.AddOutput<float>("present_cache", {2, 0, 4, 1}, {});

std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultCpuExecutionProvider());
test.Run(OpTester::ExpectResult::kExpectFailure, "is negative",
{}, nullptr, &execution_providers);
}

TEST(TensorScatterTest, Circular_ZeroSequenceLengthPreservesCache) {
OpTester test("TensorScatter", 24);
test.AddAttribute<std::string>("mode", "circular");

test.AddInput<float>("past_cache", {1, 4, 1}, {1, 2, 3, 4});
test.AddInput<float>("update", {1, 0, 1}, {});
test.AddInput<int64_t>("write_indices", {1}, {std::numeric_limits<int64_t>::max()});
test.AddOutput<float>("present_cache", {1, 4, 1}, {1, 2, 3, 4});

std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultCpuExecutionProvider());
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}

TEST(TensorScatterTest, Linear_ZeroSequenceLengthPreservesCache) {
OpTester test("TensorScatter", 24);
test.AddAttribute<std::string>("mode", "linear");

test.AddInput<float>("past_cache", {1, 4, 1}, {1, 2, 3, 4});
test.AddInput<float>("update", {1, 0, 1}, {});
test.AddInput<int64_t>("write_indices", {1}, {4});
test.AddOutput<float>("present_cache", {1, 4, 1}, {1, 2, 3, 4});

std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultCpuExecutionProvider());
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}

TEST(TensorScatterTest, Linear_ZeroSequenceLengthRejectsOutOfBoundsIndex) {
OpTester test("TensorScatter", 24);
test.AddAttribute<std::string>("mode", "linear");

test.AddInput<float>("past_cache", {1, 4, 1}, {1, 2, 3, 4});
test.AddInput<float>("update", {1, 0, 1}, {});
test.AddInput<int64_t>("write_indices", {1}, {std::numeric_limits<int64_t>::max()});
test.AddOutput<float>("present_cache", {1, 4, 1}, {1, 2, 3, 4});

std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultCpuExecutionProvider());
test.Run(OpTester::ExpectResult::kExpectFailure, "exceeds max_sequence_length",
{}, nullptr, &execution_providers);
}

static void RunZeroSequenceLengthNegativeWriteIndexTest(const std::string& mode) {
OpTester test("TensorScatter", 24);
test.AddAttribute<std::string>("mode", mode);

test.AddInput<float>("past_cache", {1, 4, 1}, {1, 2, 3, 4});
test.AddInput<float>("update", {1, 0, 1}, {});
test.AddInput<int64_t>("write_indices", {1}, {-1});
test.AddOutput<float>("present_cache", {1, 4, 1}, {1, 2, 3, 4});

std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultCpuExecutionProvider());
test.Run(OpTester::ExpectResult::kExpectFailure, "is negative",
{}, nullptr, &execution_providers);
}

TEST(TensorScatterTest, Linear_ZeroSequenceLengthRejectsNegativeIndex) {
RunZeroSequenceLengthNegativeWriteIndexTest("linear");
}

TEST(TensorScatterTest, Circular_ZeroSequenceLengthRejectsNegativeIndex) {
RunZeroSequenceLengthNegativeWriteIndexTest("circular");
}

// The CPU kernel only supports fixed-size element types (matching the CUDA
// kernel's type constraint). Non-fixed-size element types such as string are
// intentionally excluded because the kernel operates on raw memory buffers.
Expand Down
Loading