latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
tensorplay.distributed.algorithms.ddp_comm_hooks.post_localSGD_hook API
Functions 1
post_localSGD_hook
functionFull reference ↗- tensorplay.distributed.algorithms.ddp_comm_hooks.post_localSGD_hook.post_localSGD_hook(state: PostLocalSGDState, bucket: GradBucket)[source]
Run post-localSGD algorithm.
This DDP communication hook is used for running post-localSGD algorithm, by combining with a model averaging component (e.g.,
PeriodicModelAverager) that runs after the optimizer step.- Parameters:
state (PostLocalSGDState) – State information to run post-localSGD. Users mainly need to tune
start_localSGD_iterto determine when to start local SGD.bucket (dist.GradBucket) – Bucket that stores a 1D flattened gradient tensor that batches multiple per-variable tensors. Note that since DDP comm hook only supports single process single device mode, only exactly one tensor is stored in this bucket.
- Returns:
Future handler of the communication, which updates the gradients in place.
- Example::
>>> # xdoctest: +SKIP >>> state = PostLocalSGDState(process_group=process_group, subgroup=subgroup, start_localSGD_iter=10) >>> ddp_model.register_comm_hook(state, post_localSGD_hook) >>> # Also need to establish a model averaging module and run model averaging after ``optimizer.step()``.
Classes 1
PostLocalSGDState
classFull reference ↗- class tensorplay.distributed.algorithms.ddp_comm_hooks.post_localSGD_hook.PostLocalSGDState(process_group, subgroup, start_localSGD_iter, post_local_gradient_allreduce=True)[source]
Store state for all-reducing gradients globally until given step, then locally after.
Stores the state for all-reducing gradients globally using
process_groupuntil stepstart_localSGD_iter, and all-reducing gradients locally usingsubgroupafterwards.If
process_groupisNone, the global process group will be used. IfsubgroupisNone, the intra-node process group on each machine will be used.Additionally,
post_local_gradient_allreducemay be worth tuning, because both true and false may give a faster convergence.- maybe_increase_iter(bucket)[source]
Track iterations and trigger log message at start of local SGD.
Help improve this page
Found an error, an unclear step, or a missing example?
tensorplay.distributed.algorithms.ddp_comm_hooks.default_hooks API
Complete API reference for tensorplay.distributed.algorithms.ddp_comm_hooks.default_hooks, including signatures, parameters, examples and members.
tensorplay.distributed.algorithms.ddp_comm_hooks.powerSGD_hook API
Complete API reference for tensorplay.distributed.algorithms.ddp_comm_hooks.powerSGD_hook, including signatures, parameters, examples and members.

