Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 commits
Commits
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
6 changes: 4 additions & 2 deletions torchtitan/experiments/rl/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -71,7 +71,9 @@ def my_experiment() -> Controller.Config:
return Controller.Config(
model_spec=...,
rollouter=MyRollouter.Config(),
renderer=RendererConfig(...),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
trainer=PolicyTrainer.Config(...),
generator=VLLMGenerator.Config(...),
)
Expand Down Expand Up @@ -110,7 +112,7 @@ uv venv --python 3.12 titan-rl
source titan-rl/bin/activate
```

1. Install Monarch, TorchStore, and Renderers from main:
1. Install Monarch, TorchStore, and Renderers:
```bash
uv pip install -r torchtitan/experiments/rl/requirements.txt
uv pip install --no-deps "git+https://github.com/meta-pytorch/torchstore.git@main"
Expand Down
5 changes: 4 additions & 1 deletion torchtitan/experiments/rl/actors/rollout_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,8 +11,9 @@
from typing import Any

from monarch.actor import Actor, concurrent_endpoint
from torchtitan.experiments.rl.renderer import RendererConfig

from torchtitan.components.tokenizer import HuggingFaceTokenizer
from torchtitan.experiments.rl.renderer import RendererConfig
from torchtitan.experiments.rl.rollout.rollouter import RolloutWorker
from torchtitan.experiments.rl.rollout.types import RolloutGroup
from torchtitan.observability import structured_logger as sl
Expand All @@ -36,10 +37,12 @@ def __init__(
async def setup_async(
self,
*,
tokenizer_config: HuggingFaceTokenizer.Config,
renderer_config: RendererConfig,
hf_assets_path: str,
) -> None:
await self._worker.setup_async(
tokenizer_config=tokenizer_config,
renderer_config=renderer_config,
hf_assets_path=hf_assets_path,
)
Expand Down
19 changes: 12 additions & 7 deletions torchtitan/experiments/rl/controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,7 @@
from monarch.actor import ProcMesh, this_host
from monarch.spmd import setup_torch_elastic_env_async

from torchtitan.components.tokenizer import HuggingFaceTokenizer
from torchtitan.config import CompileConfig, Configurable
from torchtitan.experiments.rl.actors.generator import SamplingConfig, VLLMGenerator
from torchtitan.experiments.rl.actors.trainer import PolicyTrainer
Expand Down Expand Up @@ -297,8 +298,14 @@ class Config(Configurable.Config):
"""The rollouter: its datasets, envs, and rubric."""
# TODO: support multiple rollouters for data mixing.

tokenizer: HuggingFaceTokenizer.Config = field(
default_factory=HuggingFaceTokenizer.Config
)
"""Tokenizer loaded from `hf_assets_path`."""

renderer: RendererConfig
"""Message-to-token renderer config."""
"""The model's chat template; renders messages to token ids and parses completions
back. E.g. `RenderersLibraryConfig(renderers_config=Qwen3RendererConfig(enable_thinking=False))`."""

rollout_recorder: RolloutSampleRecorder.Config = field(
default_factory=RolloutSampleRecorder.Config
Expand Down Expand Up @@ -424,7 +431,8 @@ def __init__(self, config: Config):
log_dir=config.dump_folder,
job_config=config.to_dict(),
)
self.renderer = config.renderer.build(tokenizer_path=config.hf_assets_path)
self.tokenizer = config.tokenizer.build(tokenizer_path=config.hf_assets_path)
self.renderer = config.renderer.build(tokenizer=self.tokenizer)

# Carry the base seed and renderer stop tokens on the sampling config so
# the generator reads them off each request; the rollouter offsets the
Expand All @@ -434,10 +442,6 @@ def __init__(self, config: Config):
seed=config.generator.debug.seed,
stop_token_ids=list(self.renderer.get_stop_token_ids()),
)
# TODO: pass our own tokenizer to the renderer and read pad/eos off it
# once `renderers` supports bring-your-own-tokenizer
# (https://github.com/PrimeIntellect-ai/renderers/pull/70).
# Until then, reach into the renderer's tokenizer for the pad id (eos doubles as pad).
self._rollouter: Rollouter = config.rollouter.build()
self.rollout_recorder = config.rollout_recorder.build(
dump_dir=config.dump_folder
Expand Down Expand Up @@ -636,6 +640,7 @@ async def setup_async(
)

await self._rollouter.setup_async(
tokenizer_config=config.tokenizer,
renderer_config=config.renderer,
hf_assets_path=config.hf_assets_path,
)
Expand Down Expand Up @@ -809,7 +814,7 @@ async def run(self) -> None:
max_context_length=self.config.trainer.training.max_context_length,
num_prompts_per_train_step=async_loop.num_prompts_per_train_step,
dp_degree=self.trainer_dp_degree,
pad_id=self.renderer._tokenizer.eos_token_id,
pad_id=self.tokenizer.eos_id,
)

