Skip to content
Merged
10 changes: 10 additions & 0 deletions onnxruntime/contrib_ops/cpu/bert/rotary_embedding_helper.h
Original file line number Diff line number Diff line change
Expand Up @@ -118,6 +118,16 @@ Status CheckInputs(const T* input,
if (head_size == 0) {
ORT_RETURN_IF_ERROR(detail::CheckedMulToInt32(cache_width, 2, "head_size", head_size));
}
// Validate that the rotary embedding dimension derived from cos_cache does not exceed hidden_size,
// which would cause an out-of-bounds read on the input tensor.
int effective_rotary_dim = cache_width * 2;
if (hidden_size > 0 && effective_rotary_dim > hidden_size) {
Comment thread
tianleiwu marked this conversation as resolved.
Outdated
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"RotaryEmbedding: cos_cache dimension (", cache_width,
" * 2 = ", effective_rotary_dim,
") exceeds input hidden_size (", hidden_size,
") when rotary_embedding_dim is 0");
}
} else {
if (!transposed) {
if (num_heads <= 0) {
Expand Down
28 changes: 28 additions & 0 deletions onnxruntime/test/contrib_ops/rotary_embedding_op_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -1172,5 +1172,33 @@ TEST(RotaryEmbeddingTest, ContribRotaryEmbedding_PositionIds_Negative_WebGPU_Pas
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}

// Test that cos_cache dimension exceeding hidden_size is rejected when rotary_embedding_dim=0.
TEST(RotaryEmbeddingTest, ContribRotaryEmbedding_RejectsCosCacheExceedsHiddenSize) {
Comment thread
apsonawane marked this conversation as resolved.
Outdated
int batch_size = 1;
int sequence_length = 1;
int hidden_size = 64;
int half_rotary_dim = 64; // makes cos_cache_dims[1]*2 = 128 > hidden_size
int max_sequence_length = 2;

OpTester test("RotaryEmbedding", 1, onnxruntime::kMSDomain);
test.AddAttribute<int64_t>("interleaved", static_cast<int64_t>(0));
test.AddAttribute<int64_t>("num_heads", static_cast<int64_t>(1));
Comment thread
apsonawane marked this conversation as resolved.
Outdated

test.AddInput<float>("input", {batch_size, sequence_length, hidden_size},
std::vector<float>(hidden_size, 42.0f));
test.AddInput<int64_t>("position_ids", {1}, {0});
test.AddInput<float>("cos_cache", {max_sequence_length, half_rotary_dim},
std::vector<float>(max_sequence_length * half_rotary_dim, 0.0f));
test.AddInput<float>("sin_cache", {max_sequence_length, half_rotary_dim},
std::vector<float>(max_sequence_length * half_rotary_dim, 1.0f));
test.AddOutput<float>("output", {batch_size, sequence_length, hidden_size},
std::vector<float>(hidden_size, 0.0f));

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

} // namespace test
} // namespace onnxruntime
29 changes: 29 additions & 0 deletions onnxruntime/test/providers/cpu/llm/rotary_embedding_op_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -1412,5 +1412,34 @@ TEST(RotaryEmbeddingTest, RotaryEmbedding_PositionIds_OOB_InBatch_WebGPU_Passthr
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}

// Test that cos_cache dimension exceeding hidden_size is rejected when rotary_embedding_dim=0.
TEST(RotaryEmbeddingTest, RotaryEmbedding_RejectsCosCacheExceedsHiddenSize) {
Comment thread
tianleiwu marked this conversation as resolved.
Outdated
// hidden_size = 64, cos_cache dim1 = 64 => effective rotary dim = 128 > 64
int batch_size = 1;
int sequence_length = 1;
int hidden_size = 64;
int half_rotary_dim = 64; // makes cos_cache_dims[1]*2 = 128 > hidden_size
int max_sequence_length = 2;

OpTester test("RotaryEmbedding", 23, onnxruntime::kOnnxDomain);
test.AddAttribute<int64_t>("interleaved", static_cast<int64_t>(0));
test.AddAttribute<int64_t>("num_heads", static_cast<int64_t>(1));

test.AddInput<float>("input", {batch_size, sequence_length, hidden_size},
std::vector<float>(hidden_size, 42.0f));
test.AddInput<float>("cos_cache", {max_sequence_length, half_rotary_dim},
std::vector<float>(max_sequence_length * half_rotary_dim, 0.0f));
test.AddInput<float>("sin_cache", {max_sequence_length, half_rotary_dim},
std::vector<float>(max_sequence_length * half_rotary_dim, 1.0f));
test.AddInput<int64_t>("position_ids", {1}, {0});
test.AddOutput<float>("output", {batch_size, sequence_length, hidden_size},
std::vector<float>(hidden_size, 0.0f));

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

} // namespace test
} // namespace onnxruntime
Loading