(
self,
layer: Layer,
loader: ShardDataloader | DataLoader,
loss: Layer | Callable[..., Any] | None = None,
optimizer: Optimizer | None = None,
strategy: Strategy | None = None,
metrics: list[Metric] | None = None,
input_spec: list[list[DistributedInputSpec]] | None = None,
)
| 2967 | """ |
| 2968 | |
| 2969 | def __init__( |
| 2970 | self, |
| 2971 | layer: Layer, |
| 2972 | loader: ShardDataloader | DataLoader, |
| 2973 | loss: Layer | Callable[..., Any] | None = None, |
| 2974 | optimizer: Optimizer | None = None, |
| 2975 | strategy: Strategy | None = None, |
| 2976 | metrics: list[Metric] | None = None, |
| 2977 | input_spec: list[list[DistributedInputSpec]] | None = None, |
| 2978 | ) -> None: |
| 2979 | self._inner_strategy = self.__convert_strategy(strategy) |
| 2980 | self._structured_to_parameter_name = { |
| 2981 | k: v.name for k, v in layer.state_dict().items() |
| 2982 | } |
| 2983 | self._parameter_to_structured_name = { |
| 2984 | v: k for k, v in self._structured_to_parameter_name.items() |
| 2985 | } |
| 2986 | if os.getenv("POD_NAME"): |
| 2987 | dist.utils.log_utils.get_logger(logging.INFO).info( |
| 2988 | "Distribute training by paddle.distributed.launch" |
| 2989 | ) |
| 2990 | dist.fleet.init(is_collective=True) |
| 2991 | |
| 2992 | if ( |
| 2993 | strategy |
| 2994 | and strategy.sharding.enable_tensor_fusion |
| 2995 | and isinstance(optimizer, _ShardOptimizer) |
| 2996 | and hasattr(optimizer, '_shard_fn') |
| 2997 | and hasattr(optimizer, '_inner_opt') |
| 2998 | and use_pir_api() |
| 2999 | ): |
| 3000 | assert isinstance(optimizer._shard_fn, ShardingStage1), ( |
| 3001 | "The shard_fn should be ShardingStage1 " |
| 3002 | "when stage1 tensor fusion is enabled." |
| 3003 | ) |
| 3004 | if isinstance(optimizer._shard_fn, ShardingStage1): |
| 3005 | shard_fn = optimizer._shard_fn |
| 3006 | inner_opt = optimizer._inner_opt |
| 3007 | optimizer = ShardingOptimizerStage1( |
| 3008 | inner_opt, shard_fn, self._inner_strategy |
| 3009 | ) |
| 3010 | else: |
| 3011 | logging.warning( |
| 3012 | "Sharding tensor fusion only support ShardingStage1 now." |
| 3013 | ) |
| 3014 | |
| 3015 | self._engine = Engine( |
| 3016 | layer, loss, optimizer, metrics, strategy=self._inner_strategy |
| 3017 | ) |
| 3018 | self._mode = None |
| 3019 | self._feed_name_list = {} |
| 3020 | |
| 3021 | # convert dygraph model to static model |
| 3022 | if input_spec is not None: |
| 3023 | self._engine._inputs_spec = input_spec[0] |
| 3024 | self._engine._labels_spec = input_spec[1] |
| 3025 | elif isinstance(loader, ShardDataloader): |
| 3026 | ( |
nothing calls this directly
no test coverage detected