Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
70 commits
Select commit Hold shift + click to select a range
d399b0f
Integrate ArcticRL
sfc-gh-truwase Apr 2, 2026
844bf72
ArcticRL integgration
sfc-gh-truwase Apr 2, 2026
cadd1ee
Revert changes (#7)
sfc-gh-truwase Apr 2, 2026
6dec5a2
Merge ARL WIP
sfc-gh-truwase Apr 2, 2026
81f9837
Disable CI
sfc-gh-truwase Apr 2, 2026
d1786f5
Disable CI
sfc-gh-truwase Apr 2, 2026
6d52ea9
integration with AT ARLClient
sfc-gh-mwyatt Mar 31, 2026
a3cd7ca
align with recent changes from Karthik
sfc-gh-mwyatt Apr 1, 2026
276f5bf
add TODO
sfc-gh-mwyatt Apr 1, 2026
64cdff5
revert changes
sfc-gh-mwyatt Apr 2, 2026
8331cdc
remove unused imports
sfc-gh-mwyatt Apr 2, 2026
b1e0709
train_mini_natch -> train_global_batch
sfc-gh-truwase Apr 3, 2026
e01c9da
Merge ZoRRO work (#13)
sfc-gh-truwase Apr 3, 2026
6b1889e
Debugging multi-gpu
sfc-gh-truwase Apr 3, 2026
d271b2b
Merge branch 'arl' of https://github.com/snowflake-eng/arctic-verl in…
sfc-gh-truwase Apr 3, 2026
2a15fdc
missing
sfc-gh-sbekman Apr 3, 2026
eec0399
Merge branch 'arl' of https://github.com/snowflake-eng/arctic-verl in…
sfc-gh-truwase Apr 3, 2026
2cb52cb
Merge branch 'arl' into mwyatt/ARL-integration
sfc-gh-mwyatt Apr 3, 2026
928b66b
update to work with recent arl changes
sfc-gh-mwyatt Apr 3, 2026
f9ba22e
Multi-GPU
sfc-gh-truwase Apr 6, 2026
c3ce480
Merge remote-tracking branch 'origin/mwyatt/ARL-integration' into arl
sfc-gh-truwase Apr 6, 2026
febd285
merge
sfc-gh-sbekman Apr 6, 2026
0363c35
Merge branch 'arl' of https://github.com/snowflake-eng/arctic-verl in…
sfc-gh-truwase Apr 7, 2026
0c0d613
use_zorro env var; test zorro
sfc-gh-truwase Apr 7, 2026
8ec8c78
zorro -> log_prob
sfc-gh-sbekman Apr 7, 2026
479e7d5
WIP
sfc-gh-truwase Apr 7, 2026
8e69683
Merge branch 'arl' of https://github.com/snowflake-eng/arctic-verl in…
sfc-gh-truwase Apr 7, 2026
f793988
ArcticTraining ARLClient integration (#9)
sfc-gh-mwyatt Apr 7, 2026
bfae33f
pass temp
sfc-gh-sbekman Apr 7, 2026
5a215b8
helper
sfc-gh-sbekman Apr 7, 2026
d5bf532
Merge branch 'arl' of https://github.com/snowflake-eng/arctic-verl in…
sfc-gh-sbekman Apr 7, 2026
52c7fd4
Rebase
sfc-gh-truwase Apr 7, 2026
cde6f3b
Merge branch 'arl' of https://github.com/snowflake-eng/arctic-verl in…
sfc-gh-truwase Apr 7, 2026
27e3ada
Remove utility
sfc-gh-truwase Apr 8, 2026
dd9719e
Cleanup
sfc-gh-truwase Apr 8, 2026
b88967a
removing train batch nesting
sfc-gh-sbekman Apr 9, 2026
aa3fc1b
ref_log_prob optional
sfc-gh-truwase Apr 9, 2026
c75c667
Rebase
sfc-gh-truwase Apr 9, 2026
810b6f5
fix for ArcticTraining-dss API update (#15)
sfc-gh-mwyatt Apr 10, 2026
9c1296a
Cleaning up
sfc-gh-truwase Apr 10, 2026
7e41e46
Merge branch 'arl' of https://github.com/snowflake-eng/arctic-verl in…
sfc-gh-truwase Apr 10, 2026
44393b5
dataloader seed
sfc-gh-truwase Apr 10, 2026
bd779d2
W&B
sfc-gh-truwase Apr 10, 2026
2a4403f
split scripts
sfc-gh-sbekman Apr 10, 2026
d6f6aba
text2sql recipe
sfc-gh-truwase Apr 10, 2026
670626d
ARL integration
sfc-gh-truwase Apr 14, 2026
fd1e63d
Fix multi-gpu mismatch for train/ & log_prob
sfc-gh-truwase Apr 15, 2026
f4fc13f
Add setup
sfc-gh-truwase Apr 15, 2026
370e6aa
integrate zorro
sfc-gh-sbekman Apr 16, 2026
7a8b825
Fix batch_size to include rollout_n
sfc-gh-truwase Apr 16, 2026
ac63345
use ARL by default now
sfc-gh-sbekman Apr 16, 2026
9a60b7a
Async generate
sfc-gh-truwase Apr 16, 2026
35ea0a5
Gas support
sfc-gh-truwase Apr 17, 2026
5ce5fba
new launchers and some small cleanup in the old ones
sfc-gh-sbekman Apr 18, 2026
48ac315
Cleanup; arl yaml config
sfc-gh-truwase Apr 21, 2026
6f3713d
Remove temp arl client
sfc-gh-truwase Apr 21, 2026
62bdb5b
Disable CI (#10)
sfc-gh-truwase Apr 22, 2026
8043033
Removing TrainingWorker
sfc-gh-truwase Apr 23, 2026
0e60e10
Cleanup
sfc-gh-truwase Apr 23, 2026
436f9db
Merge branch 'main' of https://github.com/snowflake-eng/arctic-verl i…
sfc-gh-truwase Apr 23, 2026
ca9c59d
Cleanup
sfc-gh-truwase Apr 23, 2026
355efd7
Cleanup
sfc-gh-truwase Apr 23, 2026
b930595
Cleanup
sfc-gh-truwase Apr 23, 2026
4a31ea0
Cleanup
sfc-gh-truwase Apr 23, 2026
6b9e257
Add bird run script
sfc-gh-truwase Apr 23, 2026
c723dc3
Enable weight sync
sfc-gh-truwase Apr 23, 2026
b8c13bb
Remove
sfc-gh-truwase Apr 28, 2026
4f11e10
Add example
sfc-gh-truwase Apr 28, 2026
0f9e1c6
Restore CI
sfc-gh-truwase Apr 28, 2026
20ca9dd
Restore CI
sfc-gh-truwase Apr 28, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 87 additions & 0 deletions examples/arctic_rl/run_gsm8k_grpo_arl_zorro_yes.sh
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
#!/bin/bash

set -x

export PYTHONUNBUFFERED=1
export HYDRA_FULL_ERROR=1
export RAY_DEDUP_LOGS=0
export HF_HUB_OFFLINE=1
export HF_HOME=/checkpoint/huggingface
# we want to make sure this runs on non-gpu client
export CUDA_VISIBLE_DEVICES=

BSZ=1024
UBS=32
ROLL_N=5
PROMPT_LENGTH=512
RESPONSE_LENGTH=1024
MAX_STEPS=100


BSZ=8
UBS=2
ROLL_N=2
PROMPT_LENGTH=512
RESPONSE_LENGTH=1024
MAX_STEPS=4


experiment_name="qwen3-0.6B_arctic_gsm8k_grpo"

python3 -m verl.trainer.main_ppo \
algorithm.adv_estimator=grpo \
data.train_files=/code/shared/gsm8k/train.parquet \
data.val_files=/code/shared/gsm8k/test.parquet \
data.train_batch_size=$BSZ \
data.max_prompt_length=512 \
data.max_response_length=1024 \
data.filter_overlong_prompts=True \
data.truncation='error' \
data.shuffle=False \
+data.seed=42 \
actor_rollout_ref.actor.data_loader_seed=42 \
reward.num_workers=1 \
actor_rollout_ref.rollout.agent.num_workers=4 \
actor_rollout_ref.model.path=Qwen/Qwen3-0.6B \
actor_rollout_ref.actor.optim.lr=1e-6 \
actor_rollout_ref.model.use_remove_padding=False \
actor_rollout_ref.actor.ppo_mini_batch_size=$BSZ \
actor_rollout_ref.actor.ppo_micro_batch_size_per_gpu=$UBS \
actor_rollout_ref.actor.use_kl_loss=False \
actor_rollout_ref.actor.kl_loss_coef=0.001 \
actor_rollout_ref.actor.kl_loss_type=low_var_kl \
actor_rollout_ref.actor.entropy_coeff=0 \
actor_rollout_ref.model.enable_gradient_checkpointing=True \
+actor_rollout_ref.model.override_config.attn_implementation=flash_attention_3 \
actor_rollout_ref.actor.strategy=fsdp2 \
actor_rollout_ref.actor.fsdp_config.param_offload=False \
actor_rollout_ref.actor.fsdp_config.optimizer_offload=False \
actor_rollout_ref.rollout.log_prob_micro_batch_size_per_gpu=$UBS \
actor_rollout_ref.rollout.tensor_model_parallel_size=1 \
actor_rollout_ref.rollout.name=arctic \
actor_rollout_ref.rollout.gpu_memory_utilization=0.6 \
actor_rollout_ref.rollout.enforce_eager=True \
actor_rollout_ref.rollout.n=$ROLL_N \
actor_rollout_ref.ref.log_prob_micro_batch_size_per_gpu=$UBS \
actor_rollout_ref.ref.fsdp_config.param_offload=False \
actor_rollout_ref.ref.strategy=fsdp2 \
algorithm.use_kl_in_reward=False \
trainer.use_legacy_worker_impl=disable \
trainer.use_arctic_rl=True \
trainer.critic_warmup=0 \
trainer.logger=console \
trainer.experiment_name=$experiment_name \
trainer.project_name=arctic_gsm8k_grpo \
trainer.val_before_train=False \
trainer.n_gpus_per_node=1 \
trainer.nnodes=1 \
trainer.save_freq=-1 \
trainer.test_freq=-1 \
trainer.total_epochs=15 \
trainer.total_training_steps=${MAX_STEPS} \
arctic_rl.colocate=False \
arctic_rl.training_gpus=2\
arctic_rl.sampling_gpus=2\
arctic_rl.log_prob_gpus=0\
arctic_rl.use_zorro=True \
"$@" 2>&1 | tee $experiment_name.log
9 changes: 6 additions & 3 deletions verl/experimental/agent_loop/agent_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -918,13 +918,14 @@ def __init__(
worker_group: RayWorkerGroup = None,
rollout_resource_pool: RayResourcePool = None,
reward_loop_worker_handles: list[ray.actor.ActorHandle] = None,
):
**kwargs,
):
self.config = config
self.rollout_config, self.model_config = _get_rollout_and_model_config(config)
self.worker_group = worker_group
self.rollout_resource_pool = rollout_resource_pool
self.reward_loop_worker_handles = reward_loop_worker_handles

self.kwargs = kwargs
assert worker_group is not None or self.rollout_config.nnodes > 0, "nnodes must be > 0 in standalone mode"

# for recipe to change
Expand All @@ -941,9 +942,10 @@ async def create(
worker_group: RayWorkerGroup = None,
rollout_resource_pool: RayResourcePool = None,
reward_loop_worker_handles: list[ray.actor.ActorHandle] = None,
**kwargs,
):
"""Create agent loop manager."""
instance = cls(config, worker_group, rollout_resource_pool, reward_loop_worker_handles)
instance = cls(config, worker_group, rollout_resource_pool, reward_loop_worker_handles, **kwargs)
await instance._initialize_llm_servers()
await instance._init_global_load_balancer()
await instance._init_agent_loop_workers()
Expand All @@ -968,6 +970,7 @@ async def _initialize_llm_servers(self):
config=self.rollout_config,
model_config=self.model_config,
gpus_per_node=self.rollout_config.n_gpus_per_node,
**self.kwargs,
)
for replica_rank in range(num_replicas)
]
Expand Down
15 changes: 8 additions & 7 deletions verl/single_controller/ray/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -187,8 +187,9 @@ class ResourcePoolManager:
resource_pool_spec: dict[str, list[int]]
mapping: dict[int, str]
resource_pool_dict: dict[str, RayResourcePool] = field(default_factory=dict)
gpu_resource_pool_dict: dict[str, RayResourcePool] = field(default_factory=dict)

def create_resource_pool(self):
def create_resource_pool(self, use_gpu: bool = True):
"""Create Ray resource pools for distributed training.

Initializes resource pools based on the resource pool specification,
Expand All @@ -202,10 +203,11 @@ def create_resource_pool(self):
# For Megatron backend, we recommend using max_colocate_count>1
# that can utilize different WorkerGroup for differnt models
resource_pool = RayResourcePool(
process_on_nodes=process_on_nodes, use_gpu=True, max_colocate_count=3, name_prefix=resource_pool_name
process_on_nodes=process_on_nodes, use_gpu=use_gpu, max_colocate_count=3, name_prefix=resource_pool_name
)
self.resource_pool_dict[resource_pool_name] = resource_pool

if use_gpu:
self.gpu_resource_pool_dict[resource_pool_name] = resource_pool
self._check_resource_available()

def get_resource_pool(self, role) -> RayResourcePool:
Expand All @@ -214,7 +216,8 @@ def get_resource_pool(self, role) -> RayResourcePool:

def get_n_gpus(self) -> int:
"""Get the number of gpus in this cluster."""
return sum([n_gpus for process_on_nodes in self.resource_pool_spec.values() for n_gpus in process_on_nodes])
process_on_gpu_nodes = [process_on_nodes for pool_name, process_on_nodes in self.resource_pool_spec.items() if pool_name in self.gpu_resource_pool_dict]
return sum([n_gpus for process_on_nodes in process_on_gpu_nodes for n_gpus in process_on_nodes])

def _check_resource_available(self):
"""Check if the resource pool can be satisfied in this ray cluster."""
Expand All @@ -226,9 +229,7 @@ def _check_resource_available(self):

# check total required gpus can be satisfied
total_available_gpus = sum(node_available_gpus.values())
total_required_gpus = sum(
[n_gpus for process_on_nodes in self.resource_pool_spec.values() for n_gpus in process_on_nodes]
)
total_required_gpus = self.get_n_gpus()
if total_available_gpus < total_required_gpus:
raise ValueError(
f"Total available GPUs {total_available_gpus} is less than total desired GPUs {total_required_gpus}"
Expand Down
16 changes: 16 additions & 0 deletions verl/trainer/config/ppo_trainer.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -203,6 +203,9 @@ trainer:
# mode: "auto", "enable", or "disable"
use_legacy_worker_impl: auto

# whether to use arctic rl
use_arctic_rl: False

# profiler configs
global_profiler:

Expand Down Expand Up @@ -310,3 +313,16 @@ ray_kwargs:

# Path to save Ray timeline JSON for performance profiling
timeline_json_file: null


# config for arctic rl
arctic_rl:

# whether to use colocate mode
colocate: False

training_gpus: 1
sampling_gpus: 1
log_prob_gpus: 1

use_zorro: False
15 changes: 13 additions & 2 deletions verl/trainer/main_ppo.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,10 @@ def add_actor_rollout_worker(self, config):
actor_rollout_cls = ActorRolloutRefWorker
ray_worker_group_cls = RayWorkerGroup

if config.trainer.get("use_arctic_rl", False):
from verl.workers.arctic_workers import ActorRolloutRefWorker
actor_rollout_cls = ActorRolloutRefWorker

lora_rank = config.actor_rollout_ref.model.get("lora", {}).get("rank", 0)
if lora_rank <= 0:
lora_rank = config.actor_rollout_ref.model.get("lora_rank", 0)
Expand Down Expand Up @@ -340,7 +344,9 @@ def run(self, config):
train_sampler = create_rl_sampler(config.data, train_dataset)

# Initialize the PPO trainer.
trainer = RayPPOTrainer(
from verl.trainer.ppo.arctic_trainer import ArcticPPOTrainer
ppo_trainer_cls = RayPPOTrainer if not config.trainer.use_arctic_rl else ArcticPPOTrainer
trainer = ppo_trainer_cls(
config=config,
tokenizer=tokenizer,
processor=processor,
Expand All @@ -356,7 +362,12 @@ def run(self, config):
trainer.init_workers()

# Start the training process.
trainer.fit()
try:
trainer.fit()
finally:
# Ensure remote services shutdown gracefully
if hasattr(trainer, "destroy"):
trainer.destroy()


def create_rl_dataset(data_paths, data_config, tokenizer, processor, is_train=True, max_samples: int = -1):
Expand Down
Loading
Loading