# training_batch_queue
Expand Down
1 change: 1 addition & 0 deletions torchtitan/experiments/rl/environment/token.py
Original file line number Diff line number Diff line change
Expand Up @@ -161,6 +161,7 @@ async def step(self, completion: Completion) -> TokenEnvOutput:
parsed = await asyncio.to_thread(
self._renderer.parse_response,
token_ids=completion.token_ids,
tools=self._tools,
)
except Exception:
logger.exception(
Expand Down
56 changes: 42 additions & 14 deletions torchtitan/experiments/rl/examples/alphabet_sort/config_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,8 @@

import dataclasses

from renderers import GptOssRendererConfig, Qwen3RendererConfig

from torchtitan.components.checkpointer import CheckpointManager
from torchtitan.components.loss import ChunkedLossWrapper
from torchtitan.components.optimizer import default_adamw, LRSchedulersContainer
Expand Down Expand Up @@ -44,7 +46,7 @@
from torchtitan.experiments.rl.models.cast_linear import LMHeadCastConverter
from torchtitan.experiments.rl.models.vllm_registry import InferenceParallelismConfig
from torchtitan.experiments.rl.observability.metrics import MetricsProcessor
from torchtitan.experiments.rl.renderer import RendererConfig
from torchtitan.experiments.rl.renderer import RenderersLibraryConfig
from torchtitan.experiments.rl.routing.inter_generator_router import (
InterGeneratorRouter,
)
Expand Down Expand Up @@ -98,7 +100,9 @@ def rl_grpo_qwen3_0_6b_varlen() -> Controller.Config:
),
compile=CompileConfig(enable=True, backend="aot_eager"),
rollouter=AlphabetSortRollouter.Config(),
renderer=RendererConfig(name="qwen3", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
generator_router=InterGeneratorRouter.Config(
strategy=StickySessionRoutingStrategy.Config(
fallback_strategy=LeastLoadedRoutingStrategy.Config()
Expand Down Expand Up @@ -160,7 +164,9 @@ def rl_grpo_qwen3_0_6b_flex() -> Controller.Config:
),
compile=CompileConfig(enable=True, backend="aot_eager"),
rollouter=AlphabetSortRollouter.Config(),
renderer=RendererConfig(name="qwen3", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=2e-6),
Expand Down Expand Up @@ -262,7 +268,9 @@ def rl_grpo_gpt_oss_20b_varlen() -> Controller.Config:
),
compile=CompileConfig(enable=True, backend="aot_eager"),
rollouter=AlphabetSortRollouter.Config(),
renderer=RendererConfig(name="gpt_oss", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=GptOssRendererConfig(reasoning_effort="low")
),
generator_router=InterGeneratorRouter.Config(
strategy=StickySessionRoutingStrategy.Config(
fallback_strategy=LeastLoadedRoutingStrategy.Config()
Expand Down Expand Up @@ -330,7 +338,9 @@ def rl_grpo_gpt_oss_debug_varlen() -> Controller.Config:
# Debug tokenizer (vocab 2048, matches debugmodel); the gpt_oss renderer
# needs gpt-oss special tokens absent here, so use the qwen3 renderer
# like the other debug configs.
renderer=RendererConfig(name="qwen3", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=2e-6),
Expand Down Expand Up @@ -398,7 +408,9 @@ def rl_grpo_gpt_oss_debug_varlen_batch_invariant() -> Controller.Config:
# Debug tokenizer (vocab 2048, matches debugmodel); the gpt_oss renderer
# needs gpt-oss special tokens absent here, so use the qwen3 renderer
# like the other debug configs.
renderer=RendererConfig(name="qwen3", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=2e-6),
Expand Down Expand Up @@ -458,7 +470,9 @@ def rl_grpo_qwen3_1_7b() -> Controller.Config:
),
compile=CompileConfig(enable=True, backend="aot_eager"),
rollouter=AlphabetSortRollouter.Config(),
renderer=RendererConfig(name="qwen3", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=2e-6),
Expand Down Expand Up @@ -515,7 +529,9 @@ def rl_grpo_qwen3_14b() -> Controller.Config:
),
compile=CompileConfig(enable=True, backend="aot_eager"),
rollouter=AlphabetSortRollouter.Config(),
renderer=RendererConfig(name="qwen3", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=1e-6),
Expand Down Expand Up @@ -583,7 +599,9 @@ def rl_grpo_qwen3_moe_debug_varlen() -> Controller.Config:
# torch.compile and CUDA graph capture; disable both.
compile=CompileConfig(enable=False),
rollouter=AlphabetSortRollouter.Config(),
renderer=RendererConfig(name="qwen3", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=8e-4),
Expand Down Expand Up @@ -714,7 +732,9 @@ def rl_grpo_qwen3_moe_debug_varlen_batch_invariant() -> Controller.Config:
# torch.compile and CUDA graph capture; disable both.
compile=CompileConfig(enable=False),
rollouter=AlphabetSortRollouter.Config(),
renderer=RendererConfig(name="qwen3", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=8e-4),
Expand Down Expand Up @@ -782,7 +802,9 @@ def rl_grpo_qwen3_30b_a3b_varlen() -> Controller.Config:
),
compile=CompileConfig(enable=False),
rollouter=AlphabetSortRollouter.Config(),
renderer=RendererConfig(name="qwen3", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=1e-6),
Expand Down Expand Up @@ -893,7 +915,9 @@ def rl_grpo_qwen3_0_6b_varlen_batch_invariant() -> Controller.Config:
),
compile=CompileConfig(enable=True, backend="aot_eager"),
rollouter=AlphabetSortRollouter.Config(),
renderer=RendererConfig(name="qwen3", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=2e-6),
Expand Down Expand Up @@ -972,7 +996,9 @@ def rl_grpo_qwen3_5_9b_varlen() -> Controller.Config:
),
compile=CompileConfig(enable=False),
rollouter=AlphabetSortRollouter.Config(),
renderer=RendererConfig(name="qwen3", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=1e-6),
Expand Down Expand Up @@ -1056,7 +1082,9 @@ def rl_grpo_qwen3_5_debug_varlen() -> Controller.Config:
),
compile=CompileConfig(enable=False),
rollouter=AlphabetSortRollouter.Config(),
renderer=RendererConfig(name="qwen3", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=1e-6),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,8 @@

