(self, new_desc)
| 6130 | self._grad_var_to_var = None |
| 6131 | |
| 6132 | def _find_var_class_kwargs(self, new_desc): |
| 6133 | # NOTE: not all variables support shape/dtype/lod_level methods. |
| 6134 | # For example: RAW, STEP_SCOPES, etc. |
| 6135 | def get_var_desc_attr_or_none(var_desc, attr_name, allowed_types): |
| 6136 | if var_desc.type() in allowed_types: |
| 6137 | return getattr(var_desc, attr_name)() |
| 6138 | else: |
| 6139 | return None |
| 6140 | |
| 6141 | old_desc = self.desc |
| 6142 | all_new_vars = [] |
| 6143 | block_num = new_desc.num_blocks() |
| 6144 | for idx in range(block_num): |
| 6145 | if idx > (len(self.blocks) - 1): |
| 6146 | self._create_block() |
| 6147 | new_block_desc = new_desc.block(idx) |
| 6148 | all_new_vars.append([]) |
| 6149 | block_new_vars = all_new_vars[-1] |
| 6150 | for new_var_desc in new_block_desc.all_vars(): |
| 6151 | if self.blocks[idx].has_var(new_var_desc.name()): |
| 6152 | old_var = self.blocks[idx].var(new_var_desc.name()) |
| 6153 | else: |
| 6154 | old_var = None |
| 6155 | |
| 6156 | kwargs = { |
| 6157 | "type": new_var_desc.type(), |
| 6158 | "name": new_var_desc.name(), |
| 6159 | "shape": get_var_desc_attr_or_none( |
| 6160 | new_var_desc, |
| 6161 | "shape", |
| 6162 | [ |
| 6163 | core.VarDesc.VarType.DENSE_TENSOR, |
| 6164 | core.VarDesc.VarType.SELECTED_ROWS, |
| 6165 | core.VarDesc.VarType.DENSE_TENSOR_ARRAY, |
| 6166 | ], |
| 6167 | ), |
| 6168 | "dtype": get_var_desc_attr_or_none( |
| 6169 | new_var_desc, |
| 6170 | "dtype", |
| 6171 | [ |
| 6172 | core.VarDesc.VarType.DENSE_TENSOR, |
| 6173 | core.VarDesc.VarType.SELECTED_ROWS, |
| 6174 | core.VarDesc.VarType.DENSE_TENSOR_ARRAY, |
| 6175 | ], |
| 6176 | ), |
| 6177 | "lod_level": get_var_desc_attr_or_none( |
| 6178 | new_var_desc, |
| 6179 | "lod_level", |
| 6180 | [ |
| 6181 | core.VarDesc.VarType.DENSE_TENSOR, |
| 6182 | core.VarDesc.VarType.DENSE_TENSOR_ARRAY, |
| 6183 | ], |
| 6184 | ), |
| 6185 | "error_clip": ( |
| 6186 | old_var.error_clip if old_var is not None else None |
| 6187 | ), |
| 6188 | "stop_gradient": ( |
| 6189 | old_var.stop_gradient if old_var is not None else False |
no test coverage detected