TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

PostLocalSGDOptimizer

class tensorplay.distributed.optim.PostLocalSGDOptimizer(optim: Optimizer, averager: ModelAverager)[source]

Wraps an arbitrary tensorplay.optim.Optimizer and runs post-local SGD, This optimizer runs local optimizer at every step. After the warm-up stage, it averages parameters periodically after the local optimizer is applied.

Parameters:
  • optim – The local optimizer.

  • averager – A model averager instance to run post-localSGD algorithm.

Example:

>>> # xdoctest: +SKIP("undefined variables")
>>> local_optim = tp.optim.SGD(params=model.parameters(), lr=0.01)
>>> opt = PostLocalSGDOptimizer(
>>>     optim=local_optim,
>>>     averager=averagers.PeriodicModelAverager(period=4, warmup_steps=100)
>>> )
>>> for step in range(0, 200):
>>>    opt.zero_grad()
>>>    loss = loss_fn(output, labels)
>>>    loss.backward()
>>>    opt.step()
load_state_dict(state_dict)[source]

This is the same as tensorplay.optim.Optimizer load_state_dict(), but also restores model averager’s step value to the one saved in the provided state_dict.

If there is no "step" entry in state_dict, it will raise a warning and initialize the model averager’s step to 0.

state_dict()[source]

This is the same as tensorplay.optim.Optimizer state_dict(), but adds an extra entry to record model averager’s step to the checkpoint to ensure reload does not cause unnecessary warm up again.

step()[source]

Performs a single optimization step (parameter update).

On this page

Ask DeepWiki