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,27 @@
|
|
|
1
|
+
"""Arctic RL client -- HTTP client for RL training against dss-platform or local server."""
|
|
2
|
+
|
|
3
|
+
from arctic_platform.rl.client import create_arctic_rl_client
|
|
4
|
+
from arctic_platform.rl.config import ArcticRLClientConfig
|
|
5
|
+
from arctic_platform.rl.config import WeightSyncConfig
|
|
6
|
+
from arctic_platform.rl.processors import (
|
|
7
|
+
grpo_loss,
|
|
8
|
+
pack_sequences,
|
|
9
|
+
unpack_sequences,
|
|
10
|
+
register_loss_fn,
|
|
11
|
+
register_post_processor,
|
|
12
|
+
run_pipeline,
|
|
13
|
+
)
|
|
14
|
+
from arctic_platform.rl.weight_sync import WeightSyncCoordinator
|
|
15
|
+
|
|
16
|
+
__all__ = [
|
|
17
|
+
"create_arctic_rl_client",
|
|
18
|
+
"ArcticRLClientConfig",
|
|
19
|
+
"WeightSyncConfig",
|
|
20
|
+
"WeightSyncCoordinator",
|
|
21
|
+
"run_pipeline",
|
|
22
|
+
"register_loss_fn",
|
|
23
|
+
"register_post_processor",
|
|
24
|
+
"grpo_loss",
|
|
25
|
+
"pack_sequences",
|
|
26
|
+
"unpack_sequences",
|
|
27
|
+
]
|
|
@@ -0,0 +1,46 @@
|
|
|
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
|
+
"""ArcticRLClient -- a unified frontend client for HTTP and Ray clients for RL training.
|
|
17
|
+
|
|
18
|
+
Works identically against a remote dss-platform deployment or a local
|
|
19
|
+
``server.py`` instance -- the only differences are ``base_url`` and whether the
|
|
20
|
+
client launches the server.
|
|
21
|
+
|
|
22
|
+
All jobs (training, sampling, log-prob) are initialized automatically at
|
|
23
|
+
construction time.
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
from __future__ import annotations
|
|
27
|
+
|
|
28
|
+
import logging
|
|
29
|
+
|
|
30
|
+
from arctic_platform.rl.config import ArcticRLClientConfig
|
|
31
|
+
#from arctic_platform.rl.ray_server import ArcticRLRayServerState
|
|
32
|
+
from arctic_platform.rl.server import ArcticRLServerState
|
|
33
|
+
|
|
34
|
+
logger = logging.getLogger(__name__)
|
|
35
|
+
|
|
36
|
+
from arctic_platform.rl.http_client import ArcticRLHTTPClient
|
|
37
|
+
from arctic_platform.rl.ray_client import ArcticRLRayClient
|
|
38
|
+
|
|
39
|
+
def create_arctic_rl_client(config: ArcticRLClientConfig, arctic_rl_server_state: ArcticRLServerState = None):
|
|
40
|
+
if config.comm_protocol == "http":
|
|
41
|
+
return ArcticRLHTTPClient(config)
|
|
42
|
+
elif config.comm_protocol == "ray":
|
|
43
|
+
#assert arctic_rl_server_state is not None, "arctic_rl_server_state is required for comm_protocol: ray"
|
|
44
|
+
return ArcticRLRayClient(config, arctic_rl_server_state)
|
|
45
|
+
else:
|
|
46
|
+
raise ValueError(f"Invalid communication protocol: {config.comm_protocol}")
|
|
@@ -0,0 +1,150 @@
|
|
|
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
|
+
"""Configuration models for the Arctic RL client."""
|
|
17
|
+
|
|
18
|
+
from __future__ import annotations
|
|
19
|
+
|
|
20
|
+
from typing import Literal
|
|
21
|
+
from typing import Optional
|
|
22
|
+
|
|
23
|
+
from pydantic import BaseModel
|
|
24
|
+
from pydantic import Field
|
|
25
|
+
from pydantic import model_validator
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class ArcticRLClientConfig(BaseModel):
|
|
29
|
+
backend: Literal["local", "dss-platform"] = "local"
|
|
30
|
+
comm_protocol: Literal["http", "ray"] = "http"
|
|
31
|
+
checkpoint_path: Optional[str] = None
|
|
32
|
+
|
|
33
|
+
# it's best not to pass explicitly the host and port since they are auto derived from comm_protocol
|
|
34
|
+
host: Optional[str] = None
|
|
35
|
+
port: Optional[int] = None
|
|
36
|
+
|
|
37
|
+
model_name: str = Field(description="Model name or HuggingFace ID to load on all engines.")
|
|
38
|
+
ds_config: dict = Field(default_factory=dict, description="DeepSpeed config for training engine.")
|
|
39
|
+
|
|
40
|
+
# Generic training worker config — any framework uses this to configure the server's
|
|
41
|
+
# DeepSpeed training engine (optimizer, dtype, gradient_checkpointing, etc.).
|
|
42
|
+
# Standard fields: optimizer (lr, weight_decay, beta1, beta2, eps, lr_scheduler_type,
|
|
43
|
+
# gradient_clipping, warmup_steps_proportion / warmup_ratio), dtype, gradient_checkpointing,
|
|
44
|
+
# attn_impl, mb_spec (max_tokens_per_mb). Extra fields are ignored by the server.
|
|
45
|
+
training_config: Optional[dict] = Field(
|
|
46
|
+
default=None, description="Training worker config dict (optimizer, dtype, etc.)."
|
|
47
|
+
)
|
|
48
|
+
vllm_config: Optional[dict] = Field(
|
|
49
|
+
default=None, description="vLLM / ModelConfig overrides for sampling and log-prob engines."
|
|
50
|
+
)
|
|
51
|
+
log_prob_ds_config: Optional[dict] = Field(
|
|
52
|
+
default=None, description="Log-prob DeepSpeed worker config dict (batch size, dtype, etc.)."
|
|
53
|
+
)
|
|
54
|
+
ds_worker_config: Optional[dict] = Field(
|
|
55
|
+
default=None, description="Deepspeed worker config dict (optimizer, dtype, etc.)."
|
|
56
|
+
)
|
|
57
|
+
use_arctic_inference: bool = Field(
|
|
58
|
+
default=False,
|
|
59
|
+
description=(
|
|
60
|
+
"If True, set ARCTIC_INFERENCE_ENABLED=1 "
|
|
61
|
+
"and override VLLM_DISABLE_COMPILE_CACHE=0."
|
|
62
|
+
),
|
|
63
|
+
)
|
|
64
|
+
full_determinism: bool = Field(
|
|
65
|
+
default=False,
|
|
66
|
+
description="If True, the DeepSpeed worker calls enable_full_determinism for reproducible training.",
|
|
67
|
+
)
|
|
68
|
+
seed: int = Field(default=42, description="Seed used by enable_full_determinism when full_determinism=True.")
|
|
69
|
+
|
|
70
|
+
training_gpus: int = Field(default=0, description="Number of GPUs for the DeepSpeed training engine.")
|
|
71
|
+
sampling_gpus: int = Field(default=0, description="Number of GPUs for the vLLM sampling engine.")
|
|
72
|
+
log_prob_gpus: int = Field(default=0, description="Number of GPUs for the log-prob engine.")
|
|
73
|
+
log_prob_engine: Literal["vllm", "deepspeed"] = Field(
|
|
74
|
+
default="vllm", description="Engine backend for the log-prob job."
|
|
75
|
+
)
|
|
76
|
+
colocate: bool = Field(
|
|
77
|
+
default=False,
|
|
78
|
+
description="Colocate training, sampling, and log-prob workers on the same GPUs using fractional Ray resources.",
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
server_logs: bool = Field(default=True, description="Show server subprocess stdout/stderr.")
|
|
82
|
+
|
|
83
|
+
ray_auto_attach: bool = Field(
|
|
84
|
+
default=True,
|
|
85
|
+
description=(
|
|
86
|
+
"If True, the local server will attempt to attach to a pre-existing Ray cluster"
|
|
87
|
+
" (only honored when that cluster has GPU resources). Set to False to always start"
|
|
88
|
+
" a fresh Ray cluster — useful when an unrelated CPU-only Ray cluster is running."
|
|
89
|
+
),
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
startup_timeout: float = Field(
|
|
93
|
+
default=300.0, description="Seconds to wait for the local server to become healthy."
|
|
94
|
+
)
|
|
95
|
+
health_check_interval: float = Field(default=2.0, description="Seconds between health-check polls during startup.")
|
|
96
|
+
# How long to wait for each job to reach RUNNING state after /initialize.
|
|
97
|
+
job_ready_timeout: float = Field(
|
|
98
|
+
default=600.0, description="Seconds to wait for each job to become RUNNING after initialization."
|
|
99
|
+
)
|
|
100
|
+
|
|
101
|
+
# Reconnect fields — when set, ArcticRLClient skips /initialize and connects
|
|
102
|
+
# to pre-existing jobs. Populated by ArcticRLClient.reconnect_config() and
|
|
103
|
+
# consumed in ArcticRLClient.__init__. Not forwarded to /initialize.
|
|
104
|
+
training_job_id: Optional[int] = Field(default=None, exclude=True)
|
|
105
|
+
sampling_job_id: Optional[int] = Field(default=None, exclude=True)
|
|
106
|
+
log_prob_job_id: Optional[int] = Field(default=None, exclude=True)
|
|
107
|
+
|
|
108
|
+
@model_validator(mode="after")
|
|
109
|
+
def _derive_host_port(self) -> "ArcticRLClientConfig":
|
|
110
|
+
"""Derive host/port from comm_protocol unless explicitly provided.
|
|
111
|
+
|
|
112
|
+
ray comms don't use host/port (both None). http binds the RL server on
|
|
113
|
+
this node's routable IP at port 7000 so off-node Ray workers can reach
|
|
114
|
+
the driver node by IP rather than "localhost". Values passed explicitly
|
|
115
|
+
by the caller are left untouched (e.g. reconnecting to a known server).
|
|
116
|
+
"""
|
|
117
|
+
# Lazy import to avoid pulling ray in at config import time.
|
|
118
|
+
from arctic_platform.rl.ray_cluster import primary_ip
|
|
119
|
+
|
|
120
|
+
if "host" not in self.model_fields_set:
|
|
121
|
+
self.host = None if self.comm_protocol == "ray" else primary_ip()
|
|
122
|
+
if "port" not in self.model_fields_set:
|
|
123
|
+
self.port = None if self.comm_protocol == "ray" else 7000
|
|
124
|
+
return self
|
|
125
|
+
|
|
126
|
+
@model_validator(mode="after")
|
|
127
|
+
def _validate_local_gpu_counts(self) -> "ArcticRLClientConfig":
|
|
128
|
+
if self.backend != "local" or self.training_job_id is not None:
|
|
129
|
+
return self # skip validation in reconnect mode
|
|
130
|
+
for field in ("training_gpus", "sampling_gpus"):
|
|
131
|
+
if getattr(self, field) <= 0:
|
|
132
|
+
raise ValueError(f"Local backend requires {field} > 0.")
|
|
133
|
+
return self
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
class WeightSyncConfig(BaseModel):
|
|
137
|
+
"""NCCL weight-transfer topology between training GPUs and inference replicas.
|
|
138
|
+
|
|
139
|
+
Used by :class:`WeightSyncCoordinator` (standalone, not part of the HTTP client).
|
|
140
|
+
"""
|
|
141
|
+
|
|
142
|
+
training_sharding: str = Field(
|
|
143
|
+
default="dp",
|
|
144
|
+
description="Training parallelism strategy: 'dp' or 'fsdp'",
|
|
145
|
+
)
|
|
146
|
+
training_gpus: int = 1
|
|
147
|
+
inference_replicas: int = 1
|
|
148
|
+
inference_tp: int = 1
|
|
149
|
+
base_port: int = 29500
|
|
150
|
+
bucket_size: int = 256 * 1024 * 1024 # 256 MB
|