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

Method __init__

python/paddle/distributed/auto_parallel/api.py:2969–3054  ·  view source on GitHub ↗
(
        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,
    )

Source from the content-addressed store, hash-verified

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 (

Callers

nothing calls this directly

Calls 15

__convert_strategyMethod · 0.95
trainMethod · 0.95
evalMethod · 0.95
predictMethod · 0.95
use_pir_apiFunction · 0.90
EngineClass · 0.90
get_loggerMethod · 0.80
_prepare_data_specMethod · 0.80
get_flagsMethod · 0.80
itemsMethod · 0.45

Tested by

no test coverage detected