Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
Original file line number Diff line number Diff line change
Expand Up @@ -528,9 +528,18 @@ static bool MakeQDQNodeUnit(api::GraphRef& graph, const api::NodeRef& dq_node) {
inputs.push_back(zp_input.value());
}

// No zero-point: pin output_dtype to the DQ's type so the new Q doesn't default to uint8.
Comment thread
tianleiwu marked this conversation as resolved.
Outdated
std::optional<int64_t> q_output_dtype;
if (!zp_input.has_value() && IsOnnxDomain(dq_domain)) {
api::DataType dq_input_dtype = graph.GetValueInfo(dq_inputs[0])->DType();
if (dq_input_dtype != api::DataType::UNDEFINED && dq_input_dtype != api::DataType::UINT8) {
Comment thread
tianleiwu marked this conversation as resolved.
q_output_dtype = static_cast<int64_t>(dq_input_dtype);
}
}

// Add Q
auto new_q_node = MakeQuantizeOp(graph, dq_domain, inputs, axis, dq_node.GetAttributeInt("block_size"),
dq_node.GetAttributeInt("output_dtype"), dq_node.GetAttributeInt("saturate"));
q_output_dtype, dq_node.GetAttributeInt("saturate"));
new_q_node->SetLayeringAnnotation(dq_node.GetLayeringAnnotation());
auto q_node_outputs = new_q_node->Outputs();

Expand Down
26 changes: 26 additions & 0 deletions onnxruntime/test/optimizer/transpose_optimizer_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -3773,6 +3773,32 @@ TEST(TransposeOptimizerTests, TestDequantizeLinearNoAxis) {
#endif
}

// Regression test for #28716: pushing a Transpose through a zero-point-less int8 DequantizeLinear
// inserts a QuantizeLinear that must set output_dtype, else it defaults to uint8 and Resolve() fails.
TEST(TransposeOptimizerTests, TestDequantizeLinearNoZeroPoint) {
auto build_test_case = [&](ModelTestBuilder& builder) {
auto* input0_arg = MakeInput<int8_t>(builder, {{2, -1, 6, 3}}, {2, 4, 6, 3}, -128, 127);
auto* scale_arg = MakeInput<float>(builder, std::vector<int64_t>{}, std::vector<int64_t>{}, {0.05f});
auto* transpose_1_out_0 = builder.MakeIntermediate();
auto* dq_out_0 = builder.MakeIntermediate();
auto* transpose_2_out_0 = builder.MakeOutput();

auto& transpose_1 = builder.AddNode("Transpose", {input0_arg}, {transpose_1_out_0});
transpose_1.AddAttribute("perm", std::vector<int64_t>{0, 3, 1, 2});
builder.AddNode("DequantizeLinear", {transpose_1_out_0, scale_arg}, {dq_out_0}); // no zero-point
auto& transpose_2 = builder.AddNode("Transpose", {dq_out_0}, {transpose_2_out_0});
transpose_2.AddAttribute("perm", std::vector<int64_t>{0, 2, 3, 1});
};

auto check_optimized_graph = [](InferenceSessionWrapper& session) {
EXPECT_EQ(EstimateTransposeCost(session.GetGraph()), 0);
};

// output_dtype requires ONNX opset 21.
TransformerTester(build_test_case, check_optimized_graph, TransformerLevel::Default,
TransformerLevel::Level1, /*opsets*/ {21});
}

TEST(TransposeOptimizerTests, TestCast) {
auto build_test_case_1 = [&](ModelTestBuilder& builder) {
auto* input0_arg = MakeInput<int32_t>(builder, {{-1, 4, -1, 5}}, {2, 4, 6, 5}, -1, 5);
Expand Down
Loading