Skip to content
Merged
Show file tree
Hide file tree
Changes from 42 commits
Commits
Show all changes
47 commits
Select commit Hold shift + click to select a range
f968eef
Split parameter offload from z3
tjruwase Jun 10, 2022
984deaf
Format fixes
tjruwase Jun 10, 2022
1219ed8
Bug fixes
tjruwase Jun 10, 2022
37a2405
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jun 13, 2022
16088e6
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jun 14, 2022
e3e346d
Cleanup
tjruwase Jun 14, 2022
e83936e
Merge branch 'olruwase/parameter_offload' of github.com:microsoft/Dee…
tjruwase Jun 14, 2022
e122028
Remove dead code
tjruwase Jun 14, 2022
2a1b245
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jun 15, 2022
5573966
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jun 16, 2022
e86c41f
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jun 16, 2022
60111c5
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jun 16, 2022
92b90f3
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jun 20, 2022
afb6f6c
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jun 20, 2022
ecc34b1
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jun 21, 2022
8993ec4
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jun 21, 2022
59dc8f1
Release swap buffers for persisted params
tjruwase Jul 11, 2022
a175754
Merge branch 'olruwase/parameter_offload' of github.com:microsoft/Dee…
tjruwase Jul 11, 2022
a266275
Merge with master
tjruwase Jul 12, 2022
c5ae3d6
Format fixes
tjruwase Jul 12, 2022
9954b8f
Format fixes
tjruwase Jul 12, 2022
c6cfb55
Pass args correctly
tjruwase Jul 13, 2022
31ccf6d
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jul 18, 2022
2f4ccc4
Use pinned memory for nvme offload
tjruwase Jul 19, 2022
8b4a70d
Merge branch 'olruwase/parameter_offload' of github.com:microsoft/Dee…
tjruwase Jul 19, 2022
e5b7ca1
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jul 19, 2022
d3bdbc4
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jul 20, 2022
6e86f8c
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jul 25, 2022
c2a8913
Merge with masster
tjruwase Jul 27, 2022
c3fb27a
Merge with masster
tjruwase Jul 27, 2022
80a7deb
Fix missing import
tjruwase Jul 27, 2022
e858891
model pesistence params
tjruwase Jul 27, 2022
263ad48
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jul 27, 2022
ff682c9
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jul 27, 2022
7de4e46
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jul 27, 2022
88462a3
Merge branch 'master' of github.com:microsoft/DeepSpeed into olruwase…
tjruwase Jul 28, 2022
b53e913
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jul 28, 2022
fcbef1f
Merge branch 'olruwase/parameter_offload' of github.com:microsoft/Dee…
tjruwase Jul 28, 2022
0ec1f8f
Fix merge issues
tjruwase Jul 28, 2022
5db1b8c
Handle none device
tjruwase Jul 28, 2022
428964a
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jul 28, 2022
49c1f78
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jul 28, 2022
0309b01
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jul 29, 2022
bc205d8
Usse log_dist
tjruwase Jul 29, 2022
126ab13
Merge branch 'olruwase/parameter_offload' of github.com:microsoft/Dee…
tjruwase Jul 29, 2022
43b8f97
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jul 29, 2022
45a8e9e
Merge branch 'master' into olruwase/parameter_offload
tjruwase Jul 30, 2022
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
16 changes: 11 additions & 5 deletions deepspeed/runtime/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -678,6 +678,9 @@ def zero_prefetch_bucket_size(self):
def zero_param_persistence_threshold(self):
return self._config.zero_config.param_persistence_threshold

def zero_model_persistence_threshold(self):
return self._config.zero_config.model_persistence_threshold

def zero_gather_16bit_weights_on_model_save(self):
return self._config.zero_config.gather_16bit_weights_on_model_save

Expand Down Expand Up @@ -1343,7 +1346,6 @@ def _configure_bf16_optimizer(self, optimizer):

def _configure_zero_optimizer(self, optimizer):
zero_stage = self.zero_optimization_stage()
log_dist('Creating fp16 ZeRO stage {} optimizer'.format(zero_stage), ranks=[0])
assert self.communication_data_type in (torch.float16, torch.bfloat16), "ZeRO supports only 'communication_data_type': ['fp16', 'bfp16']"
timers = self.timers if self.wall_clock_breakdown() else None

Expand All @@ -1361,6 +1363,8 @@ def _configure_zero_optimizer(self, optimizer):
round_robin_gradients = self.zero_round_robin_gradients()
assert not isinstance(optimizer, DummyOptim), "zero stage 2 requires an optimizer"

