| 132 | } |
| 133 | |
| 134 | void 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 | |
| 190 | ProgramProcessor::ProgramProcessor() = default; |
| 191 | |