TensorPlay
latest (dev)
Copy
View Markdown

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 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:

  1. 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.

  2. 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:
  • optimizer (tensorplay.optim.Optimizer) – Optimizer that applies the gradients.

  • args – Any arguments.

  • kwargs – Any keyword arguments.

Warning

Closure use is not currently supported.

On this page

Ask DeepWiki