log_dist('Creating fp16 ZeRO stage {} optimizer'.format(zero_stage),
ranks=[0])
# Overlap and contiguous grads are meaningless in stage 1 and are ignored
if zero_stage == ZeroStageEnum.optimizer_states:
overlap_comm = False
Expand Down Expand Up @@ -1406,10 +1410,8 @@ def _configure_zero_optimizer(self, optimizer):

elif zero_stage == ZeroStageEnum.weights:
assert not self.has_moe_layers, "MoE not supported with Stage 3"
logger.info("Initializing ZeRO Stage 3") if dist.get_rank() == 0 else None
from deepspeed.runtime.zero.stage3 import DeepSpeedZeroOptimizer_Stage3

if isinstance(optimizer, DummyOptim):
logger.info("Creating ZeRO Offload") if dist.get_rank() == 0 else None
Comment thread
tjruwase marked this conversation as resolved.
Outdated
optimizer = DeepSpeedZeRoOffload(
self.module,
timers=timers,
Expand All @@ -1419,10 +1421,13 @@ def _configure_zero_optimizer(self, optimizer):
max_reuse_distance=self.zero_max_reuse_distance(),
max_live_parameters=self.zero_max_live_parameters(),
param_persistence_threshold=self.zero_param_persistence_threshold(),
model_persistence_threshold=self.zero_model_persistence_threshold(),
offload_param_config=self.zero_offload_param(),
mpu=self.mpu)
else:

