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

Method _build

python/paddle/distributed/auto_parallel/static/engine.py:1092–1243  ·  view source on GitHub ↗
(self, mode)

Source from the content-addressed store, hash-verified

1090 _create_dist_input_var(input_var, input_spec)
1091
1092 def _build(self, mode):
1093 if in_dynamic_mode() or self._dygraph_mode:
1094 paddle.disable_static()
1095 self._dygraph_mode = True
1096 self._logger.info("Building model with 'to_static' method.")
1097
1098 self.program_helper = ProgramHelper(
1099 self._model,
1100 self._loss,
1101 self._metrics,
1102 self._inputs_spec,
1103 self._labels_spec,
1104 )
1105 # build forward main program
1106 with utils.unique_name.guard():
1107 self.program_helper.build_program(mode)
1108
1109 self.concrete_program = self.program_helper.concrete_program
1110 serial_main_prog = self.program_helper.main_program
1111 serial_startup_prog = self.program_helper.startup_program
1112
1113 self._inputs = self.program_helper.input_vars
1114 self._labels = self.program_helper.label_vars
1115 # self._process_dist_input_specs()
1116 outputs = self.program_helper.output_vars
1117 self._losses = self.program_helper.loss_vars
1118 self._loss_names = self.program_helper.loss_names
1119 metrics = self.program_helper.metric_vars
1120
1121 paddle.enable_static()
1122 else:
1123 # build program in static mode
1124 dist_context = self._dist_contexts.get(mode, None)
1125 if dist_context is not None:
1126 return
1127
1128 outputs = []
1129 metrics = []
1130 self._losses = []
1131 serial_main_prog = self._orig_main_prog.clone()
1132 serial_startup_prog = self._orig_startup_prog.clone()
1133 if not self._skip_build:
1134 with (
1135 static.program_guard(serial_main_prog, serial_startup_prog),
1136 utils.unique_name.guard(),
1137 ):
1138 self._inputs = [
1139 s._create_feed_layer() for s in self._inputs_spec
1140 ]
1141 self._labels = [
1142 s._create_feed_layer() for s in self._labels_spec
1143 ]
1144
1145 outputs = auto_utils.to_list(self._model(*self._inputs))
1146
1147 if mode != "predict" and self._loss:
1148 assert isinstance(
1149 self._loss, paddle.nn.Layer

Callers 6

_prepare_programMethod · 0.95
_optimization_tuningMethod · 0.95
costMethod · 0.95

Calls 15

ProgramHelperClass · 0.85
new_process_groupFunction · 0.85
listFunction · 0.85
rangeFunction · 0.85
DistributedContextClass · 0.85
_create_feed_layerMethod · 0.80
infoMethod · 0.45
build_programMethod · 0.45
getMethod · 0.45
cloneMethod · 0.45
to_listMethod · 0.45

Tested by 3