from __future__ import annotations

from renderers import Qwen3RendererConfig

from torchtitan.components.checkpointer import CheckpointManager
from torchtitan.components.loss import ChunkedLossWrapper
from torchtitan.components.optimizer import default_adamw, LRSchedulersContainer
Expand All @@ -33,7 +35,7 @@
from torchtitan.experiments.rl.models.cast_linear import LMHeadCastConverter
from torchtitan.experiments.rl.models.vllm_registry import InferenceParallelismConfig
from torchtitan.experiments.rl.observability.metrics import MetricsProcessor
from torchtitan.experiments.rl.renderer import RendererConfig
from torchtitan.experiments.rl.renderer import RenderersLibraryConfig
from torchtitan.experiments.rl.routing.inter_generator_router import (
InterGeneratorRouter,
)
Expand Down Expand Up @@ -81,7 +83,9 @@ def _qwen3_4b_dapo_math_config(
),
),
),
renderer=RendererConfig(name="qwen3", enable_thinking=True),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=True)
),
num_generators=6,
generator_router=InterGeneratorRouter.Config(
strategy=LeastLoadedRoutingStrategy.Config()
Expand Down
22 changes: 14 additions & 8 deletions torchtitan/experiments/rl/examples/search_r1/config_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,8 @@

import dataclasses

from renderers import Qwen3RendererConfig

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think this is because Prime-rl "happens" to be using similar way of configuring, but the "proper" and robust way to use something in other library is to build our own wrapper class.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.


from torchtitan.components.checkpointer import CheckpointManager
from torchtitan.components.loss import ChunkedLossWrapper
from torchtitan.components.optimizer import default_adamw, LRSchedulersContainer
Expand All @@ -44,9 +46,12 @@
SearchR1Worker,
)
from torchtitan.experiments.rl.losses import DAPOLoss
from torchtitan.experiments.rl.models.muse_glimmer.renderer import (
MuseGlimmerRendererConfig,
)
from torchtitan.experiments.rl.models.vllm_registry import InferenceParallelismConfig
from torchtitan.experiments.rl.observability.metrics import MetricsProcessor
from torchtitan.experiments.rl.renderer import RendererConfig
from torchtitan.experiments.rl.renderer import RenderersLibraryConfig
from torchtitan.experiments.rl.rollout.advantage import AdvantageEstimator
from torchtitan.models.muse_glimmer import model_registry as muse_glimmer_model_registry
from torchtitan.models.muse_glimmer.state_dict_adapter import (
Expand Down Expand Up @@ -78,7 +83,9 @@ def rl_grpo_qwen3_1_7b_search_r1() -> Controller.Config:
advantage=AdvantageEstimator.Config(should_std_normalize=True),
),
),
renderer=RendererConfig(name="qwen3", enable_thinking=False),
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=1e-6),
Expand Down Expand Up @@ -206,7 +213,9 @@ def rl_grpo_qwen3_30b_a3b_deepep_search_r1_perf() -> Controller.Config:
advantage=AdvantageEstimator.Config(should_std_normalize=True),
),
),
renderer=RendererConfig(name="qwen3", enable_thinking=False), # TODO: TBD
renderer=RenderersLibraryConfig(
renderers_config=Qwen3RendererConfig(enable_thinking=False)
),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=1e-6),
Expand Down Expand Up @@ -280,12 +289,9 @@ def rl_grpo_muse_glimmer_30b_search_r1() -> Controller.Config:

