Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
59 changes: 27 additions & 32 deletions swift/optimizers/galore/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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 \
Expand Down Expand Up @@ -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 = [{
Expand All @@ -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]:
Expand Down Expand Up @@ -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,
Expand All @@ -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)
83 changes: 83 additions & 0 deletions tests/general/test_galore_schedule.py
Original file line number Diff line number Diff line change
@@ -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()
Loading