MCPcopy Create free account
hub / github.com/PaddlePaddle/Paddle / AddDepToBlockOp

Method AddDepToBlockOp

paddle/fluid/framework/program_utils.cc:134–188  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

132}
133
134void ProgramProcessor::AddDepToBlockOp(const BlockDesc &block) {
135 VLOG(3) << "Op size:" << block.AllOps().size();
136 for (OpDesc *op : block.AllOps()) {
137 if (op->HasAttr("sub_block")) {
138 auto op_type = op->Type();
139 BlockDesc *sub_block =
140 PADDLE_GET_MUTABLE(BlockDesc *, op->GetAttr("sub_block"));
141
142 // recursively processing
143 AddDepToBlockOp(*sub_block);
144
145 std::set<std::string> sub_inputs;
146 std::set<std::string> sub_outputs;
147 ProgramProcessor::GetInputsOutputsInBlock(
148 *sub_block, &sub_inputs, &sub_outputs);
149 VLOG(3) << "sub_inputs.size:" << sub_inputs.size();
150 VLOG(3) << "sub_outputs.size:" << sub_outputs.size();
151
152 auto *op_inputs = op->MutableInputs();
153 std::vector<std::string> *op_input_var_vec = nullptr;
154 VLOG(3) << "op_type:>>>>>>" << op_type;
155 if (op_type == "while") {
156 op_input_var_vec = &((*op_inputs)["kX"]);
157 } else if (op_type == "conditional_block") {
158 op_input_var_vec = &((*op_inputs)["kInputs"]);
159 } else {
160 // Only support while_op and conditional_block_op now
161 LOG(WARNING)
162 << "Currently, only support while_op and conditional_block_op.\n";
163 continue;
164 }
165
166 for (auto const &sub_input : sub_inputs) {
167 if (std::find(op_input_var_vec->begin(),
168 op_input_var_vec->end(),
169 sub_input) == op_input_var_vec->end())
170 op_input_var_vec->push_back(sub_input);
171 VLOG(3) << "modified private inputs, inputs.size():"
172 << op_input_var_vec->size();
173 }
174
175 auto *op_outputs = op->MutableOutputs();
176 auto *op_output_var_vec = &((*op_outputs)["kOutputs"]);
177
178 for (auto const &sub_output : sub_outputs) {
179 if (std::find(op_output_var_vec->begin(),
180 op_output_var_vec->end(),
181 sub_output) == op_output_var_vec->end())
182 op_output_var_vec->push_back(sub_output);
183 VLOG(3) << "modified private outputs, outputs.size():"
184 << op_output_var_vec->size();
185 }
186 }
187 }
188}
189
190ProgramProcessor::ProgramProcessor() = default;
191

Callers 1

TESTFunction · 0.80

Calls 8

AllOpsMethod · 0.80
sizeMethod · 0.45
HasAttrMethod · 0.45
TypeMethod · 0.45
GetAttrMethod · 0.45
beginMethod · 0.45
endMethod · 0.45
push_backMethod · 0.45

Tested by 1

TESTFunction · 0.64