varlen attention is used for both roles so the trainer and the vLLM generator run
one ModelSpec. The state-dict adapter handles the HF checkpoint's Q/K RoPE layout
on load, and the renderer (registered below) handles Muse Glimmer's harmony chat
on load, and the renderer handles Muse Glimmer's harmony chat
format and ATEM tool calls.
"""
# Muse Glimmer's renderer ships in torchtitan rather than the `renderers` library;
# registering makes RendererConfig(name="muse_glimmer") resolve it.

model_spec = muse_glimmer_model_registry("30B", attn_backend="varlen")
model_spec = dataclasses.replace(
model_spec, state_dict_adapter=MuseGlimmerStateDictAdapter
Expand All @@ -306,7 +312,7 @@ def rl_grpo_muse_glimmer_30b_search_r1() -> Controller.Config:
advantage=AdvantageEstimator.Config(should_std_normalize=True),
),
),
renderer=RendererConfig(name="muse_glimmer", enable_thinking=True),
renderer=MuseGlimmerRendererConfig(),
metrics=MetricsProcessor.Config(enable_wandb=True),
trainer=PolicyTrainer.Config(
optimizer=default_adamw(lr=1e-6),
Expand Down
3 changes: 2 additions & 1 deletion torchtitan/experiments/rl/generate.py
Original file line number Diff line number Diff line change
Expand Up @@ -187,7 +187,8 @@ def generate() -> None:

logger.debug("vLLM LLMEngine initialized successfully")

renderer = config.renderer.build(tokenizer_path=model_path)
tokenizer = config.tokenizer.build(tokenizer_path=model_path)
renderer = config.renderer.build(tokenizer=tokenizer)
stop_token_ids = list(renderer.get_stop_token_ids())

# Create sampling parameters from config
Expand Down
Loading
Loading