Skip to content
Merged
Show file tree
Hide file tree
Changes from all 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