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
12 changes: 12 additions & 0 deletions colossalai/zero/gemini/gemini_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -116,6 +116,18 @@ def __init__(
verbose: bool = False,
**defaults: Any,
):
if type(optim) is FusedAdam and any(device.type == "cpu" for device in module.grads_device.values()):
# FusedAdam only has a CUDA update kernel. Static placement can put
# master parameters, gradients, and optimizer states on CPU, so use
# the hybrid implementation when at least one optimizer shard is
# actually offloaded. Existing parameter groups preserve all user
# hyperparameters and per-group overrides.
get_dist_logger().warning(
"FusedAdam does not support CPU optimizer shards; switching to HybridAdam for Gemini offload.",
ranks=[0],
)
optim = HybridAdam(optim.param_groups, adamw_mode=bool(optim.adamw_mode))

super().__init__(optim)
assert isinstance(module, GeminiDDP)
assert type(optim) in _AVAIL_OPTIM_LIST, (
Expand Down
52 changes: 52 additions & 0 deletions tests/test_zero/test_gemini/test_fused_adam_offload.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
import pytest
import torch
import torch.nn as nn

import colossalai
from colossalai.booster import Booster
from colossalai.booster.plugin import GeminiPlugin
from colossalai.nn.optimizer import FusedAdam, HybridAdam
from colossalai.testing import rerun_if_address_is_in_use, spawn


class MLP(nn.Module):
def __init__(self, input_dim=1024, hidden_dim=512, num_layers=10, num_classes=10):
super().__init__()
layers = []
for index in range(num_layers):
in_dim = input_dim if index == 0 else hidden_dim
layers.extend((nn.Linear(in_dim, hidden_dim), nn.ReLU()))
layers.append(nn.Linear(hidden_dim, num_classes))
self.net = nn.Sequential(*layers)

def forward(self, x):
return self.net(x)


def run_fused_adam_offload(rank, world_size, port):
colossalai.launch(rank=rank, world_size=world_size, host="localhost", port=port, backend="nccl")
torch.manual_seed(1024)
model = MLP()
optimizer = FusedAdam(model.parameters(), lr=1e-3)
plugin = GeminiPlugin(offload_optim_frac=1.0, min_chunk_size_m=2)
booster = Booster(plugin=plugin)
model, optimizer, _, _, _ = booster.boost(model, optimizer)

assert any(device.type == "cpu" for device in model.grads_device.values())
assert type(optimizer.optim) is HybridAdam

for _ in range(2):
optimizer.zero_grad()
output = model(torch.randn(4, 1024, device="cuda", dtype=torch.float16))
booster.backward(output.float().square().mean(), optimizer)
optimizer.step()


@pytest.mark.dist
@rerun_if_address_is_in_use()
def test_fused_adam_offload():
spawn(run_fused_adam_offload, 2)


if __name__ == "__main__":
test_fused_adam_offload()