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.
Files changed (56) hide show
  1. arctic_platform/rl/README.md +7 -0
  2. arctic_platform/rl/__init__.py +27 -0
  3. arctic_platform/rl/client.py +46 -0
  4. arctic_platform/rl/config.py +150 -0
  5. arctic_platform/rl/deepspeed_worker.py +905 -0
  6. arctic_platform/rl/examples/inference_example.py +115 -0
  7. arctic_platform/rl/examples/local_server_example.py +201 -0
  8. arctic_platform/rl/examples/weight_sync_example.py +203 -0
  9. arctic_platform/rl/http_client.py +705 -0
  10. arctic_platform/rl/http_server.py +1166 -0
  11. arctic_platform/rl/processors/__init__.py +166 -0
  12. arctic_platform/rl/processors/functional.py +257 -0
  13. arctic_platform/rl/processors/grpo.py +330 -0
  14. arctic_platform/rl/processors/microbatch.py +176 -0
  15. arctic_platform/rl/processors/packing.py +90 -0
  16. arctic_platform/rl/processors/pipeline.py +880 -0
  17. arctic_platform/rl/processors/stats_tracker.py +186 -0
  18. arctic_platform/rl/processors/verl_grpo.py +479 -0
  19. arctic_platform/rl/projects/long_context_qa/README.md +87 -0
  20. arctic_platform/rl/projects/long_context_qa/download_data.py +120 -0
  21. arctic_platform/rl/projects/long_context_qa/run_qwen3_32b_longcontext_grpo_arl_zorro_yes_kl.sh +221 -0
  22. arctic_platform/rl/projects/txt2sql/README.md +176 -0
  23. arctic_platform/rl/projects/txt2sql/bird_reward.py +272 -0
  24. arctic_platform/rl/projects/txt2sql/preprocess_bird.py +684 -0
  25. arctic_platform/rl/projects/txt2sql/run_qwen3_32b_bird_grpo_arl_zorro_yes.sh +251 -0
  26. arctic_platform/rl/projects/txt2sql/run_qwen3_32b_bird_grpo_arl_zorro_yes_kl.sh +31 -0
  27. arctic_platform/rl/ray_client.py +482 -0
  28. arctic_platform/rl/ray_cluster.py +213 -0
  29. arctic_platform/rl/ray_server.py +1475 -0
  30. arctic_platform/rl/server.py +6 -0
  31. arctic_platform/rl/utils/__init__.py +16 -0
  32. arctic_platform/rl/utils/batch.py +434 -0
  33. arctic_platform/rl/utils/cuda_ipc.py +57 -0
  34. arctic_platform/rl/utils/debug.py +495 -0
  35. arctic_platform/rl/utils/ray_pg.py +210 -0
  36. arctic_platform/rl/weight_sync.py +238 -0
  37. arctic_platform/rl/zorro_train/README.md +308 -0
  38. arctic_platform/rl/zorro_train/__init__.py +23 -0
  39. arctic_platform/rl/zorro_train/actor.py +271 -0
  40. arctic_platform/rl/zorro_train/demo.py +130 -0
  41. arctic_platform/rl/zorro_train/module_patcher.py +67 -0
  42. arctic_platform/rl/zorro_train/qwen_attention_patcher.py +1071 -0
  43. arctic_platform/rl/zorro_train/qwen_model_patcher.py +1107 -0
  44. arctic_platform/rl/zorro_train/seqlen_balancing.py +917 -0
  45. arctic_platform/rl/zorro_train/test.py +7 -0
  46. arctic_platform/rl/zorro_train/test_actor_demo.py +68 -0
  47. arctic_platform/rl/zorro_train/test_forward_and_backward.py +492 -0
  48. arctic_platform/rl/zorro_train/test_perf-g1.sh +2 -0
  49. arctic_platform/rl/zorro_train/test_perf-g8.sh +2 -0
  50. arctic_platform/rl/zorro_train/test_perf.py +503 -0
  51. arctic_platform/rl/zorro_train/tests.py +711 -0
  52. arctic_platform/rl/zorro_train/zorro_train.py +2670 -0
  53. arctic_platform-0.1.1.dev0.dist-info/METADATA +84 -0
  54. arctic_platform-0.1.1.dev0.dist-info/RECORD +56 -0
  55. arctic_platform-0.1.1.dev0.dist-info/WHEEL +4 -0
  56. 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
+ # ---------------------------------------------------------------------------