arctic-platform 0.1.1.dev0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- arctic_platform/rl/README.md +7 -0
- arctic_platform/rl/__init__.py +27 -0
- arctic_platform/rl/client.py +46 -0
- arctic_platform/rl/config.py +150 -0
- arctic_platform/rl/deepspeed_worker.py +905 -0
- arctic_platform/rl/examples/inference_example.py +115 -0
- arctic_platform/rl/examples/local_server_example.py +201 -0
- arctic_platform/rl/examples/weight_sync_example.py +203 -0
- arctic_platform/rl/http_client.py +705 -0
- arctic_platform/rl/http_server.py +1166 -0
- arctic_platform/rl/processors/__init__.py +166 -0
- arctic_platform/rl/processors/functional.py +257 -0
- arctic_platform/rl/processors/grpo.py +330 -0
- arctic_platform/rl/processors/microbatch.py +176 -0
- arctic_platform/rl/processors/packing.py +90 -0
- arctic_platform/rl/processors/pipeline.py +880 -0
- arctic_platform/rl/processors/stats_tracker.py +186 -0
- arctic_platform/rl/processors/verl_grpo.py +479 -0
- arctic_platform/rl/projects/long_context_qa/README.md +87 -0
- arctic_platform/rl/projects/long_context_qa/download_data.py +120 -0
- arctic_platform/rl/projects/long_context_qa/run_qwen3_32b_longcontext_grpo_arl_zorro_yes_kl.sh +221 -0
- arctic_platform/rl/projects/txt2sql/README.md +176 -0
- arctic_platform/rl/projects/txt2sql/bird_reward.py +272 -0
- arctic_platform/rl/projects/txt2sql/preprocess_bird.py +684 -0
- arctic_platform/rl/projects/txt2sql/run_qwen3_32b_bird_grpo_arl_zorro_yes.sh +251 -0
- arctic_platform/rl/projects/txt2sql/run_qwen3_32b_bird_grpo_arl_zorro_yes_kl.sh +31 -0
- arctic_platform/rl/ray_client.py +482 -0
- arctic_platform/rl/ray_cluster.py +213 -0
- arctic_platform/rl/ray_server.py +1475 -0
- arctic_platform/rl/server.py +6 -0
- arctic_platform/rl/utils/__init__.py +16 -0
- arctic_platform/rl/utils/batch.py +434 -0
- arctic_platform/rl/utils/cuda_ipc.py +57 -0
- arctic_platform/rl/utils/debug.py +495 -0
- arctic_platform/rl/utils/ray_pg.py +210 -0
- arctic_platform/rl/weight_sync.py +238 -0
- arctic_platform/rl/zorro_train/README.md +308 -0
- arctic_platform/rl/zorro_train/__init__.py +23 -0
- arctic_platform/rl/zorro_train/actor.py +271 -0
- arctic_platform/rl/zorro_train/demo.py +130 -0
- arctic_platform/rl/zorro_train/module_patcher.py +67 -0
- arctic_platform/rl/zorro_train/qwen_attention_patcher.py +1071 -0
- arctic_platform/rl/zorro_train/qwen_model_patcher.py +1107 -0
- arctic_platform/rl/zorro_train/seqlen_balancing.py +917 -0
- arctic_platform/rl/zorro_train/test.py +7 -0
- arctic_platform/rl/zorro_train/test_actor_demo.py +68 -0
- arctic_platform/rl/zorro_train/test_forward_and_backward.py +492 -0
- arctic_platform/rl/zorro_train/test_perf-g1.sh +2 -0
- arctic_platform/rl/zorro_train/test_perf-g8.sh +2 -0
- arctic_platform/rl/zorro_train/test_perf.py +503 -0
- arctic_platform/rl/zorro_train/tests.py +711 -0
- arctic_platform/rl/zorro_train/zorro_train.py +2670 -0
- arctic_platform-0.1.1.dev0.dist-info/METADATA +84 -0
- arctic_platform-0.1.1.dev0.dist-info/RECORD +56 -0
- arctic_platform-0.1.1.dev0.dist-info/WHEEL +4 -0
- arctic_platform-0.1.1.dev0.dist-info/licenses/LICENSE +201 -0
|
@@ -0,0 +1,905 @@
|
|
|
1
|
+
# Copyright 2025 Snowflake Inc.
|
|
2
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
3
|
+
#
|
|
4
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
5
|
+
# you may not use this file except in compliance with the License.
|
|
6
|
+
# You may obtain a copy of the License at
|
|
7
|
+
#
|
|
8
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
9
|
+
#
|
|
10
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
11
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
12
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
13
|
+
# See the License for the specific language governing permissions and
|
|
14
|
+
# limitations under the License.
|
|
15
|
+
|
|
16
|
+
"""Local RL server matching the dss-platform sftp_server HTTP API.
|
|
17
|
+
|
|
18
|
+
Uses Ray to manage DeepSpeed workers and ArcticInference ReplicaPools.
|
|
19
|
+
|
|
20
|
+
Usage::
|
|
21
|
+
|
|
22
|
+
python -m arctic_platform.rl.server \\
|
|
23
|
+
--training-gpus 4 --sampling-gpus 2 --log-prob-gpus 2
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
from __future__ import annotations
|
|
27
|
+
|
|
28
|
+
import io
|
|
29
|
+
import logging
|
|
30
|
+
import os
|
|
31
|
+
import time
|
|
32
|
+
from typing import Any
|
|
33
|
+
|
|
34
|
+
import deepspeed
|
|
35
|
+
import ray
|
|
36
|
+
import torch
|
|
37
|
+
import torch.distributed as dist
|
|
38
|
+
import uvicorn
|
|
39
|
+
from arctic_inference.server.weight_sync.sender import WeightSender
|
|
40
|
+
from deepspeed.accelerator import get_accelerator
|
|
41
|
+
from transformers import AutoModelForCausalLM
|
|
42
|
+
import numbers
|
|
43
|
+
from arctic_platform.rl.processors import run_pipeline
|
|
44
|
+
from arctic_platform.rl.utils import (
|
|
45
|
+
unpack_batch,
|
|
46
|
+
merge_dict_shards,
|
|
47
|
+
combine_metric_microbatches,
|
|
48
|
+
split_dict,
|
|
49
|
+
log_dp_shard_tokens,
|
|
50
|
+
)
|
|
51
|
+
from arctic_platform.rl.ray_cluster import primary_ip
|
|
52
|
+
from arctic_platform.rl.utils.debug import enable_full_determinism
|
|
53
|
+
from arctic_platform.rl.utils.debug import see_memory_usage, pr, pr0
|
|
54
|
+
|
|
55
|
+
logger = logging.getLogger(__name__)
|
|
56
|
+
|
|
57
|
+
# ---------------------------------------------------------------------------
|
|
58
|
+
# Request / response models (mirrors dss-platform sftp_server)
|
|
59
|
+
# ---------------------------------------------------------------------------
|
|
60
|
+
|
|
61
|
+
ENABLE_TIMERS = True
|
|
62
|
+
if ENABLE_TIMERS:
|
|
63
|
+
from arctic_platform.rl.utils.debug import SynchronizedWallClockTimerSimple
|
|
64
|
+
timers = SynchronizedWallClockTimerSimple(wall_clock_breakdown=True)
|
|
65
|
+
else:
|
|
66
|
+
from arctic_platform.rl.utils.debug import SynchronizedWallClockTimerSimpleDummy
|
|
67
|
+
timers = SynchronizedWallClockTimerSimpleDummy(wall_clock_breakdown=True)
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def make_model_gradient_checkpointing_compatible(model):
|
|
71
|
+
# Taken from arctic_platform/model/hf_factory.py
|
|
72
|
+
if hasattr(model, "enable_input_require_grads"):
|
|
73
|
+
model.enable_input_require_grads()
|
|
74
|
+
elif hasattr(model, "get_input_embeddings"):
|
|
75
|
+
|
|
76
|
+
def make_inputs_require_grad(module, input, output):
|
|
77
|
+
output.requires_grad_(True)
|
|
78
|
+
|
|
79
|
+
model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)
|
|
80
|
+
return model
|
|
81
|
+
|
|
82
|
+
# ---------------------------------------------------------------------------
|
|
83
|
+
# DeepSpeed training actor
|
|
84
|
+
# ---------------------------------------------------------------------------
|
|
85
|
+
|
|
86
|
+
import socket
|
|
87
|
+
@ray.remote
|
|
88
|
+
class DeepSpeedWorker:
|
|
89
|
+
"""Single-GPU worker for DeepSpeed training."""
|
|
90
|
+
|
|
91
|
+
def __init__(self, rank: int, world_size: int, master_port: int):
|
|
92
|
+
self.rank = rank
|
|
93
|
+
self.world_size = world_size
|
|
94
|
+
self.my_addr = socket.gethostname()
|
|
95
|
+
self.master_addr = primary_ip()
|
|
96
|
+
self.master_port = master_port
|
|
97
|
+
self.engine = None
|
|
98
|
+
self._weight_sender: WeightSender | None = None
|
|
99
|
+
self._on_gpu = True
|
|
100
|
+
|
|
101
|
+
def get_ip(self) -> str:
|
|
102
|
+
return self.my_addr
|
|
103
|
+
|
|
104
|
+
def initialize(self, master_addr: str, job_config: dict) -> bool:
|
|
105
|
+
self.master_addr = master_addr
|
|
106
|
+
os.environ.update(
|
|
107
|
+
{
|
|
108
|
+
"RANK": str(self.rank),
|
|
109
|
+
"LOCAL_RANK": "0",
|
|
110
|
+
"WORLD_SIZE": str(self.world_size),
|
|
111
|
+
"MASTER_ADDR": self.master_addr,
|
|
112
|
+
"MASTER_PORT": str(self.master_port),
|
|
113
|
+
}
|
|
114
|
+
)
|
|
115
|
+
|
|
116
|
+
if job_config.get("full_determinism", False):
|
|
117
|
+
enable_full_determinism(seed=job_config.get("seed", 42))
|
|
118
|
+
|
|
119
|
+
deepspeed.init_distributed()
|
|
120
|
+
|
|
121
|
+
model_name = job_config["model_name"]
|
|
122
|
+
ds_config = job_config.get("ds_config") or {}
|
|
123
|
+
self.job_type = job_config.get("job_type")
|
|
124
|
+
pr0(f"{self.job_type=} {job_config=}")
|
|
125
|
+
pr0(f"ds_worker[before_modify]: {self.job_type=} {ds_config=}")
|
|
126
|
+
|
|
127
|
+
ds_worker_config = job_config.get("ds_worker_config") or {}
|
|
128
|
+
ds_worker_config["world_size"] = self.world_size
|
|
129
|
+
self.ds_worker_config = ds_worker_config
|
|
130
|
+
|
|
131
|
+
# Build the DeepSpeed config per job type. Training engines get an
|
|
132
|
+
# optimizer; the reference/log-prob engine is forward-only and is
|
|
133
|
+
# configured from log_prob_config with no optimizer state.
|
|
134
|
+
if self.job_type == "log_prob":
|
|
135
|
+
log_prob_config = job_config.get("log_prob_config") or {}
|
|
136
|
+
ds_config = self.ds_inference_config(log_prob_config, ds_worker_config)
|
|
137
|
+
self._has_optimizer = False
|
|
138
|
+
else:
|
|
139
|
+
ds_config = self.ds_training_config(job_config, ds_config, ds_worker_config)
|
|
140
|
+
self._has_optimizer = True
|
|
141
|
+
|
|
142
|
+
pr0(f"ds_worker[after_modify]: {self.job_type=} {ds_config=} {ds_worker_config=}")
|
|
143
|
+
|
|
144
|
+
attn_implementation = ds_worker_config.get("attn_implementation", "flash_attention_2")
|
|
145
|
+
|
|
146
|
+
model = AutoModelForCausalLM.from_pretrained(
|
|
147
|
+
model_name,
|
|
148
|
+
attn_implementation=attn_implementation,
|
|
149
|
+
dtype=torch.bfloat16,
|
|
150
|
+
)
|
|
151
|
+
|
|
152
|
+
if ds_worker_config.get("use_liger", False):
|
|
153
|
+
pr0(f"Using liger kernel w/ {attn_implementation=}")
|
|
154
|
+
from liger_kernel.transformers import AutoLigerKernelForCausalLM
|
|
155
|
+
# Apply Liger kernel to the model if use_liger is set to True
|
|
156
|
+
from liger_kernel.transformers.monkey_patch import _apply_liger_kernel_to_instance
|
|
157
|
+
|
|
158
|
+
_apply_liger_kernel_to_instance(
|
|
159
|
+
model=model,
|
|
160
|
+
cross_entropy=False,
|
|
161
|
+
fused_linear_cross_entropy=False,
|
|
162
|
+
rope=True,
|
|
163
|
+
rms_norm=True,
|
|
164
|
+
swiglu=True,
|
|
165
|
+
)
|
|
166
|
+
# if ds_worker_config.get("use_liger", False):
|
|
167
|
+
# pr0(f"Using liger kernel w/ {attn_implementation=}")
|
|
168
|
+
# from liger_kernel.transformers import AutoLigerKernelForCausalLM
|
|
169
|
+
# model = AutoLigerKernelForCausalLM.from_pretrained(
|
|
170
|
+
# model_name,
|
|
171
|
+
# attn_implementation=attn_implementation,
|
|
172
|
+
# dtype=torch.bfloat16,
|
|
173
|
+
# )
|
|
174
|
+
|
|
175
|
+
# else:
|
|
176
|
+
# model = AutoModelForCausalLM.from_pretrained(
|
|
177
|
+
# model_name,
|
|
178
|
+
# attn_implementation=attn_implementation,
|
|
179
|
+
# dtype=torch.bfloat16,
|
|
180
|
+
# )
|
|
181
|
+
|
|
182
|
+
zorro_train_enable = ds_worker_config.get("zorro_train_enable", False)
|
|
183
|
+
if zorro_train_enable:
|
|
184
|
+
self.model_patch_in_zorro(model, ds_worker_config)
|
|
185
|
+
|
|
186
|
+
init_kwargs = dict(model=model, config=ds_config)
|
|
187
|
+
if self._has_optimizer:
|
|
188
|
+
# Forward-only (log-prob) engines are initialized without an
|
|
189
|
+
# optimizer so DeepSpeed allocates no optimizer state.
|
|
190
|
+
init_kwargs["model_parameters"] = model.parameters()
|
|
191
|
+
self.engine, _, _, _ = deepspeed.initialize(**init_kwargs)
|
|
192
|
+
self._device = get_accelerator().device_name(self.engine.local_rank)
|
|
193
|
+
|
|
194
|
+
gpu_id = torch.cuda.current_device()
|
|
195
|
+
gpu_uuid = torch.cuda.get_device_properties(gpu_id).uuid
|
|
196
|
+
logger.info("Rank %d initialized on GPU %d (uuid=%s, device=%s)",
|
|
197
|
+
self.rank, gpu_id, gpu_uuid, self._device)
|
|
198
|
+
self.cpu_device = torch.device("cpu")
|
|
199
|
+
|
|
200
|
+
enable_gradient_checkpointing = ds_worker_config.get("enable_gradient_checkpointing", True)
|
|
201
|
+
if enable_gradient_checkpointing:
|
|
202
|
+
model.gradient_checkpointing_enable()
|
|
203
|
+
|
|
204
|
+
pr0(f"ds_worker[after_initialize]: {self.job_type=} {self.engine.global_steps=} {zorro_train_enable=} {model_name=}")
|
|
205
|
+
|
|
206
|
+
return True
|
|
207
|
+
|
|
208
|
+
|
|
209
|
+
def ds_training_config(self, job_config: dict, ds_config: dict, ds_worker_config: dict) -> dict:
|
|
210
|
+
"""Build the DeepSpeed config for a trainable engine (with optimizer).
|
|
211
|
+
|
|
212
|
+
Prefer the high-level training_config when provided (the framework sends
|
|
213
|
+
training_config and the server owns the DeepSpeed details); otherwise
|
|
214
|
+
fall back to a default AdamW optimizer.
|
|
215
|
+
"""
|
|
216
|
+
training_config = job_config.get("training_config")
|
|
217
|
+
if training_config is not None:
|
|
218
|
+
opt_cfg = training_config.get("optimizer", {})
|
|
219
|
+
ds_config.setdefault("optimizer", {
|
|
220
|
+
"type": "AdamW",
|
|
221
|
+
"params": {
|
|
222
|
+
"lr": opt_cfg.get("lr", 1e-5),
|
|
223
|
+
"betas": [opt_cfg.get("beta1", 0.9), opt_cfg.get("beta2", 0.999)],
|
|
224
|
+
"eps": 1e-8,
|
|
225
|
+
"weight_decay": opt_cfg.get("weight_decay", 0.0),
|
|
226
|
+
},
|
|
227
|
+
})
|
|
228
|
+
if "gradient_accumulation_steps" in training_config:
|
|
229
|
+
ds_config.setdefault("gradient_accumulation_steps", training_config["gradient_accumulation_steps"])
|
|
230
|
+
if "gradient_clipping" in opt_cfg:
|
|
231
|
+
ds_config.setdefault("gradient_clipping", opt_cfg["gradient_clipping"])
|
|
232
|
+
|
|
233
|
+
ds_config.setdefault("train_micro_batch_size_per_gpu", 1)
|
|
234
|
+
|
|
235
|
+
if ds_worker_config.get("use_autocast", False):
|
|
236
|
+
ds_config.setdefault("torch_autocast", {"enabled": True, "dtype": "bfloat16"})
|
|
237
|
+
ds_config.setdefault("bf16", {"enabled": True, "bf16_master_weights_and_grads": True, "bf16_optimizer_states": True})
|
|
238
|
+
else:
|
|
239
|
+
ds_config.setdefault("bf16", {"enabled": True})
|
|
240
|
+
|
|
241
|
+
ds_config.setdefault(
|
|
242
|
+
"optimizer",
|
|
243
|
+
{
|
|
244
|
+
"type": "AdamW",
|
|
245
|
+
"params": {"lr": 1e-5, "betas": [0.9, 0.999], "eps": 1e-8},
|
|
246
|
+
},
|
|
247
|
+
)
|
|
248
|
+
return ds_config
|
|
249
|
+
|
|
250
|
+
|
|
251
|
+
def ds_inference_config(self, log_prob_config: dict, ds_worker_config: dict) -> dict:
|
|
252
|
+
"""Build a forward-only DeepSpeed config (no optimizer state).
|
|
253
|
+
|
|
254
|
+
Used for the reference / log-prob engine. Keeps ZeRO param sharding
|
|
255
|
+
(stage + offload_param) and bf16, but omits the optimizer, gradient
|
|
256
|
+
accumulation, gradient clipping, train_batch_size, and any
|
|
257
|
+
offload_optimizer so DeepSpeed allocates no optimizer state.
|
|
258
|
+
"""
|
|
259
|
+
src = dict(log_prob_config or {})
|
|
260
|
+
zero = dict(src.get("zero_optimization", {}) or {})
|
|
261
|
+
zero.pop("offload_optimizer", None)
|
|
262
|
+
|
|
263
|
+
cfg: dict = {
|
|
264
|
+
"train_micro_batch_size_per_gpu": src.get("train_micro_batch_size_per_gpu", 1),
|
|
265
|
+
}
|
|
266
|
+
if zero:
|
|
267
|
+
cfg["zero_optimization"] = zero
|
|
268
|
+
if "sequence_parallel_size" in src:
|
|
269
|
+
cfg["sequence_parallel_size"] = src["sequence_parallel_size"]
|
|
270
|
+
if ds_worker_config.get("use_autocast", False):
|
|
271
|
+
cfg["torch_autocast"] = {"enabled": True, "dtype": "bfloat16"}
|
|
272
|
+
cfg["bf16"] = {"enabled": True}
|
|
273
|
+
return cfg
|
|
274
|
+
|
|
275
|
+
|
|
276
|
+
def model_patch_in_zorro(self, model, ds_worker_config):
|
|
277
|
+
from arctic_platform.rl.zorro_train.qwen_model_patcher import Qwen3ModelOncePatcher
|
|
278
|
+
|
|
279
|
+
#pr0(f"Patching ZoRRO")
|
|
280
|
+
|
|
281
|
+
response_len = ds_worker_config.get("response_len")
|
|
282
|
+
max_token_len = ds_worker_config.get("max_token_len")
|
|
283
|
+
rollout_n = ds_worker_config.get("rollout_n")
|
|
284
|
+
temperature = ds_worker_config.get("temperature")
|
|
285
|
+
logits_optimization = ds_worker_config.get("logits_optimization", "none")
|
|
286
|
+
logits_optimization_peak_mem_size_in_gib = ds_worker_config.get("logits_optimization_peak_mem_size_in_gib", 4)
|
|
287
|
+
logits_compute_from_fp32_inputs = ds_worker_config.get("logits_compute_from_fp32_inputs", False)
|
|
288
|
+
logits_compute_in_fp32 = ds_worker_config.get("logits_compute_in_fp32", False)
|
|
289
|
+
use_unpad = ds_worker_config.get("use_unpad")
|
|
290
|
+
world_size = ds_worker_config.get("world_size")
|
|
291
|
+
|
|
292
|
+
self.dedup_actor_model_once_patcher = Qwen3ModelOncePatcher(model, response_len=response_len, max_token_len=max_token_len, rollout_n=rollout_n, temperature=temperature, logits_optimization=logits_optimization, logits_optimization_peak_mem_size_in_gib=logits_optimization_peak_mem_size_in_gib, logits_compute_from_fp32_inputs=logits_compute_from_fp32_inputs, logits_compute_in_fp32=logits_compute_in_fp32, use_unpad=use_unpad, world_size=world_size)
|
|
293
|
+
self.dedup_actor_model_once_patcher.patch_forward()
|
|
294
|
+
|
|
295
|
+
# move batch to device
|
|
296
|
+
def _move_batch_to_device(self, batch: Any, device: torch.device):
|
|
297
|
+
if isinstance(batch, dict):
|
|
298
|
+
return {k: self._move_batch_to_device(v, device) for k, v in batch.items()}
|
|
299
|
+
elif isinstance(batch, (list, tuple)):
|
|
300
|
+
return [self._move_batch_to_device(v, device) for v in batch]
|
|
301
|
+
elif isinstance(batch, torch.Tensor):
|
|
302
|
+
return batch.to(device)
|
|
303
|
+
return batch
|
|
304
|
+
|
|
305
|
+
def _forward_maybe_backward(self, batch: dict, backward: bool) -> dict:
|
|
306
|
+
#torch.autograd.set_detect_anomaly(True)
|
|
307
|
+
|
|
308
|
+
pr0(f"_forward_maybe_backward mode: {backward=}")
|
|
309
|
+
PROFILE = False
|
|
310
|
+
# if backward:
|
|
311
|
+
# PROFILE = True
|
|
312
|
+
if PROFILE:
|
|
313
|
+
torch.cuda.memory._record_memory_history(max_entries=int(1e12))
|
|
314
|
+
see_memory_usage("_forward_maybe_backward start", force=True)
|
|
315
|
+
|
|
316
|
+
args, batch_data, meta_data, processing = unpack_batch(batch)
|
|
317
|
+
batch_data = self._move_batch_to_device(batch_data, self._device)
|
|
318
|
+
|
|
319
|
+
tag = "forward_only" if not backward else "forward_backward"
|
|
320
|
+
|
|
321
|
+
log_dp_shard_tokens(self.rank, f"{tag} shard", batch_data, meta_data)
|
|
322
|
+
|
|
323
|
+
pr0(f"[DeepSpeedWorker] {tag}: {batch_data.keys()=} {meta_data.keys()=} {processing.keys()=}")
|
|
324
|
+
|
|
325
|
+
for k, v in batch_data.items():
|
|
326
|
+
pr0(f"[DeepSpeedWorker] {tag}: {k=}: {v.shape=}")
|
|
327
|
+
|
|
328
|
+
grad_accum_steps = self.engine.gradient_accumulation_steps()
|
|
329
|
+
micro_batch_data = split_dict(batch_data, grad_accum_steps)
|
|
330
|
+
num_micro_batches = len(micro_batch_data)
|
|
331
|
+
pipeline_micro_batch_outputs = []
|
|
332
|
+
return_tensors = meta_data.get("worker_return_tensors", False)
|
|
333
|
+
|
|
334
|
+
pr0(f"mbs {len(micro_batch_data)=} {grad_accum_steps=}")
|
|
335
|
+
|
|
336
|
+
for i, micro_batch in enumerate(micro_batch_data):
|
|
337
|
+
import time
|
|
338
|
+
#time.sleep(1)
|
|
339
|
+
#pr0(f"{i=}")
|
|
340
|
+
#pr0(f"{micro_batch.keys()=}")
|
|
341
|
+
|
|
342
|
+
log_dp_shard_tokens(
|
|
343
|
+
self.rank, f"{tag} micro_batch {i}/{num_micro_batches}", micro_batch, meta_data,
|
|
344
|
+
)
|
|
345
|
+
|
|
346
|
+
DEBUG = False
|
|
347
|
+
if DEBUG:
|
|
348
|
+
from arctic_platform.rl.zorro_train import analyze_normal_batch_via_attention_mask
|
|
349
|
+
analyze_normal_batch_via_attention_mask(micro_batch["input_ids"], micro_batch["attention_mask"], response_len=meta_data["max_response_len"])
|
|
350
|
+
|
|
351
|
+
#die
|
|
352
|
+
see_memory_usage(f"_forward_maybe_backward mb {i=}", force=True)
|
|
353
|
+
if i == 0:
|
|
354
|
+
pr0(f"[DeepSpeedWorker] {tag}: {i=}/{num_micro_batches=} {meta_data.keys()=} {processing.keys()=}")
|
|
355
|
+
|
|
356
|
+
micro_batch_output = run_pipeline(
|
|
357
|
+
self.engine, args, micro_batch, meta_data, processing,
|
|
358
|
+
device=self._device, backward=backward,
|
|
359
|
+
pack=False,
|
|
360
|
+
return_tensors=return_tensors,
|
|
361
|
+
)
|
|
362
|
+
|
|
363
|
+
if i == 0:
|
|
364
|
+
pr0(f"[DeepSpeedWorker] {tag}: {i=}/{num_micro_batches=} {micro_batch_output.keys()=}")
|
|
365
|
+
pipeline_micro_batch_outputs.append(micro_batch_output)
|
|
366
|
+
|
|
367
|
+
# DS requires matching steps for backward pass
|
|
368
|
+
if backward and i < num_micro_batches - 1:
|
|
369
|
+
self.engine.step()
|
|
370
|
+
|
|
371
|
+
pipeline_outputs = dict()
|
|
372
|
+
for k, v in pipeline_micro_batch_outputs[0].items():
|
|
373
|
+
if k == "metrics" and isinstance(v, dict):
|
|
374
|
+
# Per-microbatch loss-fn metrics are emitted as paired
|
|
375
|
+
# ``{name}.sum`` / ``{name}.tokens`` scalars (plus a few
|
|
376
|
+
# passthrough numerics like ``kl_coef``). Sum them across
|
|
377
|
+
# this rank's microbatches so each rank returns one scalar
|
|
378
|
+
# per metric; ``ray_server.fwd_bwd`` / ``http_server.fwd_bwd``
|
|
379
|
+
# then sums across DP ranks and collapses the paired keys
|
|
380
|
+
# into a single global token-mean per metric per mini-batch.
|
|
381
|
+
pipeline_outputs[k] = combine_metric_microbatches(
|
|
382
|
+
[r[k] for r in pipeline_micro_batch_outputs]
|
|
383
|
+
)
|
|
384
|
+
elif isinstance(v, dict):
|
|
385
|
+
pipeline_outputs[k] = merge_dict_shards([r[k] for r in pipeline_micro_batch_outputs])
|
|
386
|
+
elif isinstance(v, numbers.Number):
|
|
387
|
+
# TODO: weight average needs to be implemented
|
|
388
|
+
pipeline_outputs[k] = sum([r[k] for r in pipeline_micro_batch_outputs]) / len(pipeline_micro_batch_outputs)
|
|
389
|
+
|
|
390
|
+
pipeline_outputs = self._move_batch_to_device(pipeline_outputs, self.cpu_device)
|
|
391
|
+
|
|
392
|
+
see_memory_usage("_forward_maybe_backward end", force=True)
|
|
393
|
+
if PROFILE:
|
|
394
|
+
dir = "/tmp/mem-prof"
|
|
395
|
+
rank = 0 # torch.distributed.get_rank()
|
|
396
|
+
from pathlib import Path
|
|
397
|
+
Path(dir).mkdir(exist_ok=True)
|
|
398
|
+
torch.cuda.memory._dump_snapshot(f"{dir}/rank-{rank}.pickle")
|
|
399
|
+
exit()
|
|
400
|
+
|
|
401
|
+
pr0(f"[DeepSpeedWorker] {tag}: {pipeline_outputs.keys()=}")
|
|
402
|
+
return pipeline_outputs
|
|
403
|
+
|
|
404
|
+
def forward_backward(self, batch: dict) -> dict:
|
|
405
|
+
tname = timers.start("forward_backward")
|
|
406
|
+
results = self._forward_maybe_backward(batch, backward=True)
|
|
407
|
+
timers.stop_and_print_elapsed(tname);
|
|
408
|
+
return results
|
|
409
|
+
|
|
410
|
+
def forward_no_grad(self, batch: dict) -> dict:
|
|
411
|
+
tname = timers.start("forward_no_grad")
|
|
412
|
+
results = self._forward_maybe_backward(batch, backward=False)
|
|
413
|
+
timers.stop_and_print_elapsed(tname);
|
|
414
|
+
return results
|
|
415
|
+
|
|
416
|
+
def step(self) -> dict:
|
|
417
|
+
self.engine.step()
|
|
418
|
+
# Pull grad_norm out of DeepSpeed so it can be logged by the trainer.
|
|
419
|
+
# rename_dict in ray_trainer turns "grad_norm" -> "actor/grad_norm",
|
|
420
|
+
# matching the FSDP baseline path in verl/workers/actor/dp_actor.py.
|
|
421
|
+
grad_norm = self.engine.get_global_grad_norm()
|
|
422
|
+
if isinstance(grad_norm, torch.Tensor):
|
|
423
|
+
grad_norm = grad_norm.item()
|
|
424
|
+
metrics = dict(
|
|
425
|
+
last_lr=self.engine.get_lr()[0],
|
|
426
|
+
)
|
|
427
|
+
if grad_norm is not None:
|
|
428
|
+
metrics["grad_norm"] = grad_norm
|
|
429
|
+
return dict(metrics=metrics, batch=dict())
|
|
430
|
+
|
|
431
|
+
def save_checkpoint(self, path: str) -> bool:
|
|
432
|
+
self.engine.save_checkpoint(path)
|
|
433
|
+
return True
|
|
434
|
+
|
|
435
|
+
def compute_log_probs(self, batch_bytes: bytes) -> bytes:
|
|
436
|
+
batch = torch.load(io.BytesIO(batch_bytes), map_location=self._device)
|
|
437
|
+
args, kwargs, _, _ = unpack_batch(batch)
|
|
438
|
+
with torch.no_grad():
|
|
439
|
+
logits = self.engine(*args, **kwargs).logits
|
|
440
|
+
log_probs = torch.log_softmax(logits, dim=-1)
|
|
441
|
+
shifted_ids = kwargs["input_ids"][:, 1:]
|
|
442
|
+
token_log_probs = log_probs[:, :-1].gather(-1, shifted_ids.unsqueeze(-1)).squeeze(-1)
|
|
443
|
+
buf = io.BytesIO()
|
|
444
|
+
torch.save(token_log_probs.cpu(), buf)
|
|
445
|
+
return buf.getvalue()
|
|
446
|
+
|
|
447
|
+
def max_param_bytes(self) -> int:
|
|
448
|
+
max_bytes = 0
|
|
449
|
+
for p in self.engine.module.parameters():
|
|
450
|
+
elem_size = p.data.element_size()
|
|
451
|
+
numel = p.ds_numel if hasattr(p, "ds_id") else p.data.numel()
|
|
452
|
+
max_bytes = max(max_bytes, numel * elem_size)
|
|
453
|
+
return max_bytes
|
|
454
|
+
|
|
455
|
+
|
|
456
|
+
def init_weight_sender(self, group, schedule, master_addr, base_port, bucket_size) -> bool:
|
|
457
|
+
self._weight_sender = WeightSender(
|
|
458
|
+
group=group,
|
|
459
|
+
schedule=schedule,
|
|
460
|
+
master_addr=master_addr,
|
|
461
|
+
base_port=base_port,
|
|
462
|
+
device=torch.device(self._device),
|
|
463
|
+
bucket_size=bucket_size,
|
|
464
|
+
)
|
|
465
|
+
return True
|
|
466
|
+
|
|
467
|
+
def get_weights(self) -> list[tuple[str, torch.Tensor]]:
|
|
468
|
+
weights = []
|
|
469
|
+
for n, p in self.engine.module.named_parameters():
|
|
470
|
+
if hasattr(p, "ds_id"):
|
|
471
|
+
with deepspeed.zero.GatheredParameters([p], enabled=True):
|
|
472
|
+
weights.append((n, p.data))
|
|
473
|
+
else:
|
|
474
|
+
weights.append((n, p.data))
|
|
475
|
+
return weights
|
|
476
|
+
|
|
477
|
+
def send_weights(self) -> dict:
|
|
478
|
+
weights = self.get_weights()
|
|
479
|
+
if self._weight_sender is not None:
|
|
480
|
+
return self._weight_sender.send(weights)
|
|
481
|
+
return {"status": "not_initialized"}
|
|
482
|
+
|
|
483
|
+
|
|
484
|
+
def send_weights_ipc(self, group_id: int) -> dict:
|
|
485
|
+
"""Save weights to shared memory for colocated (same-GPU) transfer."""
|
|
486
|
+
from arctic_inference.server.weight_sync.ipc_engine import save_weights_to_shm
|
|
487
|
+
weights = [(n, p.data) for n, p in self.engine.module.named_parameters()]
|
|
488
|
+
return save_weights_to_shm(weights, group_id)
|
|
489
|
+
|
|
490
|
+
def get_cuda_ipc_handles(self) -> dict:
|
|
491
|
+
"""Create CUDA IPC handles for all model parameters.
|
|
492
|
+
|
|
493
|
+
Returns a dict with names, dtypes, shapes and pickled IPC handles
|
|
494
|
+
that can be opened by another process on the same GPU.
|
|
495
|
+
|
|
496
|
+
For ZeRO-2 this reads p.data directly (full params on every rank).
|
|
497
|
+
For ZeRO-3 this is incorrect — use gather_cuda_ipc_handles instead.
|
|
498
|
+
"""
|
|
499
|
+
import base64
|
|
500
|
+
import pickle
|
|
501
|
+
from torch.multiprocessing.reductions import reduce_tensor
|
|
502
|
+
|
|
503
|
+
gpu_uuid = str(torch.cuda.get_device_properties(
|
|
504
|
+
torch.cuda.current_device()).uuid)
|
|
505
|
+
|
|
506
|
+
names, dtype_names, shapes = [], [], []
|
|
507
|
+
handles = []
|
|
508
|
+
self._ipc_tensor_refs = []
|
|
509
|
+
|
|
510
|
+
for name, p in self.engine.module.named_parameters():
|
|
511
|
+
weight = p.data.detach().contiguous()
|
|
512
|
+
self._ipc_tensor_refs.append(weight)
|
|
513
|
+
handle = reduce_tensor(weight)
|
|
514
|
+
handles.append({gpu_uuid: handle})
|
|
515
|
+
names.append(name)
|
|
516
|
+
dtype_names.append(str(weight.dtype).split(".")[-1])
|
|
517
|
+
shapes.append(list(weight.shape))
|
|
518
|
+
|
|
519
|
+
torch.cuda.synchronize()
|
|
520
|
+
pickled = base64.b64encode(pickle.dumps(handles)).decode("utf-8")
|
|
521
|
+
return {
|
|
522
|
+
"names": names,
|
|
523
|
+
"dtype_names": dtype_names,
|
|
524
|
+
"shapes": shapes,
|
|
525
|
+
"ipc_handles_pickled": pickled,
|
|
526
|
+
"num_params": len(names),
|
|
527
|
+
}
|
|
528
|
+
|
|
529
|
+
def gather_cuda_ipc_handles(self) -> dict:
|
|
530
|
+
"""Gather ZeRO-3 partitioned params and create CUDA IPC handles.
|
|
531
|
+
|
|
532
|
+
All ranks must call this collectively (GatheredParameters is a
|
|
533
|
+
collective op). Inside the context manager the full param lives
|
|
534
|
+
on every rank's GPU — each rank clones it and creates an IPC handle
|
|
535
|
+
keyed by its own GPU UUID before the context manager frees the
|
|
536
|
+
gathered tensor. Every rank returns its payload; the server merges
|
|
537
|
+
them so each colocated inference replica finds a handle for its GPU.
|
|
538
|
+
"""
|
|
539
|
+
import base64
|
|
540
|
+
import pickle
|
|
541
|
+
|
|
542
|
+
import deepspeed
|
|
543
|
+
from torch.multiprocessing.reductions import reduce_tensor
|
|
544
|
+
|
|
545
|
+
gpu_uuid = str(torch.cuda.get_device_properties(
|
|
546
|
+
torch.cuda.current_device()).uuid)
|
|
547
|
+
|
|
548
|
+
t0 = time.monotonic()
|
|
549
|
+
names, dtype_names, shapes = [], [], []
|
|
550
|
+
handles = []
|
|
551
|
+
self._ipc_tensor_refs = []
|
|
552
|
+
model = self.engine.module
|
|
553
|
+
|
|
554
|
+
# Every rank builds IPC handles for the full (gathered) weight on its
|
|
555
|
+
# OWN physical GPU, keyed by that GPU's UUID. With colocated multi-GPU
|
|
556
|
+
# inference, each vLLM replica lives on a distinct physical GPU (bundle
|
|
557
|
+
# r == training rank r), so it needs a handle for *its* GPU. The server
|
|
558
|
+
# merges these per-rank dicts so each replica finds its GPU's handle.
|
|
559
|
+
# (Previously only rank 0 produced handles, which only worked when a
|
|
560
|
+
# single sampling GPU was colocated with rank 0.)
|
|
561
|
+
for name, p in model.named_parameters():
|
|
562
|
+
if hasattr(p, "ds_id"):
|
|
563
|
+
with deepspeed.zero.GatheredParameters([p], enabled=True):
|
|
564
|
+
weight = p.data.detach().clone().contiguous()
|
|
565
|
+
self._ipc_tensor_refs.append(weight)
|
|
566
|
+
handle = reduce_tensor(weight)
|
|
567
|
+
handles.append({gpu_uuid: handle})
|
|
568
|
+
names.append(name)
|
|
569
|
+
dtype_names.append(str(weight.dtype).split(".")[-1])
|
|
570
|
+
shapes.append(list(weight.shape))
|
|
571
|
+
else:
|
|
572
|
+
weight = p.data.detach().contiguous()
|
|
573
|
+
self._ipc_tensor_refs.append(weight)
|
|
574
|
+
handle = reduce_tensor(weight)
|
|
575
|
+
handles.append({gpu_uuid: handle})
|
|
576
|
+
names.append(name)
|
|
577
|
+
dtype_names.append(str(weight.dtype).split(".")[-1])
|
|
578
|
+
shapes.append(list(weight.shape))
|
|
579
|
+
|
|
580
|
+
if hasattr(self.engine, 'empty_partition_cache'):
|
|
581
|
+
self.engine.empty_partition_cache()
|
|
582
|
+
torch.cuda.synchronize()
|
|
583
|
+
|
|
584
|
+
elapsed = time.monotonic() - t0
|
|
585
|
+
pickled = base64.b64encode(pickle.dumps(handles)).decode("utf-8")
|
|
586
|
+
logger.info("Rank %d gathered IPC handles in %.2fs (%d params, gpu=%s)",
|
|
587
|
+
self.rank, elapsed, len(names), gpu_uuid)
|
|
588
|
+
return {
|
|
589
|
+
"names": names,
|
|
590
|
+
"dtype_names": dtype_names,
|
|
591
|
+
"shapes": shapes,
|
|
592
|
+
"ipc_handles_pickled": pickled,
|
|
593
|
+
"num_params": len(names),
|
|
594
|
+
"gpu_uuid": gpu_uuid,
|
|
595
|
+
}
|
|
596
|
+
|
|
597
|
+
def release_ipc_handles(self) -> bool:
|
|
598
|
+
"""Release tensor references held for IPC handles."""
|
|
599
|
+
self._ipc_tensor_refs = []
|
|
600
|
+
torch.cuda.ipc_collect()
|
|
601
|
+
torch.cuda.synchronize()
|
|
602
|
+
return True
|
|
603
|
+
|
|
604
|
+
def get_parameter_names(self) -> list:
|
|
605
|
+
"""Return the model's parameter names in deterministic order.
|
|
606
|
+
|
|
607
|
+
Used to drive the low-memory streaming weight sync one parameter at a
|
|
608
|
+
time. Only names cross the Ray boundary; the live parameter is resolved
|
|
609
|
+
on each rank inside ``get_cuda_ipc_handle`` so ZeRO-3 ``ds_id`` / live
|
|
610
|
+
storage is preserved.
|
|
611
|
+
"""
|
|
612
|
+
return [name for name, _ in self.engine.module.named_parameters()]
|
|
613
|
+
|
|
614
|
+
def _param_by_name(self, name: str):
|
|
615
|
+
"""Resolve this rank's live module parameter for ``name`` (cached)."""
|
|
616
|
+
cache = getattr(self, "_param_name_cache", None)
|
|
617
|
+
if cache is None:
|
|
618
|
+
cache = dict(self.engine.module.named_parameters())
|
|
619
|
+
self._param_name_cache = cache
|
|
620
|
+
return cache[name]
|
|
621
|
+
|
|
622
|
+
def get_cuda_ipc_handle(self, name: str) -> dict:
|
|
623
|
+
"""Create a CUDA IPC handle payload for a single parameter on this
|
|
624
|
+
rank's GPU.
|
|
625
|
+
|
|
626
|
+
Memory-efficient counterpart to ``gather_cuda_ipc_handles``: only ONE
|
|
627
|
+
full parameter is materialized at a time instead of the whole model.
|
|
628
|
+
For ZeRO-3 params all ranks must call this collectively with the same
|
|
629
|
+
``name`` (``GatheredParameters`` is a collective op). The caller must
|
|
630
|
+
invoke ``release_ipc_handles`` between params so peak extra GPU memory
|
|
631
|
+
stays at one full parameter per GPU.
|
|
632
|
+
"""
|
|
633
|
+
import base64
|
|
634
|
+
import pickle
|
|
635
|
+
|
|
636
|
+
import deepspeed
|
|
637
|
+
from torch.multiprocessing.reductions import reduce_tensor
|
|
638
|
+
|
|
639
|
+
p = self._param_by_name(name)
|
|
640
|
+
|
|
641
|
+
gpu_uuid = str(torch.cuda.get_device_properties(
|
|
642
|
+
torch.cuda.current_device()).uuid)
|
|
643
|
+
|
|
644
|
+
if hasattr(p, "ds_id"):
|
|
645
|
+
with deepspeed.zero.GatheredParameters([p], enabled=True):
|
|
646
|
+
weight = p.data.detach().clone().contiguous()
|
|
647
|
+
else:
|
|
648
|
+
weight = p.data.detach().contiguous()
|
|
649
|
+
|
|
650
|
+
# Hold exactly one source tensor alive until release_ipc_handles().
|
|
651
|
+
self._ipc_tensor_refs = [weight]
|
|
652
|
+
handle = reduce_tensor(weight)
|
|
653
|
+
torch.cuda.synchronize()
|
|
654
|
+
|
|
655
|
+
return {
|
|
656
|
+
"names": [name],
|
|
657
|
+
"dtype_names": [str(weight.dtype).split(".")[-1]],
|
|
658
|
+
"shapes": [list(weight.shape)],
|
|
659
|
+
"ipc_handles_pickled": base64.b64encode(
|
|
660
|
+
pickle.dumps([{gpu_uuid: handle}])).decode("utf-8"),
|
|
661
|
+
"num_params": 1,
|
|
662
|
+
"gpu_uuid": gpu_uuid,
|
|
663
|
+
}
|
|
664
|
+
|
|
665
|
+
def save_state_dict_to_path(self, path: str) -> dict:
|
|
666
|
+
"""Save model state dict to a file.
|
|
667
|
+
|
|
668
|
+
Works for both ZeRO-2 (params are full) and ZeRO-3 (params are
|
|
669
|
+
partitions — we just save whatever is local). For ZeRO-3, the
|
|
670
|
+
caller must ensure all ranks call this collectively and only
|
|
671
|
+
rank 0's output is used.
|
|
672
|
+
"""
|
|
673
|
+
t0 = time.monotonic()
|
|
674
|
+
weights = [(n, p.data.cpu()) for n, p in self.engine.module.named_parameters()]
|
|
675
|
+
if self.rank == 0:
|
|
676
|
+
os.makedirs(os.path.dirname(path), exist_ok=True)
|
|
677
|
+
torch.save(weights, path)
|
|
678
|
+
num_params = len(weights)
|
|
679
|
+
del weights
|
|
680
|
+
import gc
|
|
681
|
+
gc.collect()
|
|
682
|
+
elapsed = time.monotonic() - t0
|
|
683
|
+
logger.info("Rank %d saved state dict to %s in %.2fs (%d params)",
|
|
684
|
+
self.rank, path, elapsed, num_params)
|
|
685
|
+
return {"num_params": num_params, "elapsed": elapsed}
|
|
686
|
+
|
|
687
|
+
def gather_and_save_state_dict(self, path: str) -> dict:
|
|
688
|
+
"""Gather ZeRO-3 partitioned params and save full state dict.
|
|
689
|
+
|
|
690
|
+
All ranks must call this collectively. Every parameter is wrapped
|
|
691
|
+
in GatheredParameters so the all-gather runs on all ranks. Only
|
|
692
|
+
rank 0 copies the full tensor and writes to disk.
|
|
693
|
+
"""
|
|
694
|
+
import deepspeed
|
|
695
|
+
t0 = time.monotonic()
|
|
696
|
+
model = self.engine.module
|
|
697
|
+
weights = []
|
|
698
|
+
for n, p in model.named_parameters():
|
|
699
|
+
if hasattr(p, "ds_id"):
|
|
700
|
+
with deepspeed.zero.GatheredParameters([p], enabled=True):
|
|
701
|
+
if self.rank == 0:
|
|
702
|
+
weights.append((n, p.data.cpu()))
|
|
703
|
+
else:
|
|
704
|
+
if self.rank == 0:
|
|
705
|
+
weights.append((n, p.data.cpu()))
|
|
706
|
+
if self.rank == 0:
|
|
707
|
+
os.makedirs(os.path.dirname(path), exist_ok=True)
|
|
708
|
+
torch.save(weights, path)
|
|
709
|
+
num_params = len(weights)
|
|
710
|
+
del weights
|
|
711
|
+
import gc
|
|
712
|
+
gc.collect()
|
|
713
|
+
if hasattr(self.engine, 'empty_partition_cache'):
|
|
714
|
+
self.engine.empty_partition_cache()
|
|
715
|
+
torch.cuda.empty_cache()
|
|
716
|
+
elapsed = time.monotonic() - t0
|
|
717
|
+
logger.info("Rank %d gathered+saved state dict in %.2fs (%d params)",
|
|
718
|
+
self.rank, elapsed, num_params)
|
|
719
|
+
return {"num_params": num_params, "elapsed": elapsed}
|
|
720
|
+
|
|
721
|
+
def _log_mem(self, label):
|
|
722
|
+
from deepspeed.runtime.utils import see_memory_usage
|
|
723
|
+
see_memory_usage(f"[Rank {self.rank}] {label}", force=True)
|
|
724
|
+
|
|
725
|
+
def _ds_offload(self, include):
|
|
726
|
+
"""Offload states using DeepSpeed native API.
|
|
727
|
+
|
|
728
|
+
engine.offload_states works for ZeRO-3/ZeRO-2 without
|
|
729
|
+
offload_optimizer. When offload_optimizer is configured,
|
|
730
|
+
DeepSpeed raises AssertionError ("Moving states across devices
|
|
731
|
+
is not supported"); fall back to optimizer.offload_states which
|
|
732
|
+
bypasses the engine assertion (see DeepSpeed issue #6596).
|
|
733
|
+
"""
|
|
734
|
+
from deepspeed.runtime.zero.config import OffloadDeviceEnum
|
|
735
|
+
try:
|
|
736
|
+
self.engine.offload_states(include=include)
|
|
737
|
+
except AssertionError as e:
|
|
738
|
+
if "Moving states across devices" not in str(e):
|
|
739
|
+
raise
|
|
740
|
+
opt = getattr(self.engine, "optimizer", None)
|
|
741
|
+
if opt is None:
|
|
742
|
+
# Forward-only (inference) engine has no optimizer to fall back
|
|
743
|
+
# to; the engine-level assertion only fires with offload_optimizer.
|
|
744
|
+
raise
|
|
745
|
+
opt.offload_states(
|
|
746
|
+
include=include,
|
|
747
|
+
device=OffloadDeviceEnum.cpu,
|
|
748
|
+
pin_memory=True,
|
|
749
|
+
)
|
|
750
|
+
|
|
751
|
+
def _ds_reload(self):
|
|
752
|
+
"""Reload state to GPU.
|
|
753
|
+
|
|
754
|
+
engine.reload_states works for most configs. When
|
|
755
|
+
offload_optimizer is configured, falls back to optimizer
|
|
756
|
+
directly. ZeRO-2 reload_states requires .grad on fp32 param
|
|
757
|
+
partitions; create zero grads if missing. Forward-only engines
|
|
758
|
+
have no optimizer, so the optimizer-specific handling is skipped.
|
|
759
|
+
"""
|
|
760
|
+
opt = getattr(self.engine, "optimizer", None)
|
|
761
|
+
if opt is not None and hasattr(opt, 'single_partition_of_fp32_groups'):
|
|
762
|
+
for fp32_group in opt.single_partition_of_fp32_groups:
|
|
763
|
+
if fp32_group.grad is None:
|
|
764
|
+
fp32_group.grad = torch.zeros_like(fp32_group)
|
|
765
|
+
try:
|
|
766
|
+
self.engine.reload_states()
|
|
767
|
+
except AssertionError as e:
|
|
768
|
+
if "Moving states across devices" not in str(e):
|
|
769
|
+
raise
|
|
770
|
+
if opt is None:
|
|
771
|
+
raise
|
|
772
|
+
opt.reload_states()
|
|
773
|
+
|
|
774
|
+
def _move_params(self, device) -> None:
|
|
775
|
+
"""Move model parameters to ``device`` for a forward-only engine.
|
|
776
|
+
|
|
777
|
+
DeepSpeed's offload_states/reload_states require a real optimizer, so a
|
|
778
|
+
no-optimizer (inference) engine can't use them. Instead move each
|
|
779
|
+
parameter's storage directly. Under ZeRO-3 the local shard lives in
|
|
780
|
+
``param.ds_tensor``; otherwise it is ``param.data``. This is only ever
|
|
781
|
+
called while the engine is idle (the client wakes it before any forward
|
|
782
|
+
and sleeps it after), so it never races a parameter all-gather.
|
|
783
|
+
"""
|
|
784
|
+
for p in self.engine.module.parameters():
|
|
785
|
+
shard = getattr(p, "ds_tensor", None)
|
|
786
|
+
if shard is not None:
|
|
787
|
+
shard.data = shard.data.to(device, non_blocking=True)
|
|
788
|
+
else:
|
|
789
|
+
p.data = p.data.to(device, non_blocking=True)
|
|
790
|
+
|
|
791
|
+
def offload_to_cpu(self) -> dict:
|
|
792
|
+
"""Offload engine state to CPU using DeepSpeed native API.
|
|
793
|
+
|
|
794
|
+
Trainable engines offload optimizer/grad states plus model params via
|
|
795
|
+
the DeepSpeed offload_states API. A forward-only (no-optimizer) engine
|
|
796
|
+
has no optimizer for that API, so it moves only its model params.
|
|
797
|
+
"""
|
|
798
|
+
if not self._on_gpu:
|
|
799
|
+
return {"status": "already_offloaded"}
|
|
800
|
+
t0 = time.monotonic()
|
|
801
|
+
self._log_mem("before offload_to_cpu")
|
|
802
|
+
|
|
803
|
+
if getattr(self, "_has_optimizer", True):
|
|
804
|
+
from deepspeed.runtime.zero.offload_states import OffloadStateTypeEnum
|
|
805
|
+
self._ds_offload(include=[
|
|
806
|
+
OffloadStateTypeEnum.hp_params,
|
|
807
|
+
OffloadStateTypeEnum.lp_params,
|
|
808
|
+
OffloadStateTypeEnum.lp_grads,
|
|
809
|
+
OffloadStateTypeEnum.contiguous_grad_buffer,
|
|
810
|
+
])
|
|
811
|
+
else:
|
|
812
|
+
self._move_params(self.cpu_device)
|
|
813
|
+
|
|
814
|
+
torch.cuda.synchronize()
|
|
815
|
+
torch.cuda.empty_cache()
|
|
816
|
+
self._on_gpu = False
|
|
817
|
+
elapsed = time.monotonic() - t0
|
|
818
|
+
mem_mb = torch.cuda.memory_allocated() / 1e6
|
|
819
|
+
logger.info("Rank %d offloaded to CPU in %.2fs (%.0f MB GPU remaining)",
|
|
820
|
+
self.rank, elapsed, mem_mb)
|
|
821
|
+
return {"status": "offloaded", "elapsed": elapsed, "gpu_mb": mem_mb}
|
|
822
|
+
|
|
823
|
+
def backload_to_gpu(self) -> dict:
|
|
824
|
+
"""Reload engine state to GPU using DeepSpeed native API.
|
|
825
|
+
|
|
826
|
+
Forward-only engines move only their model params back (no optimizer
|
|
827
|
+
state to reload)."""
|
|
828
|
+
if self._on_gpu:
|
|
829
|
+
return {"status": "already_on_gpu"}
|
|
830
|
+
t0 = time.monotonic()
|
|
831
|
+
self._log_mem("before reload_states")
|
|
832
|
+
|
|
833
|
+
if getattr(self, "_has_optimizer", True):
|
|
834
|
+
self._ds_reload()
|
|
835
|
+
else:
|
|
836
|
+
self._move_params(torch.device(self._device))
|
|
837
|
+
torch.cuda.empty_cache()
|
|
838
|
+
self._log_mem("after reload_states + empty_cache")
|
|
839
|
+
|
|
840
|
+
torch.cuda.synchronize()
|
|
841
|
+
self._on_gpu = True
|
|
842
|
+
elapsed = time.monotonic() - t0
|
|
843
|
+
logger.info("Rank %d backloaded to GPU in %.2fs", self.rank, elapsed)
|
|
844
|
+
return {"status": "on_gpu", "elapsed": elapsed}
|
|
845
|
+
|
|
846
|
+
def offload_non_lp_states(self) -> dict:
|
|
847
|
+
"""Offload everything except bf16 params (for CUDA IPC sync)."""
|
|
848
|
+
t0 = time.monotonic()
|
|
849
|
+
self._log_mem("before offload_non_lp")
|
|
850
|
+
|
|
851
|
+
from deepspeed.runtime.zero.offload_states import OffloadStateTypeEnum
|
|
852
|
+
self._ds_offload(include=[
|
|
853
|
+
OffloadStateTypeEnum.hp_params,
|
|
854
|
+
OffloadStateTypeEnum.lp_grads,
|
|
855
|
+
OffloadStateTypeEnum.contiguous_grad_buffer,
|
|
856
|
+
])
|
|
857
|
+
|
|
858
|
+
torch.cuda.synchronize()
|
|
859
|
+
torch.cuda.empty_cache()
|
|
860
|
+
elapsed = time.monotonic() - t0
|
|
861
|
+
mem_mb = torch.cuda.memory_allocated() / 1e6
|
|
862
|
+
logger.info("Rank %d offloaded non-lp states in %.2fs (%.0f MB GPU)",
|
|
863
|
+
self.rank, elapsed, mem_mb)
|
|
864
|
+
return {"status": "offloaded_non_lp", "elapsed": elapsed, "gpu_mb": mem_mb}
|
|
865
|
+
|
|
866
|
+
def offload_lp_params(self) -> dict:
|
|
867
|
+
"""Offload bf16 model params to CPU (after CUDA IPC sync)."""
|
|
868
|
+
t0 = time.monotonic()
|
|
869
|
+
|
|
870
|
+
from deepspeed.runtime.zero.offload_states import OffloadStateTypeEnum
|
|
871
|
+
self._ds_offload(include=[OffloadStateTypeEnum.lp_params])
|
|
872
|
+
|
|
873
|
+
torch.cuda.synchronize()
|
|
874
|
+
torch.cuda.empty_cache()
|
|
875
|
+
self._on_gpu = False
|
|
876
|
+
elapsed = time.monotonic() - t0
|
|
877
|
+
mem_mb = torch.cuda.memory_allocated() / 1e6
|
|
878
|
+
logger.info("Rank %d offloaded lp_params in %.2fs (%.0f MB GPU remaining)",
|
|
879
|
+
self.rank, elapsed, mem_mb)
|
|
880
|
+
return {"status": "offloaded_lp", "elapsed": elapsed, "gpu_mb": mem_mb}
|
|
881
|
+
|
|
882
|
+
def empty_cache(self) -> dict:
|
|
883
|
+
"""Release ZeRO-3 partition cache and PyTorch cached memory."""
|
|
884
|
+
if hasattr(self.engine, 'empty_partition_cache'):
|
|
885
|
+
self.engine.empty_partition_cache()
|
|
886
|
+
import gc
|
|
887
|
+
gc.collect()
|
|
888
|
+
torch.cuda.empty_cache()
|
|
889
|
+
mem_mb = torch.cuda.memory_allocated() / 1e6
|
|
890
|
+
logger.info("Rank %d empty_cache: %.0f MB GPU remaining", self.rank, mem_mb)
|
|
891
|
+
return {"gpu_mb": mem_mb}
|
|
892
|
+
|
|
893
|
+
def destroy(self) -> bool:
|
|
894
|
+
if self._weight_sender is not None:
|
|
895
|
+
self._weight_sender.destroy()
|
|
896
|
+
self._weight_sender = None
|
|
897
|
+
if dist.is_initialized():
|
|
898
|
+
dist.destroy_process_group()
|
|
899
|
+
self.engine = None
|
|
900
|
+
return True
|
|
901
|
+
|
|
902
|
+
|
|
903
|
+
# ---------------------------------------------------------------------------
|
|
904
|
+
# Helpers
|
|
905
|
+
# ---------------------------------------------------------------------------
|