From cc6d6337ffdcd7ade88b735c39a33ba989c1ad0a Mon Sep 17 00:00:00 2001 From: Chenghao Liu Date: Tue, 15 Sep 2026 20:07:45 +0800 Subject: [PATCH] fix(train): initialize GaLore through split optimizer hooks Signed-off-by: Chenghao Liu --- swift/optimizers/galore/utils.py | 59 +++++++++---------- tests/general/test_galore_schedule.py | 83 +++++++++++++++++++++++++++ 2 files changed, 110 insertions(+), 32 deletions(-) create mode 100644 tests/general/test_galore_schedule.py diff --git a/swift/optimizers/galore/utils.py b/swift/optimizers/galore/utils.py index 330b4f957a..163bf855b9 100644 --- a/swift/optimizers/galore/utils.py +++ b/swift/optimizers/galore/utils.py @@ -8,7 +8,6 @@ from transformers import get_scheduler from typing import TYPE_CHECKING, Any, Dict, List, Tuple, Union -from swift.trainers import calculate_max_steps from swift.utils import get_logger from ..base import OptimizerCallback @@ -82,8 +81,7 @@ def step(self, *args, **kwargs) -> None: self._last_lr = lr_scheduler.get_last_lr() -def _create_optimizer_and_scheduler(model: nn.Module, args: 'TrainingArguments', config: GaLoreConfig, max_steps, - **defaults): +def _create_optimizer(model: nn.Module, args: 'TrainingArguments', config: GaLoreConfig, **defaults): galore_params = [] for module_name, module in model.named_modules(): if not isinstance(module, (nn.Linear, nn.Embedding)) or \ @@ -124,19 +122,7 @@ def _create_optimizer_and_scheduler(model: nn.Module, args: 'TrainingArguments', else: optimizer_dict[p] = optim_cls([{'params': [p], **defaults}], **optim_kwargs) - # get scheduler dict - scheduler_dict = {} - for p in model.parameters(): - if p.requires_grad: - scheduler_dict[p] = get_scheduler( - optimizer=optimizer_dict[p], - name=args.lr_scheduler_type, - num_training_steps=max_steps * 2, - num_warmup_steps=args.warmup_steps * 2, - scheduler_specific_kwargs=args.lr_scheduler_kwargs, - ) - - return GaloreOptimizerWrapper(optimizer_dict), GaloreSchedulerWrapper(scheduler_dict) + return GaloreOptimizerWrapper(optimizer_dict) else: decay_parameters = HfTrainer.get_decay_parameter_names(None, model) param_groups = [{ @@ -161,15 +147,7 @@ def _create_optimizer_and_scheduler(model: nn.Module, args: 'TrainingArguments', 0.0, }, ]) - optim = optim_cls(param_groups, **optim_kwargs) - scheduler = get_scheduler( - optimizer=optim, - name=args.lr_scheduler_type, - num_training_steps=max_steps, - num_warmup_steps=args.warmup_steps, - scheduler_specific_kwargs=args.lr_scheduler_kwargs, - ) - return optim, scheduler + return optim_cls(param_groups, **optim_kwargs) def get_optimizer(args: 'TrainingArguments', config: GaLoreConfig) -> Tuple[Any, Any]: @@ -221,10 +199,10 @@ def get_optimizer(args: 'TrainingArguments', config: GaLoreConfig) -> Tuple[Any, class GaloreOptimizerCallback(OptimizerCallback): - def create_optimizer_and_scheduler(self, num_training_steps: int): - trainer = self.trainer + def create_optimizer(self, model=None): args = self.args - training_steps = calculate_max_steps(args, trainer.train_dataset) + if model is None: + model = self.trainer.model galore_config = GaLoreConfig( target_modules=args.galore_target_modules, rank=args.galore_rank, @@ -240,7 +218,24 @@ def create_optimizer_and_scheduler(self, num_training_steps: int): gamma_proj=args.galore_gamma_proj, queue_size=args.galore_queue_size, ) - optimizer, lr_scheduler = _create_optimizer_and_scheduler( - trainer.model, args, galore_config, training_steps, lr=args.learning_rate, weight_decay=args.weight_decay) - trainer.optimizer = optimizer - trainer.lr_scheduler = lr_scheduler + return _create_optimizer(model, args, galore_config, lr=args.learning_rate, weight_decay=args.weight_decay) + + def create_scheduler(self, num_training_steps: int, optimizer): + args = self.args + + def make_scheduler(optimizer): + return get_scheduler( + optimizer=optimizer, + name=args.lr_scheduler_type, + num_training_steps=num_training_steps, + num_warmup_steps=args.get_warmup_steps(num_training_steps), + scheduler_specific_kwargs=args.lr_scheduler_kwargs, + ) + + if isinstance(optimizer, GaloreOptimizerWrapper): + return GaloreSchedulerWrapper({p: make_scheduler(opt) for p, opt in optimizer.optimizers.items()}) + return make_scheduler(optimizer) + + def create_optimizer_and_scheduler(self, num_training_steps: int): + self.trainer.optimizer = self.create_optimizer() + self.trainer.lr_scheduler = self.create_scheduler(num_training_steps, self.trainer.optimizer) diff --git a/tests/general/test_galore_schedule.py b/tests/general/test_galore_schedule.py new file mode 100644 index 0000000000..50001b2cea --- /dev/null +++ b/tests/general/test_galore_schedule.py @@ -0,0 +1,83 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +import tempfile +import torch +import unittest +from torch import nn +from transformers import Trainer +from types import SimpleNamespace + +from swift.optimizers.galore.utils import GaloreOptimizerCallback +from swift.trainers import TrainingArguments + + +class TestGaloreSchedule(unittest.TestCase): + + def test_split_hooks_create_galore_optimizer(self): + from swift.optimizers.galore import GaLoreAdamW + with tempfile.TemporaryDirectory() as output_dir: + args = TrainingArguments( + output_dir=output_dir, + use_cpu=True, + report_to=[], + optim='adamw_torch', + galore_target_modules=['0'], + galore_rank=2) + trainer = Trainer(args=args, model=nn.Sequential(nn.Linear(4, 4))) + callback = GaloreOptimizerCallback(args, trainer) + optimizer = callback.create_optimizer() + self.assertIsInstance(optimizer, GaLoreAdamW) + self.assertTrue(any('rank' in group for group in optimizer.param_groups)) + scheduler = callback.create_scheduler(5, optimizer) + self.assertIs(scheduler.optimizer, optimizer) + + def test_learning_rate_follows_training_steps_and_warmup(self): + cases = ((5, 0.4, 0), (5, 0.0, 2), (5, 0.0, 0), (4, 0.0, 0)) + for per_parameter in (False, True): + for total_steps, warmup_ratio, warmup_steps in cases: + with self.subTest( + per_parameter=per_parameter, + total_steps=total_steps, + warmup_ratio=warmup_ratio, + warmup_steps=warmup_steps): + with tempfile.TemporaryDirectory() as output_dir: + args = TrainingArguments( + output_dir=output_dir, + use_cpu=True, + report_to=[], + optim='adamw_torch', + learning_rate=0.01, + weight_decay=0.0, + lr_scheduler_type='linear', + warmup_ratio=warmup_ratio, + warmup_steps=warmup_steps, + per_device_train_batch_size=2, + gradient_accumulation_steps=2, + num_train_epochs=1, + galore_target_modules=['0'], + galore_rank=2, + galore_optim_per_parameter=per_parameter) + model = nn.Sequential(nn.Linear(4, 4)) + trainer = SimpleNamespace(args=args, model=model, train_dataset=list(range(18))) + GaloreOptimizerCallback(args, trainer).create_optimizer_and_scheduler(total_steps) + optimizers = list( + trainer.optimizer.optimizers.values()) if per_parameter else [trainer.optimizer] + resolved_warmup = args.get_warmup_steps(total_steps) + learning_rates = [] + for step in range(total_steps + 1): + rates = [group['lr'] for optimizer in optimizers for group in optimizer.param_groups] + expected_factor = ( + step / max(1, resolved_warmup) if step < resolved_warmup else max( + 0.0, (total_steps - step) / max(1, total_steps - resolved_warmup))) + learning_rates.append((rates, args.learning_rate * expected_factor)) + if step < total_steps: + model(torch.ones(2, 4)).square().mean().backward() + trainer.optimizer.step() + trainer.lr_scheduler.step() + trainer.optimizer.zero_grad() + for rates, expected in learning_rates: + for actual in rates: + self.assertAlmostEqual(actual, expected, places=10) + + +if __name__ == '__main__': + unittest.main()