latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
ShardedGradScaler
- class tensorplay.distributed.fsdp.sharded_grad_scaler.ShardedGradScaler(device: str = 'cuda', init_scale: float = 65536.0, backoff_factor: float = 0.5, growth_factor: float = 2.0, growth_interval: int = 2000, enabled: bool = True, process_group: Any = None)[source]
- get_backoff_factor() → float
Return a Python float containing the scale backoff factor.
- get_growth_factor() → float
Return a Python float containing the scale growth factor.
- get_growth_interval() → int
Return a Python int containing the growth interval.
- get_scale() → float
Return a Python float containing the current scale, or 1.0 if scaling is disabled.
- is_enabled() → bool
Return a bool indicating whether this instance is enabled.
- load_state_dict(state_dict: dict[str, Any]) → None
Load the scaler state.
If this instance is disabled,
load_state_dict()is a no-op.- Parameters:
state_dict (dict) – scaler state. Should be an object returned from a call to
state_dict().
- set_backoff_factor(new_factor: float) → None
Set a new scale backoff factor.
- Parameters:
new_scale (float) – Value to use as the new scale backoff factor.
- set_growth_factor(new_factor: float) → None
Set a new scale growth factor.
- Parameters:
new_scale (float) – Value to use as the new scale growth factor.
- set_growth_interval(new_interval: int) → None
Set a new growth interval.
- Parameters:
new_interval (int) – Value to use as the new growth interval.
- state_dict() → dict[str, Any]
Return the state of the scaler as a
dict.It contains five entries:
"scale"- a Python float containing the current scale"growth_factor"- a Python float containing the current growth factor"backoff_factor"- a Python float containing the current backoff factor"growth_interval"- a Python int containing the current growth interval"_growth_tracker"- a Python int containing the number of recent consecutive unskipped steps.
If this instance is not enabled, returns an empty dict.
Note
If you wish to checkpoint the scaler’s state after a particular iteration,
state_dict()should be called afterupdate().
- step(optimizer: Optimizer, *args: Any, **kwargs: Any) → Any
Invoke
unscale_(optimizer)followed by parameter update, if gradients are not infs/NaN.step()carries out the following two operations:Internally invokes
unscale_(optimizer)(unlessunscale_()was explicitly called foroptimizerearlier in the iteration). As part of theunscale_(), gradients are checked for infs/NaNs.If no inf/NaN gradients are found, invokes
optimizer.step()using the unscaled gradients. Otherwise,optimizer.step()is skipped to avoid corrupting the params.
*argsand**kwargsare forwarded tooptimizer.step().Returns the return value of
optimizer.step(*args, **kwargs).- Parameters:
optimizer (tensorplay.optim.Optimizer) – Optimizer that applies the gradients.
args – Any arguments.
kwargs – Any keyword arguments.
Warning
Closure use is not currently supported.
Help improve this page
Found an error, an unclear step, or a missing example?

