diff --git a/distributed_shampoo/distributed_shampoo.py b/distributed_shampoo/distributed_shampoo.py index 0bbbc84..e28e8e2 100644 --- a/distributed_shampoo/distributed_shampoo.py +++ b/distributed_shampoo/distributed_shampoo.py @@ -754,6 +754,7 @@ def _preconditioner_config_to_list_cls( return SpectralDescentPreconditionerList( block_list=state_lists[DISTRIBUTOR].local_blocked_params, preconditioner_config=preconditioner_config, + shampoo_pt2_compile_config=self._shampoo_pt2_compile_config, ) case _: raise NotImplementedError(f"{preconditioner_config=} not supported!") diff --git a/distributed_shampoo/preconditioner/spectral_descent_preconditioner_list.py b/distributed_shampoo/preconditioner/spectral_descent_preconditioner_list.py index 58b7f7c..2510f46 100644 --- a/distributed_shampoo/preconditioner/spectral_descent_preconditioner_list.py +++ b/distributed_shampoo/preconditioner/spectral_descent_preconditioner_list.py @@ -7,15 +7,60 @@ """ +import torch +from collections.abc import Callable +from dataclasses import asdict from distributed_shampoo.preconditioner.matrix_functions import matrix_orthogonalization +from distributed_shampoo.preconditioner.matrix_functions_types import ( + NewtonSchulzOrthogonalizationConfig, +) from distributed_shampoo.preconditioner.preconditioner_list import ( PreconditionerList, profile_decorator, ) -from distributed_shampoo.shampoo_types import SpectralDescentPreconditionerConfig +from distributed_shampoo.shampoo_types import ( + ShampooPT2CompileConfig, + SpectralDescentPreconditionerConfig, +) +from torch._higher_order_ops import foreach_map from torch import Tensor +def _newton_schulz( + A: Tensor, + a: float, + b: float, + c: float, + num_iterations: int, +) -> Tensor: + transpose = A.shape[0] > A.shape[1] + X = A.T if transpose else A + X = X / X.norm().clamp(min=1e-8) + for _ in range(num_iterations): + gram = X @ X.T + gram_update = torch.addmm(gram, gram, gram, beta=b, alpha=c) + X = torch.addmm(X, gram_update, X, beta=a) + return X.T if transpose else X + + +def _foreach_newton_schulz( + grads: tuple[Tensor, ...], + coefficients: tuple[float, float, float], + num_iterations: int, + scales: tuple[float, ...], +) -> tuple[Tensor, ...]: + a, b, c = coefficients + orthogonalized = foreach_map( + _newton_schulz, + grads, + a, + b, + c, + num_iterations, + ) + return tuple(result.mul(scale) for result, scale in zip(orthogonalized, scales)) + + class SpectralDescentPreconditionerList(PreconditionerList): """Preconditioner list for spectral descent. @@ -33,6 +78,7 @@ def __init__( self, block_list: tuple[Tensor, ...], preconditioner_config: SpectralDescentPreconditionerConfig, + shampoo_pt2_compile_config: ShampooPT2CompileConfig | None = None, ) -> None: if any(block.dim() != 2 for block in block_list): raise ValueError( @@ -41,6 +87,14 @@ def __init__( ) super().__init__(block_list) self._preconditioner_config = preconditioner_config + self._foreach_newton_schulz: Callable[..., tuple[Tensor, ...]] = ( + torch.compile( + _foreach_newton_schulz, + **asdict(shampoo_pt2_compile_config), + ) + if shampoo_pt2_compile_config is not None + else _foreach_newton_schulz + ) @profile_decorator def update_preconditioners( @@ -53,6 +107,21 @@ def update_preconditioners( @profile_decorator def precondition(self, masked_grad_list: tuple[Tensor, ...]) -> tuple[Tensor, ...]: + config = self._preconditioner_config.orthogonalization_config + if ( + masked_grad_list + and isinstance(config, NewtonSchulzOrthogonalizationConfig) + and all(grad.dtype is torch.bfloat16 for grad in masked_grad_list) + ): + return self._foreach_newton_schulz( + masked_grad_list, + config.coefficients, + config.num_iterations, + tuple( + config.scale_by_dims_fn(grad.shape[1], grad.shape[0]) + for grad in masked_grad_list + ), + ) return tuple( # An error will be raised when grad is not 2D. matrix_orthogonalization( diff --git a/distributed_shampoo/preconditioner/tests/spectral_descent_preconditioner_list_test.py b/distributed_shampoo/preconditioner/tests/spectral_descent_preconditioner_list_test.py index fb0f4bc..b8590c6 100644 --- a/distributed_shampoo/preconditioner/tests/spectral_descent_preconditioner_list_test.py +++ b/distributed_shampoo/preconditioner/tests/spectral_descent_preconditioner_list_test.py @@ -8,9 +8,11 @@ """ import re +from unittest import mock from typing import Any import torch +from distributed_shampoo.preconditioner.matrix_functions import matrix_orthogonalization from distributed_shampoo.preconditioner.matrix_functions_types import ( DefaultNewtonSchulzOrthogonalizationConfig, OrthogonalizationConfig, @@ -80,6 +82,32 @@ def test_precondition_non_square_matrix( ) preconditioner_list.precondition(masked_grad_list=masked_grad_list) + def test_precondition_bfloat16_uses_foreach_map(self) -> None: + block_list = ( + torch.randn(3, 2, dtype=torch.bfloat16), + torch.randn(2, 3, dtype=torch.bfloat16), + ) + preconditioner_list = SpectralDescentPreconditionerList( + block_list=block_list, + preconditioner_config=DefaultSpectralDescentPreconditionerConfig, + ) + expected = tuple(matrix_orthogonalization(block) for block in block_list) + with mock.patch( + "distributed_shampoo.preconditioner.spectral_descent_preconditioner_list.foreach_map", + wraps=torch._higher_order_ops.foreach_map, + ) as foreach_map_mock: + actual = preconditioner_list.precondition(masked_grad_list=block_list) + foreach_map_mock.assert_called_once() + for actual_block, expected_block in zip(actual, expected): + torch.testing.assert_close(actual_block, expected_block) + + def test_precondition_empty_list(self) -> None: + preconditioner_list = SpectralDescentPreconditionerList( + block_list=(), + preconditioner_config=DefaultSpectralDescentPreconditionerConfig, + ) + self.assertEqual(preconditioner_list.precondition(masked_grad_list=()), ()) + @parametrize( "block_list", (