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,7 @@
1
+ # Arctic RL
2
+
3
+ TODO: add quick start, etc.
4
+
5
+ Here are a few recipes using Arctic RL:
6
+ * [Txt2SQL](projects/txt2sql)
7
+ * ...
@@ -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