From 3ddf7434b6c9673aefa9cd26603ce36a88a84fb7 Mon Sep 17 00:00:00 2001 From: Akshay Sonawane Date: Tue, 11 Aug 2026 22:53:55 -0700 Subject: [PATCH 1/3] Avoid overflow in CPU TensorScatter indices --- .../core/providers/cpu/llm/tensorscatter.cc | 6 ++-- .../cpu/llm/tensorscatter_op_test.cc | 31 +++++++++++++++++++ 2 files changed, 35 insertions(+), 2 deletions(-) diff --git a/onnxruntime/core/providers/cpu/llm/tensorscatter.cc b/onnxruntime/core/providers/cpu/llm/tensorscatter.cc index a90c596ecf7ec..2b354ac89be3a 100644 --- a/onnxruntime/core/providers/cpu/llm/tensorscatter.cc +++ b/onnxruntime/core/providers/cpu/llm/tensorscatter.cc @@ -134,7 +134,7 @@ Status TensorScatter::Compute(OpKernelContext* context) const { uint8_t* cache_base = dst_bytes + cache_offset; if (!circular_) { - ORT_ENFORCE(wi + sequence_length <= max_sequence_length, + 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, ")"); // Single contiguous memcpy for the whole slice. @@ -143,8 +143,10 @@ Status TensorScatter::Compute(OpKernelContext* context) const { 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(SafeInt(cache_pos) * suffix_bytes); ptrdiff_t src_off = static_cast(SafeInt(s) * suffix_bytes); memcpy(cache_base + dst_off, update_base + src_off, suffix_bytes); diff --git a/onnxruntime/test/providers/cpu/llm/tensorscatter_op_test.cc b/onnxruntime/test/providers/cpu/llm/tensorscatter_op_test.cc index 5b2af54d309da..8a35586357e8a 100644 --- a/onnxruntime/test/providers/cpu/llm/tensorscatter_op_test.cc +++ b/onnxruntime/test/providers/cpu/llm/tensorscatter_op_test.cc @@ -1,6 +1,8 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. +#include + #include "gtest/gtest.h" #include "test/providers/provider_test_utils.h" @@ -335,6 +337,21 @@ TEST(TensorScatterTest, Linear_OutOfBoundsWriteIndex) { {}, nullptr, &execution_providers); } +TEST(TensorScatterTest, Linear_WriteIndexAdditionOverflow) { + OpTester test("TensorScatter", 24); + test.AddAttribute("mode", "linear"); + + test.AddInput("past_cache", {1, 4, 1}, {0, 0, 0, 0}); + test.AddInput("update", {1, 2, 1}, {1, 2}); + test.AddInput("write_indices", {1}, {std::numeric_limits::max()}); + test.AddOutput("present_cache", {1, 4, 1}, {0, 0, 0, 0}); + + std::vector> 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) { @@ -354,6 +371,20 @@ TEST(TensorScatterTest, Circular_NegativeWriteIndex) { {}, nullptr, &execution_providers); } +TEST(TensorScatterTest, Circular_LargeWriteIndexWrapsWithoutOverflow) { + OpTester test("TensorScatter", 24); + test.AddAttribute("mode", "circular"); + + test.AddInput("past_cache", {1, 4, 1}, {0, 0, 0, 0}); + test.AddInput("update", {1, 2, 1}, {1, 2}); + test.AddInput("write_indices", {1}, {std::numeric_limits::max()}); + test.AddOutput("present_cache", {1, 4, 1}, {2, 0, 0, 1}); + + std::vector> execution_providers; + execution_providers.push_back(DefaultCpuExecutionProvider()); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); +} + // 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. From b2f04fb9eef1fd239675071fbdf360fd27f6c904 Mon Sep 17 00:00:00 2001 From: Akshay Sonawane Date: Thu, 13 Aug 2026 14:18:22 -0700 Subject: [PATCH 2/3] Handle empty TensorScatter updates --- .../core/providers/cpu/llm/tensorscatter.cc | 15 +++- .../cpu/llm/tensorscatter_op_test.cc | 80 +++++++++++++++++++ 2 files changed, 94 insertions(+), 1 deletion(-) diff --git a/onnxruntime/core/providers/cpu/llm/tensorscatter.cc b/onnxruntime/core/providers/cpu/llm/tensorscatter.cc index 2b354ac89be3a..4eb4ceafd8274 100644 --- a/onnxruntime/core/providers/cpu/llm/tensorscatter.cc +++ b/onnxruntime/core/providers/cpu/llm/tensorscatter.cc @@ -83,12 +83,25 @@ Status TensorScatter::Compute(OpKernelContext* context) const { const size_t total_bytes = SafeInt(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) { + for (int64_t batch_idx = 0; batch_idx < batch_size; ++batch_idx) { + const int64_t wi = write_indices != nullptr ? write_indices[batch_idx] : 0; + ORT_ENFORCE(wi >= 0, "TensorScatter: write_indices[", batch_idx, "] = ", wi, " is negative"); + if (!circular_) { + ORT_ENFORCE(wi <= max_sequence_length, + "TensorScatter linear mode: write_indices[", batch_idx, "] + sequence_length (", + wi, " + 0) exceeds max_sequence_length (", max_sequence_length, ")"); + } + } + 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}) diff --git a/onnxruntime/test/providers/cpu/llm/tensorscatter_op_test.cc b/onnxruntime/test/providers/cpu/llm/tensorscatter_op_test.cc index 8a35586357e8a..f13f4904f81ce 100644 --- a/onnxruntime/test/providers/cpu/llm/tensorscatter_op_test.cc +++ b/onnxruntime/test/providers/cpu/llm/tensorscatter_op_test.cc @@ -385,6 +385,86 @@ TEST(TensorScatterTest, Circular_LargeWriteIndexWrapsWithoutOverflow) { test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); } +TEST(TensorScatterTest, Circular_ZeroSequenceLengthIsNoOp) { + OpTester test("TensorScatter", 24); + test.AddAttribute("mode", "circular"); + + test.AddInput("past_cache", {1, 0, 1}, {}); + test.AddInput("update", {1, 0, 1}, {}); + test.AddInput("write_indices", {1}, {std::numeric_limits::max()}); + test.AddOutput("present_cache", {1, 0, 1}, {}); + + std::vector> execution_providers; + execution_providers.push_back(DefaultCpuExecutionProvider()); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); +} + +TEST(TensorScatterTest, Circular_ZeroSequenceLengthPreservesCache) { + OpTester test("TensorScatter", 24); + test.AddAttribute("mode", "circular"); + + test.AddInput("past_cache", {1, 4, 1}, {1, 2, 3, 4}); + test.AddInput("update", {1, 0, 1}, {}); + test.AddInput("write_indices", {1}, {std::numeric_limits::max()}); + test.AddOutput("present_cache", {1, 4, 1}, {1, 2, 3, 4}); + + std::vector> 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("mode", "linear"); + + test.AddInput("past_cache", {1, 4, 1}, {1, 2, 3, 4}); + test.AddInput("update", {1, 0, 1}, {}); + test.AddInput("write_indices", {1}, {4}); + test.AddOutput("present_cache", {1, 4, 1}, {1, 2, 3, 4}); + + std::vector> 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("mode", "linear"); + + test.AddInput("past_cache", {1, 4, 1}, {1, 2, 3, 4}); + test.AddInput("update", {1, 0, 1}, {}); + test.AddInput("write_indices", {1}, {std::numeric_limits::max()}); + test.AddOutput("present_cache", {1, 4, 1}, {1, 2, 3, 4}); + + std::vector> 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("mode", mode); + + test.AddInput("past_cache", {1, 4, 1}, {1, 2, 3, 4}); + test.AddInput("update", {1, 0, 1}, {}); + test.AddInput("write_indices", {1}, {-1}); + test.AddOutput("present_cache", {1, 4, 1}, {1, 2, 3, 4}); + + std::vector> 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. From c999d3d4c005f4dc3aadbbb16bc532526908d6de Mon Sep 17 00:00:00 2001 From: Akshay Sonawane Date: Mon, 17 Aug 2026 13:16:32 -0700 Subject: [PATCH 3/3] Validate TensorScatter indices before updates --- .../core/providers/cpu/llm/tensorscatter.cc | 25 +++++++-------- .../cpu/llm/tensorscatter_op_test.cc | 31 +++++++++++++++++++ 2 files changed, 43 insertions(+), 13 deletions(-) diff --git a/onnxruntime/core/providers/cpu/llm/tensorscatter.cc b/onnxruntime/core/providers/cpu/llm/tensorscatter.cc index 4eb4ceafd8274..4f0262db06798 100644 --- a/onnxruntime/core/providers/cpu/llm/tensorscatter.cc +++ b/onnxruntime/core/providers/cpu/llm/tensorscatter.cc @@ -75,6 +75,18 @@ Status TensorScatter::Compute(OpKernelContext* context) const { write_indices = write_indices_tensor->Data(); } + 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); @@ -90,15 +102,6 @@ Status TensorScatter::Compute(OpKernelContext* context) const { } if (sequence_length == 0) { - for (int64_t batch_idx = 0; batch_idx < batch_size; ++batch_idx) { - const int64_t wi = write_indices != nullptr ? write_indices[batch_idx] : 0; - ORT_ENFORCE(wi >= 0, "TensorScatter: write_indices[", batch_idx, "] = ", wi, " is negative"); - if (!circular_) { - ORT_ENFORCE(wi <= max_sequence_length, - "TensorScatter linear mode: write_indices[", batch_idx, "] + sequence_length (", - wi, " + 0) exceeds max_sequence_length (", max_sequence_length, ")"); - } - } return Status::OK(); } @@ -139,7 +142,6 @@ 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(SafeInt(p) * update_axis_stride); ptrdiff_t cache_offset = static_cast(SafeInt(p) * cache_axis_stride); @@ -147,9 +149,6 @@ Status TensorScatter::Compute(OpKernelContext* context) const { uint8_t* cache_base = dst_bytes + cache_offset; 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, ")"); // Single contiguous memcpy for the whole slice. ptrdiff_t wi_offset = static_cast(SafeInt(wi) * suffix_bytes); size_t copy_len = SafeInt(sequence_length) * suffix_bytes; diff --git a/onnxruntime/test/providers/cpu/llm/tensorscatter_op_test.cc b/onnxruntime/test/providers/cpu/llm/tensorscatter_op_test.cc index f13f4904f81ce..7d3d010c93084 100644 --- a/onnxruntime/test/providers/cpu/llm/tensorscatter_op_test.cc +++ b/onnxruntime/test/providers/cpu/llm/tensorscatter_op_test.cc @@ -399,6 +399,37 @@ TEST(TensorScatterTest, Circular_ZeroSequenceLengthIsNoOp) { test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); } +TEST(TensorScatterTest, ZeroSequenceLengthWithoutWriteIndicesSkipsLargeBatch) { + OpTester test("TensorScatter", 24); + test.AddAttribute("mode", "circular"); + + constexpr int64_t large_batch = std::numeric_limits::max(); + test.AddInput("past_cache", {large_batch, 0, 1}, {}); + test.AddInput("update", {large_batch, 0, 1}, {}); + test.AddOptionalInputEdge(); + test.AddOutput("present_cache", {large_batch, 0, 1}, {}); + + std::vector> execution_providers; + execution_providers.push_back(DefaultCpuExecutionProvider()); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); +} + +TEST(TensorScatterTest, ExplicitWriteIndicesValidatedWithZeroPrefixDimension) { + OpTester test("TensorScatter", 24); + test.AddAttribute("axis", 2); + test.AddAttribute("mode", "circular"); + + test.AddInput("past_cache", {2, 0, 4, 1}, {}); + test.AddInput("update", {2, 0, 1, 1}, {}); + test.AddInput("write_indices", {2}, {-1, 0}); + test.AddOutput("present_cache", {2, 0, 4, 1}, {}); + + std::vector> 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("mode", "circular");