TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

tensorplay.distributed.algorithms.ddp_comm_hooks.post_localSGD_hook.post_localSGD_hook

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_iter to 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()``.

On this page

Ask DeepWiki