Optimizer for distributed training. For the distributed training, this method would rebuild a new instance of DistributedOptimizer. Which has basic Optimizer function and special features for distributed training. Args: optimizer(Optimizer): The executo
(
self,
optimizer: Optimizer,
strategy: DistributedStrategy | None = None,
)
| 1475 | self._runtime_handle._shrink(threshold) |
| 1476 | |
| 1477 | def distributed_optimizer( |
| 1478 | self, |
| 1479 | optimizer: Optimizer, |
| 1480 | strategy: DistributedStrategy | None = None, |
| 1481 | ) -> Self: |
| 1482 | """ |
| 1483 | Optimizer for distributed training. |
| 1484 | |
| 1485 | For the distributed training, this method would rebuild a new instance of DistributedOptimizer. |
| 1486 | Which has basic Optimizer function and special features for distributed training. |
| 1487 | |
| 1488 | Args: |
| 1489 | optimizer(Optimizer): The executor to run for init server. |
| 1490 | strategy(DistributedStrategy): Extra properties for distributed optimizer. |
| 1491 | It is recommended to use DistributedStrategy in fleet.init(). The strategy |
| 1492 | here is for compatibility. If the strategy in fleet.distributed_optimizer() |
| 1493 | is not None, then it will overwrite the DistributedStrategy in fleet.init(), |
| 1494 | which will take effect in distributed training. |
| 1495 | |
| 1496 | Returns: |
| 1497 | Fleet: instance of fleet. |
| 1498 | |
| 1499 | Examples: |
| 1500 | |
| 1501 | .. code-block:: pycon |
| 1502 | |
| 1503 | >>> import paddle |
| 1504 | >>> import paddle.distributed.fleet as fleet |
| 1505 | >>> fleet.init(is_collective=True) |
| 1506 | >>> linear = paddle.nn.Linear(10, 10) |
| 1507 | >>> strategy = fleet.DistributedStrategy() |
| 1508 | >>> optimizer = paddle.optimizer.SGD(learning_rate=0.001, parameters=linear.parameters()) |
| 1509 | >>> optimizer = fleet.distributed_optimizer(optimizer, strategy=strategy) |
| 1510 | |
| 1511 | """ |
| 1512 | self.user_defined_optimizer = optimizer |
| 1513 | |
| 1514 | if strategy is not None: |
| 1515 | if self._is_collective: |
| 1516 | logger.warning( |
| 1517 | "It is recommended to use DistributedStrategy " |
| 1518 | "in fleet.init(). The strategy here is only for compatibility. " |
| 1519 | "If the strategy in fleet.distributed_optimizer() is " |
| 1520 | "not None, then it will overwrite the DistributedStrategy in fleet.init(), " |
| 1521 | "which will take effect in distributed training." |
| 1522 | ) |
| 1523 | self._user_defined_strategy = copy.deepcopy(strategy) |
| 1524 | |
| 1525 | self._context = {} |
| 1526 | |
| 1527 | return self |
| 1528 | |
| 1529 | def _get_amp_optimizer(self): |
| 1530 | # imitate target optimizer retrieval |