API reference
Documentation / API reference latest (dev) Latest development documentation · Updated 2026-10-08
tensorplay.distributed.fsdp.sharded_grad_scaler API Classes 1
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 after update() .
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) (unless unscale_() was explicitly called for optimizer
earlier in the iteration). As part of the unscale_() , 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.
*args and **kwargs are forwarded to optimizer.step() .
Returns the return value of optimizer.step(*args, **kwargs) .
Parameters:
Warning
Closure use is not currently supported.
Help improve this page Found an error, an unclear step, or a missing example?