From 3455580780bdfcd74108d4f5e23889e924c8e584 Mon Sep 17 00:00:00 2001 From: Maxim Vafin Date: Thu, 12 Oct 2023 11:28:03 +0200 Subject: [PATCH] [PT FE] Fix pad default value (#20401) --- .../pytorch/src/transforms/prim_list_construct_pad.cpp | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/src/frontends/pytorch/src/transforms/prim_list_construct_pad.cpp b/src/frontends/pytorch/src/transforms/prim_list_construct_pad.cpp index 2147ddc17d9..74b4e8cc03b 100644 --- a/src/frontends/pytorch/src/transforms/prim_list_construct_pad.cpp +++ b/src/frontends/pytorch/src/transforms/prim_list_construct_pad.cpp @@ -71,13 +71,19 @@ PrimListConstructPadReplacer::PrimListConstructPadReplacer() { input_node = pad_op->input_value(0); padding = pad_op->input_value(1); auto mode_node = pad_op->input_value(2).get_node_shared_ptr(); - pad_value = pad_op->input_value(3); if (const auto& fw_node_mode = cast_fw_node(mode_node, "prim::Constant")) { const auto& attrs = fw_node_mode->get_attrs(); if (attrs.find("string_value") != attrs.end()) { mode = attrs.at("string_value"); } } + pad_value = pad_op->input_value(3); + if (const auto& fw_node_mode = cast_fw_node(pad_value.get_node_shared_ptr(), "prim::Constant")) { + const auto& attrs = fw_node_mode->get_attrs(); + if (attrs.find("none_value") != attrs.end()) { + pad_value = v0::Constant::create(element::f32, Shape{}, {0}); + } + } } else if ((pad_op = cast_fw_node(m.get_match_root(), "aten::reflection_pad2d"))) { mode = "reflect"; input_node = pad_op->input_value(0);