log_dist('Creating fp16 ZeRO stage {} optimizer'.format(zero_stage),
ranks=[0])
from deepspeed.runtime.zero.stage3 import DeepSpeedZeroOptimizer_Stage3
optimizer = DeepSpeedZeroOptimizer_Stage3(
self.module,
optimizer,
Expand All @@ -1438,6 +1443,7 @@ def _configure_zero_optimizer(self, optimizer):
max_reuse_distance=self.zero_max_reuse_distance(),
max_live_parameters=self.zero_max_live_parameters(),
param_persistence_threshold=self.zero_param_persistence_threshold(),
model_persistence_threshold=self.zero_model_persistence_threshold(),
dp_process_group=self.data_parallel_group,
reduce_scatter=self.zero_reduce_scatter(),
overlap_comm=self.zero_overlap_comm(),
Expand Down
4 changes: 4 additions & 0 deletions deepspeed/runtime/zero/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
"""

from pydantic import Field, validator
import sys
from typing import Optional
from enum import Enum
from deepspeed.runtime.config_utils import get_scalar_param, DeepSpeedConfigModel
Expand Down Expand Up @@ -114,6 +115,9 @@ class DeepSpeedZeroConfig(DeepSpeedConfigModel):
param_persistence_threshold: int = Field(1e5,
ge=0,
alias="stage3_param_persistence_threshold")
model_persistence_threshold: int = Field(sys.maxsize,
ge=0,
alias="stage3_model_persistence_threshold")
max_live_parameters: int = Field(1e9, ge=0, alias="stage3_max_live_parameters")
max_reuse_distance: int = Field(1e9, ge=0, alias="stage3_max_reuse_distance")
gather_16bit_weights_on_model_save: bool = Field(
Expand Down
22 changes: 16 additions & 6 deletions deepspeed/runtime/zero/parameter_offload.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
Licensed under the MIT license.
"""

import sys
import torch
from torch.cuda import Stream
from collections import OrderedDict
Expand Down Expand Up @@ -173,10 +174,11 @@ def __init__(self,
max_reuse_distance=1000000000,
max_live_parameters=1000000000,
param_persistence_threshold=100000,
model_persistence_threshold=sys.maxsize,
offload_param_config=None,
mpu=None):

see_memory_usage("TensorOffload initialize beginning", force=True)
see_memory_usage("DeepSpeedZeRoOffload initialize [begin]", force=True)

print_rank_0(f"initialized {__class__.__name__} with args: {locals()}",
force=False)
Expand All @@ -196,8 +198,11 @@ def __init__(self,

_inject_parameters(module, ZeROOrderedDict)

self.persistence_threshold = int(param_persistence_threshold)
self.persistent_parameters = self.mark_persistent_parameters()
self.param_numel_persistence_threshold = int(param_persistence_threshold)
self.model_persistence_threshold = int(model_persistence_threshold)
self.persistent_parameters = self.mark_persistent_parameters(
self.param_numel_persistence_threshold,
self.model_persistence_threshold)

self.param_coordinators = {}
self._prefetch_bucket_sz = int(prefetch_bucket_size)
Expand All @@ -213,6 +218,8 @@ def __init__(self,
f'Created module hooks: forward = {len(self.forward_hooks)}, backward = {len(self.backward_hooks)}',
force=False)

see_memory_usage("DeepSpeedZeRoOffload initialize [end]", force=True)

@instrument_w_nvtx
def partition_all_parameters(self):
"""Partitioning Parameters that were not partitioned usually if parameters
Expand Down Expand Up @@ -291,20 +298,23 @@ def _end_of_forward_hook(module, *args):
global FWD_MODULE_STACK
FWD_MODULE_STACK.append(self.module)

def mark_persistent_parameters(self):
def mark_persistent_parameters(self, param_threshold, model_threshold):
persistent_params = []
total_persistent_parameters = 0
params_count = 0
for _, param in self.module.named_parameters(recurse=True):
if param.ds_numel < self.persistence_threshold:
if param.ds_numel + total_persistent_parameters > model_threshold:
continue

if param.ds_numel < param_threshold:
params_count += 1
param.ds_persist = True
persistent_params.append(param)
total_persistent_parameters += param.ds_numel

print_rank_0(
f"Parameter Offload: Total persistent parameters: {total_persistent_parameters} in {params_count} params",
force=False)
force=True)

return persistent_params

Expand Down
10 changes: 7 additions & 3 deletions deepspeed/runtime/zero/partition_parameters.py
Original file line number Diff line number Diff line change
Expand Up @@ -669,9 +669,13 @@ def get_model():

# Remote device is the device where parameter partitions are stored
# It can be same as local_device or it could be CPU or NVMe.
self.remote_device = self.local_device if remote_device is None else remote_device
self.pin_memory = pin_memory if (self.remote_device
== OffloadDeviceEnum.cpu) else False
self.remote_device = self.local_device if remote_device in [
None,
Comment thread
mrwyattii marked this conversation as resolved.
OffloadDeviceEnum.none
] else remote_device
self.pin_memory = pin_memory if (
self.remote_device in [OffloadDeviceEnum.cpu,
OffloadDeviceEnum.nvme]) else False

# Enable fp16 param swapping to NVMe
if self.remote_device == OffloadDeviceEnum.nvme:
Expand Down
10 changes: 10 additions & 0 deletions deepspeed/runtime/zero/partitioned_param_coordinator.py
Original file line number Diff line number Diff line change
Expand Up @@ -400,6 +400,16 @@ def __all_gather_params(self, params: Set[Parameter]) -> None:
assert param.ds_status == ZeroParamStatus.INFLIGHT, param.ds_summary()
self.__inflight_param_registry[param] = handle

# Release swap buffers for persisted params on nvme since they will never be partitioned or evicted from GPU
swap_persisted_params = [
p for p in partitioned_params
if p.ds_persist and p.ds_tensor.final_location == OffloadDeviceEnum.nvme
]
if swap_persisted_params:
swap_persisted_params[
0].nvme_swapper.remove_partition_and_release_buffers(
swap_persisted_params)

@instrument_w_nvtx
def __release_param(self, param: Parameter) -> None:
if param.ds_status == ZeroParamStatus.AVAILABLE and not param.ds_active_sub_modules:
Expand Down
23 changes: 14 additions & 9 deletions deepspeed/runtime/zero/stage3.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
Licensed under the MIT license.
"""

import sys
import gc
import collections
from typing import Deque, Dict, Tuple
Expand Down Expand Up @@ -88,6 +89,7 @@ def __init__(self,
max_reuse_distance=1000000000,
max_live_parameters=1000000000,
param_persistence_threshold=100000,
model_persistence_threshold=sys.maxsize,
dp_process_group=None,
reduce_scatter=True,
overlap_comm=False,
Expand Down Expand Up @@ -146,15 +148,18 @@ def __init__(self,
self.params_in_nvme_and_cpu = False
self.max_params_in_cpu = 0

self.parameter_offload = DeepSpeedZeRoOffload(module,
timers,
ds_config,
overlap_comm,
prefetch_bucket_size,
max_reuse_distance,
max_live_parameters,
param_persistence_threshold,
offload_param_config)
self.parameter_offload = DeepSpeedZeRoOffload(
module=module,
timers=timers,
ds_config=ds_config,
overlap_comm=overlap_comm,
prefetch_bucket_size=prefetch_bucket_size,
max_reuse_distance=max_reuse_distance,
max_live_parameters=max_live_parameters,
param_persistence_threshold=param_persistence_threshold,
model_persistence_threshold=model_persistence_threshold,
offload_param_config=offload_optimizer_config)

self.persistent_parameters = self.parameter_offload.persistent_parameters
self._configure_offloading(offload_optimizer_config, offload_param_config)

Expand Down