diff --git a/torchtitan/experiments/rl/README.md b/torchtitan/experiments/rl/README.md index 5df4088888..cad2b5f4d7 100644 --- a/torchtitan/experiments/rl/README.md +++ b/torchtitan/experiments/rl/README.md @@ -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(...), ) @@ -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" diff --git a/torchtitan/experiments/rl/actors/rollout_worker.py b/torchtitan/experiments/rl/actors/rollout_worker.py index a292c3dbab..6084769204 100644 --- a/torchtitan/experiments/rl/actors/rollout_worker.py +++ b/torchtitan/experiments/rl/actors/rollout_worker.py @@ -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 @@ -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, ) diff --git a/torchtitan/experiments/rl/controller.py b/torchtitan/experiments/rl/controller.py index 2ece881348..ab7dad0022 100644 --- a/torchtitan/experiments/rl/controller.py +++ b/torchtitan/experiments/rl/controller.py @@ -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 @@ -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 @@ -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 @@ -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 @@ -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, ) @@ -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 diff --git a/torchtitan/experiments/rl/environment/token.py b/torchtitan/experiments/rl/environment/token.py index 3a528a4b88..29d1dba86c 100644 --- a/torchtitan/experiments/rl/environment/token.py +++ b/torchtitan/experiments/rl/environment/token.py @@ -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( diff --git a/torchtitan/experiments/rl/examples/alphabet_sort/config_registry.py b/torchtitan/experiments/rl/examples/alphabet_sort/config_registry.py index e5ff428a95..08454dc76f 100644 --- a/torchtitan/experiments/rl/examples/alphabet_sort/config_registry.py +++ b/torchtitan/experiments/rl/examples/alphabet_sort/config_registry.py @@ -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 @@ -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, ) @@ -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() @@ -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), @@ -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() @@ -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), @@ -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), @@ -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), @@ -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), @@ -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), @@ -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), @@ -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), @@ -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), @@ -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), @@ -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), diff --git a/torchtitan/experiments/rl/examples/dapo_math/config_registry.py b/torchtitan/experiments/rl/examples/dapo_math/config_registry.py index 6ccc2f8edd..4ccce8b6e7 100644 --- a/torchtitan/experiments/rl/examples/dapo_math/config_registry.py +++ b/torchtitan/experiments/rl/examples/dapo_math/config_registry.py @@ -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 @@ -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, ) @@ -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() diff --git a/torchtitan/experiments/rl/examples/search_r1/config_registry.py b/torchtitan/experiments/rl/examples/search_r1/config_registry.py index 6f71e4ddf2..cc6153b26a 100644 --- a/torchtitan/experiments/rl/examples/search_r1/config_registry.py +++ b/torchtitan/experiments/rl/examples/search_r1/config_registry.py @@ -18,6 +18,8 @@ import dataclasses +from renderers import Qwen3RendererConfig + from torchtitan.components.checkpointer import CheckpointManager from torchtitan.components.loss import ChunkedLossWrapper from torchtitan.components.optimizer import default_adamw, LRSchedulersContainer @@ -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 ( @@ -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), @@ -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), @@ -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 @@ -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), diff --git a/torchtitan/experiments/rl/generate.py b/torchtitan/experiments/rl/generate.py index 5b3f5ab45f..805fdeb382 100755 --- a/torchtitan/experiments/rl/generate.py +++ b/torchtitan/experiments/rl/generate.py @@ -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 diff --git a/torchtitan/experiments/rl/models/muse_glimmer/renderer.py b/torchtitan/experiments/rl/models/muse_glimmer/renderer.py index 2be5350f70..f0b7c4208b 100644 --- a/torchtitan/experiments/rl/models/muse_glimmer/renderer.py +++ b/torchtitan/experiments/rl/models/muse_glimmer/renderer.py @@ -28,18 +28,13 @@ reasoning states. Treat that test as the spec -- if the template changes upstream, it fails first. -Implements the ``renderers.Renderer`` Protocol. ``register()`` installs it into the -``renderers`` library's public registry (``RENDERER_REGISTRY`` / ``_CONFIG_BY_NAME``), -which is that library's supported extension path -- no fork or upstream change needed. - -Every other TorchTitan model resolves to a renderer that lives in -PrimeIntellect-ai/renderers. This one ships here because Muse Glimmer is not in that -library yet. +Implements the ``renderers.Renderer`` Protocol. Muse Glimmer is not in the library yet, so +``MuseGlimmerRendererConfig`` is a TorchTitan ``RendererConfig`` whose ``build`` constructs +this class directly. TODO: upstream this to PrimeIntellect-ai/renderers (renderer -> renderers/muse_glimmer.py, -atem.py -> a tool parser in renderers/parsers.py), then delete both files and -``register()``, leaving only the ``_RENDERER_BY_MODEL`` entry in -experiments/rl/renderer.py. +atem.py -> a tool parser in renderers/parsers.py), then delete both files and select it +through ``RenderersLibraryConfig`` like the other renderers. It lives under ``experiments/rl`` rather than ``torchtitan/models/muse_glimmer`` because RL is its only consumer and ``renderers`` is an RL-only optional dependency; keeping it @@ -51,8 +46,10 @@ import datetime import json import re -from typing import ClassVar, Literal, NamedTuple +from dataclasses import dataclass, replace +from typing import NamedTuple +from renderers import Renderer from renderers.base import ( extract_message_tool_names, ParsedResponse, @@ -63,31 +60,18 @@ should_rerender_for_thinking_retention, trim_to_turn_close, ) -from renderers.configs import BaseRendererConfig +from renderers.configs import ThinkingRetention -from .atem import parse_atem_tool_calls, render_atem_tool_call +from torchtitan.components.tokenizer import HuggingFaceTokenizer +from torchtitan.experiments.rl.renderer import RendererConfig, RendererTokenizerWrapper -RENDERER_NAME = "muse_glimmer" +from .atem import parse_atem_tool_calls, render_atem_tool_call -class MuseGlimmerRendererConfig(BaseRendererConfig): +@dataclass(kw_only=True, slots=True) +class MuseGlimmerRendererConfig(RendererConfig): """Muse Glimmer (harmony chat format + ATEM tool calls) renderer config.""" - name: Literal["muse_glimmer"] = RENDERER_NAME - - # renderers validates in BaseRendererConfig.__pydantic_init_subclass__ that every - # non-base field is classified as either a chat-template kwarg or a renderer-internal - # knob; the two sets must be disjoint and together cover all of them. Declared - # unconditionally -- versions without the validator ignore these ClassVars, so this - # is compatible with both. The template fields mirror kwargs the published - # chat_template.jinja reads, which is what the library's parity matrix varies. - _template_fields: ClassVar[frozenset[str]] = frozenset( - {"reasoning_strength", "knowledge_cutoff", "current_date"} - ) - _internal_fields: ClassVar[frozenset[str]] = frozenset( - {"retain_reasoning_in_history", "answer_from_reasoning_fallback"} - ) - reasoning_strength: str | None = None """Sizes the reasoning budget, rendered as ``Reasoning strength: .`` @@ -131,6 +115,34 @@ class MuseGlimmerRendererConfig(BaseRendererConfig): empty ``content`` is unscoreable and the answer is often the final reasoning line. """ + thinking_retention: ThinkingRetention | None = None + """The library-wide bridge policy override (`renderers.BaseRendererConfig.thinking_retention`). + ``None`` keeps the template's implied policy; ``"tool_cycle"`` re-renders at a new user query.""" + + def __post_init__(self) -> None: + # A dataclass does not validate values; these knobs change bridging and reward + # scoring with nothing visible in the rendered prompt, so check them here. + for name in ("reasoning_strength", "knowledge_cutoff", "current_date"): + value = getattr(self, name) + if value is not None and not isinstance(value, str): + raise TypeError( + f"{name} must be str | None, got {type(value).__name__}" + ) + for name in ("retain_reasoning_in_history", "answer_from_reasoning_fallback"): + value = getattr(self, name) + if type(value) is not bool: + raise TypeError(f"{name} must be bool, got {type(value).__name__}") + if self.thinking_retention not in (None, "tool_cycle", "all"): + raise ValueError( + "thinking_retention must be None, 'tool_cycle' or 'all', " + f"got {self.thinking_retention!r}" + ) + + def build(self, *, tokenizer: HuggingFaceTokenizer) -> Renderer: + # Snapshot the config, as `Configurable.Config.build` does, so later edits to the + # recipe object cannot desynchronize full renders from bridging. + return MuseGlimmerRenderer(RendererTokenizerWrapper(tokenizer), replace(self)) + # Muse Glimmer special tokens. The ids are checked against the tokenizer in __init__ # rather than trusted, since they are baked into parse_response and the loss mask. @@ -343,17 +355,13 @@ class _Piece(NamedTuple): class MuseGlimmerRenderer: def __init__(self, tokenizer, config: MuseGlimmerRendererConfig | None = None): - # (tokenizer, config) is the renderers-library constructor contract, so - # ``create_renderer`` can instantiate this from RENDERER_REGISTRY. + # Match the `(tokenizer, config)` constructor used by library renderers. self._tok = tokenizer self._config = config or MuseGlimmerRendererConfig() - # The controller reads renderer._tokenizer (e.g. for pad_id=eos_token_id). - self._tokenizer = tokenizer self._bos = tokenizer.bos_token or "" - # BaseRendererConfig.thinking_retention is the library-wide knob every renderer - # is expected to honour in its bridge. Muse Glimmer's published chat template - # renders reasoning_content for every assistant turn unconditionally -- no - # query-boundary drop like gpt-oss's auto_drop_analysis or Qwen3's think-block + # `thinking_retention` is the library-wide bridge knob. Muse Glimmer's published + # chat template renders reasoning_content for every assistant turn unconditionally + # -- no query-boundary drop like gpt-oss's auto_drop_analysis or Qwen3's think-block # stripping -- so "all" is the template-faithful implied policy. An explicit # thinking_retention on the config overrides it. self.effective_thinking_retention = resolve_thinking_retention( @@ -758,31 +766,3 @@ def parse_response(self, token_ids, *, tools=None) -> ParsedResponse: reasoning_content=reasoning, tool_calls=tool_calls, ) - - -def register() -> None: - """Install the muse_glimmer renderer into the ``renderers`` library registry. - - Uses the library's public extension surface -- implement the ``Renderer`` - Protocol, then add the class to ``RENDERER_REGISTRY`` and its config to - ``_CONFIG_BY_NAME`` -- so ``create_renderer(config_from_name("muse_glimmer"))`` - resolves it. Also maps the ``muse_glimmer`` TorchTitan model name to it, which is what - ``RendererConfig(name="muse_glimmer")`` looks up. - - Idempotent. Delete this once the renderer is upstreamed to - PrimeIntellect-ai/renderers (only the _RENDERER_BY_MODEL entry stays). - """ - from renderers import base as renderers_base, configs as renderers_configs - - from torchtitan.experiments.rl.renderer import _RENDERER_BY_MODEL - - # Populate the library's built-ins first: _populate_registry() early-returns if - # RENDERER_REGISTRY is already non-empty, so registering before it runs would - # suppress every built-in renderer. - renderers_base._populate_registry() - - renderers_configs._CONFIG_BY_NAME.setdefault( - RENDERER_NAME, MuseGlimmerRendererConfig - ) - renderers_base.RENDERER_REGISTRY[RENDERER_NAME] = MuseGlimmerRenderer - _RENDERER_BY_MODEL["muse_glimmer"] = RENDERER_NAME diff --git a/torchtitan/experiments/rl/renderer.py b/torchtitan/experiments/rl/renderer.py index 013f945ae7..76279fa282 100644 --- a/torchtitan/experiments/rl/renderer.py +++ b/torchtitan/experiments/rl/renderer.py @@ -6,109 +6,142 @@ from __future__ import annotations -import logging -from dataclasses import dataclass, fields +from dataclasses import dataclass +from typing import Annotated, Any -from renderers import config_from_name, create_renderer, Renderer +import tyro +from renderers import create_renderer, Renderer +from renderers.configs import BaseRendererConfig +from torchtitan.components.tokenizer import HuggingFaceTokenizer from torchtitan.config import Configurable -logger = logging.getLogger(__name__) - -# Map a TorchTitan model name to its `renderers` renderer. Models not listed fall -# back to "auto" (renderers resolves from the tokenizer) -# https://github.com/PrimeIntellect-ai/renderers/blob/942449c37ab6e9fab26d59b40336514c8baa6b13/renderers/configs.py#L404 -_RENDERER_BY_MODEL = { - "qwen3": "qwen3", - "qwen3_vl": "qwen3-vl", - "gpt_oss": "gpt-oss", - "deepseek_v3": "deepseek-v3", - # TODO: upstream the Muse Glimmer renderer to PrimeIntellect-ai/renderers, then - # delete its `register()` and point this at the library's name (hyphenated, like - # the entries above). It ships in torchtitan and self-registers only because the - # library has no Muse Glimmer renderer yet; every other model here resolves to one - # the library owns. See rl/models/muse_glimmer/renderer.py. - "muse_glimmer": "muse_glimmer", - "default": "default", # llama3 - "auto": "auto", # ignores knobs, resolves from tokenizer, -} - @dataclass(kw_only=True, slots=True) class RendererConfig(Configurable.Config): - """Selects the renderer used for chat message <-> token conversion. + """Base config of a renderer; `build` returns a `renderers.Renderer` on TorchTitan's tokenizer. + + Subclasses: `RenderersLibraryConfig` for a renderer from the `renderers` library, and + in-tree renderers such as `MuseGlimmerRendererConfig`. + """ - Wraps `PrimeIntellect-ai/renderers`. `build` loads a tokenizer from - `tokenizer_path`, maps the model `name` to a renderer, and forwards any - supported knobs. + def build(self, *, tokenizer: HuggingFaceTokenizer) -> Renderer: + raise NotImplementedError - Args: - name: TorchTitan model name (e.g. `"qwen3"`, `"llama3"`), mapped to a - `renderers` renderer via `_RENDERER_BY_MODEL`. `None` (the default) - resolves the renderer from the tokenizer. - tool_parser: Tool-call parser name, when the renderer supports it. - reasoning_parser: Reasoning parser name, when the renderer supports it. - enable_thinking: Let the model emit reasoning, when supported. - preserve_all_thinking: Keep historical reasoning in future prompts. - preserve_thinking_between_tool_calls: Keep reasoning during tool loops. - Every field defaults to `None`; a non-`None` value overrides that knob on the - chosen renderer's config, otherwise the renderer keeps its own default. +@dataclass(kw_only=True, slots=True) +class RenderersLibraryConfig(RendererConfig): + """Builds one of the `renderers` library's renderers on TorchTitan's tokenizer. Example: - renderer = RendererConfig(name="qwen3").build(tokenizer_path="./Qwen3-0.6B") + from renderers import Qwen3RendererConfig + + from torchtitan.components.tokenizer import HuggingFaceTokenizer + from torchtitan.experiments.rl.renderer import RenderersLibraryConfig + + renderer = RenderersLibraryConfig( + renderers_config=Qwen3RendererConfig(enable_thinking=False) + ).build(tokenizer=HuggingFaceTokenizer(tokenizer_path="./Qwen3-0.6B")) prompt_ids = renderer.render_ids( - [{"role": "user", "content": "hi"}], add_generation_prompt=True + [{"role": "user", "content": "hi"}], + add_generation_prompt=True, ) """ - name: str | None = None - tool_parser: str | None = None - reasoning_parser: str | None = None - enable_thinking: bool | None = None - preserve_all_thinking: bool | None = None - preserve_thinking_between_tool_calls: bool | None = None - - def build(self, *, tokenizer_path: str) -> Renderer: - # TODO(renderers#70): use TorchTitan's tokenizer once `renderers` supports - # bring-your-own-tokenizer (PR adds a Tokenizer protocol; drops transformers). - from transformers import AutoTokenizer - - tokenizer = AutoTokenizer.from_pretrained(tokenizer_path) - - # `name=None` (or "auto") -> let `create_renderer` resolve from the tokenizer. - renderer_name = _RENDERER_BY_MODEL.get(self.name, self.name) - if renderer_name == "muse_glimmer": - # TODO: temporary. Delete this block once the Muse Glimmer renderer is - # upstreamed to PrimeIntellect-ai/renderers -- the library registers its - # own renderers in _populate_registry(), so no torchtitan-side hook is - # needed for any other model here. - # - # Until then it has to live in build(), not in the config registry: the - # renderer is constructed inside a Monarch-spawned RolloutWorker, which - # only receives the serialized config and never imports the recipe module, - # so a register() call there never runs in that process. - from torchtitan.experiments.rl.models.muse_glimmer import ( - renderer as _muse_glimmer_renderer, + renderers_config: Annotated[BaseRendererConfig, tyro.conf.Suppress] + """The library's typed config for the model, e.g. `Qwen3RendererConfig(enable_thinking=False)`. + Renderers and their options: + https://github.com/PrimeIntellect-ai/renderers/blob/renderers-v0.1.11/docs/renderer-config.md""" + + def to_dict(self) -> dict[str, Any]: + return {"renderers_config": self.renderers_config.model_dump(mode="json")} + + def build(self, *, tokenizer: HuggingFaceTokenizer) -> Renderer: + if self.renderers_config.name == "auto": + raise ValueError( + f"AutoRendererConfig resolves by exact match of tokenizer.name_or_path ({tokenizer.tokenizer_path!r}) " + "against renderers' MODEL_RENDERER_MAP, else falls back to DefaultRenderer (unsupported here). " + "Pick the model's renderer, e.g. Qwen3RendererConfig(...)." + ) + if self.renderers_config.name == "default": + raise ValueError( + "DefaultRenderer needs Hugging Face apply_chat_template; TorchTitan's template rendering lacks " + "its special-token variables (bos_token, ...) and would silently produce different tokens. " + "Pick the model's renderer, e.g. Qwen3RendererConfig(...)." ) + return create_renderer( + tokenizer=RendererTokenizerWrapper(tokenizer), config=self.renderers_config + ) + + +class RendererTokenizerWrapper: + """Adapt TorchTitan's loaded tokenizer to `renderers.OffsetTokenizer`. + + Protocol and bring-your-own-tokenizer guide: + https://github.com/PrimeIntellect-ai/renderers/blob/renderers-v0.1.11/renderers/base.py#L668-L699 + https://github.com/PrimeIntellect-ai/renderers/blob/renderers-v0.1.11/README.md#install + + `renderers` needs Hugging Face-style special-token attributes, raw encoding + without automatic BOS/EOS, token-to-id lookup, and character offsets. The + offsets identify tokens from message content (`is_content`). This adapter + exposes that interface from TorchTitan's underlying `tokenizers.Tokenizer`; + it does not load a second tokenizer. + + Example: + + from torchtitan.components.tokenizer import HuggingFaceTokenizer + from torchtitan.experiments.rl.renderer import RendererTokenizerWrapper + + tokenizer = RendererTokenizerWrapper( + HuggingFaceTokenizer(tokenizer_path="./Qwen3-0.6B") + ) + tokenizer.encode("hi") # [6023] + tokenizer( + "hi", + add_special_tokens=False, + return_offsets_mapping=True, + ) # ids + character offsets + tokenizer.convert_tokens_to_ids("<|im_end|>") # 151645 + """ + + def __init__(self, tokenizer: HuggingFaceTokenizer): + # The `tokenizers.Tokenizer` inside; it has the offsets and token -> id lookup. + self._tokenizer_backend = tokenizer.tokenizer + self.name_or_path = tokenizer.tokenizer_path + self.bos_token = tokenizer.bos_token + self.eos_token = tokenizer.eos_token + self.bos_token_id = tokenizer.bos_id + self.eos_token_id = tokenizer.eos_id + # `tokenizers` returns None for unknown tokens; it has no unk id. + self.unk_token_id = None + + def encode( + self, text: str, add_special_tokens: bool = False, **kwargs + ) -> list[int]: + return self._tokenizer_backend.encode( + text, add_special_tokens=add_special_tokens + ).ids + + def decode(self, token_ids, skip_special_tokens: bool = False, **kwargs) -> str: + return self._tokenizer_backend.decode( + list(token_ids), skip_special_tokens=skip_special_tokens + ) - _muse_glimmer_renderer.register() - renderer_config = config_from_name(renderer_name) if renderer_name else None - if renderer_config is None: - return create_renderer(tokenizer, None) - - # Rebuild the typed config and pass parameters - # that are not None and are supported - config_type = type(renderer_config) - args = { - field.name: getattr(self, field.name) # {key: value} - for field in fields(self) - if field.name != "name" # Get all self.fields, except name - and getattr(self, field.name) is not None # Only consider provided fields - and field.name in config_type.model_fields # Config supports this field - } - logger.info( - f"Using renderer {renderer_name}, of type {config_type}, with args {args}" + def convert_tokens_to_ids( + self, tokens: str | list[str] + ) -> int | None | list[int | None]: + if isinstance(tokens, str): + return self._tokenizer_backend.token_to_id(tokens) + return [self._tokenizer_backend.token_to_id(token) for token in tokens] + + def __call__( + self, text: str, *, add_special_tokens: bool, return_offsets_mapping: bool + ) -> dict: + encoding = self._tokenizer_backend.encode( + text, add_special_tokens=add_special_tokens ) - return create_renderer(tokenizer, config_type(**args)) + output = {"input_ids": encoding.ids} + if return_offsets_mapping: + output["offset_mapping"] = encoding.offsets + return output diff --git a/torchtitan/experiments/rl/requirements.txt b/torchtitan/experiments/rl/requirements.txt index 9ef9658732..46b8ca477c 100644 --- a/torchtitan/experiments/rl/requirements.txt +++ b/torchtitan/experiments/rl/requirements.txt @@ -3,4 +3,4 @@ opentelemetry-sdk opentelemetry-exporter-otlp-proto-http pygtrie portpicker -git+https://github.com/PrimeIntellect-ai/renderers.git@main +renderers==0.1.11 diff --git a/torchtitan/experiments/rl/rollout/rollouter.py b/torchtitan/experiments/rl/rollout/rollouter.py index 44e7be7c19..1f468284c7 100644 --- a/torchtitan/experiments/rl/rollout/rollouter.py +++ b/torchtitan/experiments/rl/rollout/rollouter.py @@ -13,8 +13,10 @@ from monarch.actor import ProcMesh, this_host +from torchtitan.components.tokenizer import HuggingFaceTokenizer from torchtitan.config import Configurable from torchtitan.experiments.rl.environment import MessageEnv, TokenEnv +from torchtitan.experiments.rl.renderer import RendererConfig from torchtitan.experiments.rl.rollout.advantage import AdvantageEstimator from torchtitan.experiments.rl.rollout.types import ( GenerateFn, @@ -33,7 +35,6 @@ from torchtitan.experiments.rl.actors.generator import SamplingConfig from torchtitan.experiments.rl.actors.rollout_worker import RolloutWorkerActor - from torchtitan.experiments.rl.renderer import RendererConfig logger = logging.getLogger(__name__) @@ -135,6 +136,7 @@ def get_validation_sample(self) -> object: async def setup_async( self, *, + tokenizer_config: HuggingFaceTokenizer.Config, renderer_config: RendererConfig, hf_assets_path: str, ) -> None: @@ -155,6 +157,7 @@ async def setup_async( num_threads=self._config.num_threads_per_worker, ) await self._worker_actors.setup_async.call( + tokenizer_config=tokenizer_config, renderer_config=renderer_config, hf_assets_path=hf_assets_path, ) @@ -241,11 +244,13 @@ def __init__(self, config: Config) -> None: async def setup_async( self, *, + tokenizer_config: HuggingFaceTokenizer.Config, renderer_config: RendererConfig, hf_assets_path: str, ) -> None: """Build runtime dependencies after the worker actor is spawned.""" - self._renderer = renderer_config.build(tokenizer_path=hf_assets_path) + tokenizer = tokenizer_config.build(tokenizer_path=hf_assets_path) + self._renderer = renderer_config.build(tokenizer=tokenizer) def make_env_group( self, diff --git a/torchtitan/experiments/rl/tests/integration_tests.py b/torchtitan/experiments/rl/tests/integration_tests.py index 11a1d5f92e..cdf429005a 100644 --- a/torchtitan/experiments/rl/tests/integration_tests.py +++ b/torchtitan/experiments/rl/tests/integration_tests.py @@ -46,7 +46,6 @@ def build_rl_test_list() -> list[OverrideDefinitions]: "--async-loop.num-samples-per-prompt 2", "--trainer.training.max_context_length 1024", "--trainer.training.num_tokens_per_microbatch_per_dp_rank 2048", - "--renderer.enable-thinking False", "--generator.sampling.max_tokens 256", "--trainer.debug.no_batch_invariant", "--generator.debug.no_batch_invariant", @@ -75,7 +74,6 @@ def build_rl_test_list() -> list[OverrideDefinitions]: "--async-loop.num-samples-per-prompt 2", "--trainer.training.max_context_length 1024", "--trainer.training.num_tokens_per_microbatch_per_dp_rank 2048", - "--renderer.enable-thinking False", "--generator.sampling.max_tokens 256", "--trainer.debug.no_batch_invariant", "--generator.debug.no_batch_invariant", @@ -104,7 +102,6 @@ def build_rl_test_list() -> list[OverrideDefinitions]: "--async-loop.num-samples-per-prompt 2", "--trainer.training.max_context_length 1024", "--trainer.training.num_tokens_per_microbatch_per_dp_rank 2048", - "--renderer.enable-thinking False", "--generator.sampling.max_tokens 256", "--trainer.debug.no_batch_invariant", "--generator.debug.no_batch_invariant", @@ -145,7 +142,6 @@ def build_rl_test_list() -> list[OverrideDefinitions]: "--async-loop.num-samples-per-prompt 2", "--trainer.training.max_context_length 1024", "--trainer.training.num_tokens_per_microbatch_per_dp_rank 2048", - "--renderer.enable-thinking False", "--generator.sampling.max_tokens 256", "--trainer.debug.no_batch_invariant", "--generator.debug.no_batch_invariant", @@ -166,7 +162,6 @@ def build_rl_test_list() -> list[OverrideDefinitions]: "--async-loop.num-samples-per-prompt 2", "--trainer.training.max_context_length 1024", "--trainer.training.num_tokens_per_microbatch_per_dp_rank 2048", - "--renderer.enable-thinking False", "--generator.sampling.max_tokens 256", "--trainer.debug.no_batch_invariant", "--generator.debug.no_batch_invariant", @@ -199,7 +194,6 @@ def build_rl_test_list() -> list[OverrideDefinitions]: "--async-loop.num-samples-per-prompt 2", "--trainer.training.max_context_length 1024", "--trainer.training.num_tokens_per_microbatch_per_dp_rank 2048", - "--renderer.enable-thinking False", "--generator.sampling.max_tokens 128", "--metrics.no-enable-wandb", ], @@ -220,7 +214,6 @@ def build_rl_test_list() -> list[OverrideDefinitions]: "--async-loop.num-samples-per-prompt 2", "--trainer.training.max_context_length 1024", "--trainer.training.num_tokens_per_microbatch_per_dp_rank 2048", - "--renderer.enable-thinking False", "--generator.sampling.max_tokens 256", "--trainer.checkpoint.no-enable", # use random-init weights "--generator.checkpoint.no-enable", @@ -247,7 +240,6 @@ def build_rl_test_list() -> list[OverrideDefinitions]: "--async-loop.num-samples-per-prompt 2", "--trainer.training.max_context_length 1024", "--trainer.training.num_tokens_per_microbatch_per_dp_rank 1024", - "--renderer.enable-thinking False", "--generator.sampling.max_tokens 128", "--trainer.checkpoint.no-enable", # random-init weights "--generator.checkpoint.no-enable", diff --git a/torchtitan/experiments/rl/tests/test_alphabet_sort.py b/torchtitan/experiments/rl/tests/test_alphabet_sort.py index d9665a9535..8d4305e93b 100644 --- a/torchtitan/experiments/rl/tests/test_alphabet_sort.py +++ b/torchtitan/experiments/rl/tests/test_alphabet_sort.py @@ -11,6 +11,9 @@ import asyncio import pytest +from renderers import Qwen3RendererConfig + +from torchtitan.components.tokenizer import HuggingFaceTokenizer from torchtitan.experiments.rl.examples.alphabet_sort import ( AlphabetSortDataset, @@ -22,6 +25,7 @@ ) from torchtitan.experiments.rl.examples.alphabet_sort.env import AlphabetSortEnv from torchtitan.experiments.rl.examples.alphabet_sort.rubric import score_sorted_list +from torchtitan.experiments.rl.renderer import RenderersLibraryConfig from torchtitan.experiments.rl.rollout import Rollout, RolloutStatus, RolloutTurn from torchtitan.experiments.rl.types import RolloutTurnID @@ -369,10 +373,6 @@ def test_env_walks_through_follow_up_turns() -> None: def test_rollouter_builds_one_env_per_group_member( monkeypatch: pytest.MonkeyPatch, ) -> None: - class _RendererConfig: - def build(self, *, tokenizer_path: str): - return None - _patch_names(monkeypatch) config = AlphabetSortRollouter.Config() rollouter = AlphabetSortRollouter(config) @@ -381,8 +381,11 @@ def build(self, *, tokenizer_path: str): assert isinstance(worker, AlphabetSortWorker) asyncio.run( worker.setup_async( - renderer_config=_RendererConfig(), - hf_assets_path="hf_assets_path", + tokenizer_config=HuggingFaceTokenizer.Config(), + renderer_config=RenderersLibraryConfig( + renderers_config=Qwen3RendererConfig(enable_thinking=False) + ), + hf_assets_path="tests/assets/tokenizer", ) ) sample = rollouter.get_training_sample() diff --git a/torchtitan/experiments/rl/tests/test_muse_glimmer_renderer.py b/torchtitan/experiments/rl/tests/test_muse_glimmer_renderer.py index 9a0410d845..bacdeec3f8 100644 --- a/torchtitan/experiments/rl/tests/test_muse_glimmer_renderer.py +++ b/torchtitan/experiments/rl/tests/test_muse_glimmer_renderer.py @@ -23,6 +23,7 @@ import pytest +from torchtitan.components.tokenizer import HuggingFaceTokenizer from torchtitan.experiments.rl.models.muse_glimmer.renderer import ( EOM_ID, EOT_ID, @@ -227,6 +228,46 @@ def _renderer(tokenizer, **overrides): return MuseGlimmerRenderer(tokenizer, MuseGlimmerRendererConfig(**overrides)) +@pytest.mark.parametrize( + ("kwargs", "error"), + [ + ({"thinking_retention": "everything"}, ValueError), + ({"retain_reasoning_in_history": "false"}, TypeError), + ({"answer_from_reasoning_fallback": 1}, TypeError), + ({"reasoning_strength": 123}, TypeError), + ], +) +def test_config_rejects_invalid_values(kwargs, error): + with pytest.raises(error): + MuseGlimmerRendererConfig(**kwargs) + + +def test_build_snapshots_config(): + path = os.environ.get("MUSE_GLIMMER_TOKENIZER", DEFAULT_TOKENIZER) + if not os.path.isdir(path): + pytest.skip("HuggingFaceTokenizer needs a local tokenizer directory") + config = MuseGlimmerRendererConfig(retain_reasoning_in_history=True) + renderer = config.build(tokenizer=HuggingFaceTokenizer(tokenizer_path=path)) + config.retain_reasoning_in_history = False + assert renderer._config.retain_reasoning_in_history is True + assert renderer.effective_thinking_retention == "all" + + +def test_config_build_matches_hf_tokenizer_path(tokenizer): + path = os.environ.get("MUSE_GLIMMER_TOKENIZER", DEFAULT_TOKENIZER) + if not os.path.isdir(path): + pytest.skip("HuggingFaceTokenizer needs a local tokenizer directory") + titan = MuseGlimmerRendererConfig(reasoning_strength="low").build( + tokenizer=HuggingFaceTokenizer(tokenizer_path=path) + ) + assert isinstance(titan, MuseGlimmerRenderer) + hf = _renderer(tokenizer, reasoning_strength="low") + messages = [{"role": "user", "content": "search for bob"}] + assert titan.render_ids( + messages, tools=TOOLS, add_generation_prompt=True + ) == hf.render_ids(messages, tools=TOOLS, add_generation_prompt=True) + + def _template_kwargs(config: MuseGlimmerRendererConfig) -> dict: """Only forward knobs the caller set, so the template applies its own defaults.""" return { diff --git a/torchtitan/experiments/rl/tests/test_renderer.py b/torchtitan/experiments/rl/tests/test_renderer.py new file mode 100644 index 0000000000..ac99a0caa5 --- /dev/null +++ b/torchtitan/experiments/rl/tests/test_renderer.py @@ -0,0 +1,132 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import json +from dataclasses import dataclass + +import pytest +from renderers import ( + AutoRendererConfig, + DefaultRendererConfig, + OffsetTokenizer, + Qwen3RendererConfig, + Tokenizer, +) + +from torchtitan.components.tokenizer import HuggingFaceTokenizer +from torchtitan.config import Configurable +from torchtitan.experiments.rl.renderer import ( + RendererConfig, + RenderersLibraryConfig, + RendererTokenizerWrapper, +) + +_TOKENIZER_PATH = "tests/assets/tokenizer" + + +# --- RenderersLibraryConfig --- + + +def test_build_renders_with_titan_tokenizer() -> None: + tokenizer = HuggingFaceTokenizer(tokenizer_path=_TOKENIZER_PATH) + renderer = RenderersLibraryConfig( + renderers_config=Qwen3RendererConfig(enable_thinking=False) + ).build(tokenizer=tokenizer) + rendered = renderer.render( + [{"role": "user", "content": "hi"}], add_generation_prompt=True + ) + assert rendered.token_ids[0] == tokenizer.token_to_id("<|im_start|>") + assert renderer.get_stop_token_ids() == [ + tokenizer.token_to_id("<|im_end|>"), + tokenizer.token_to_id("<|endoftext|>"), + ] + assert len(rendered.is_content) == len(rendered.token_ids) + + +@pytest.mark.parametrize( + ("renderers_config", "reason"), + [ + (AutoRendererConfig(), "MODEL_RENDERER_MAP"), + (DefaultRendererConfig(), "special-token variables"), + ], +) +def test_auto_and_default_are_refused(renderers_config, reason: str) -> None: + tokenizer = HuggingFaceTokenizer(tokenizer_path=_TOKENIZER_PATH) + with pytest.raises(ValueError) as error: + RenderersLibraryConfig(renderers_config=renderers_config).build( + tokenizer=tokenizer + ) + assert reason in str(error.value) + assert "Pick the model's renderer" in str(error.value) + + +def test_config_to_dict_is_json() -> None: + # The controller logs `Controller.Config.to_dict()` as the job config. + @dataclass(kw_only=True, slots=True) + class _Holder(Configurable.Config): + renderer: RendererConfig + + holder = _Holder( + renderer=RenderersLibraryConfig( + renderers_config=Qwen3RendererConfig(enable_thinking=False) + ) + ) + value = holder.to_dict() + assert value["renderer"]["renderers_config"]["name"] == "qwen3" + assert value["renderer"]["renderers_config"]["enable_thinking"] is False + json.dumps(value) + + +# --- RendererTokenizerWrapper --- + + +def test_renderer_tokenizer_satisfies_offset_protocol() -> None: + tokenizer = HuggingFaceTokenizer(tokenizer_path=_TOKENIZER_PATH) + renderer_tokenizer = RendererTokenizerWrapper(tokenizer) + assert isinstance(renderer_tokenizer, Tokenizer) + assert isinstance(renderer_tokenizer, OffsetTokenizer) + assert renderer_tokenizer.eos_token_id == tokenizer.eos_id + assert renderer_tokenizer.convert_tokens_to_ids( + "<|im_end|>" + ) == tokenizer.token_to_id("<|im_end|>") + encoding = renderer_tokenizer( + "hi there", add_special_tokens=False, return_offsets_mapping=True + ) + assert encoding["input_ids"] == renderer_tokenizer.encode("hi there") + assert len(encoding["offset_mapping"]) == len(encoding["input_ids"]) + + +def test_encode_never_adds_bos() -> None: + # The debug tokenizer has a BOS token; renderers place special tokens themselves. + tokenizer = HuggingFaceTokenizer(tokenizer_path=_TOKENIZER_PATH) + assert tokenizer.bos_id is not None + assert tokenizer.bos_id not in RendererTokenizerWrapper(tokenizer).encode("hi") + + +def test_render_matches_hf_tokenizer_path() -> None: + transformers = pytest.importorskip("transformers") + from renderers import create_renderer + + messages = [ + {"role": "system", "content": "Sort names."}, + {"role": "user", "content": "Zed, Amy <|im_end|> tricky"}, + {"role": "assistant", "reasoning_content": "think", "content": "Amy, Zed"}, + {"role": "user", "content": "Add Bob."}, + ] + config = Qwen3RendererConfig(enable_thinking=False) + hf = create_renderer( + transformers.AutoTokenizer.from_pretrained(_TOKENIZER_PATH), config + ) + titan = create_renderer( + RendererTokenizerWrapper(HuggingFaceTokenizer(tokenizer_path=_TOKENIZER_PATH)), + config, + ) + expected = hf.render(messages, add_generation_prompt=True) + actual = titan.render(messages, add_generation_prompt=True) + assert actual.token_ids == expected.token_ids + assert actual.is_content == expected.is_content + assert actual.sampled_mask == expected.sampled_mask + assert actual.message_indices == expected.message_indices diff --git a/torchtitan/experiments/rl/tests/test_rollout_worker.py b/torchtitan/experiments/rl/tests/test_rollout_worker.py index ec5fba8c09..d8a9dafc68 100644 --- a/torchtitan/experiments/rl/tests/test_rollout_worker.py +++ b/torchtitan/experiments/rl/tests/test_rollout_worker.py @@ -9,8 +9,13 @@ import asyncio from types import SimpleNamespace +from renderers import Qwen3RendererConfig + +from torchtitan.components.tokenizer import HuggingFaceTokenizer + from torchtitan.experiments.rl.actors.generator import SamplingConfig from torchtitan.experiments.rl.environment.token import TokenEnvOutput +from torchtitan.experiments.rl.renderer import RenderersLibraryConfig from torchtitan.experiments.rl.rollout import RolloutStatus from torchtitan.experiments.rl.rollout.rollouter import RolloutWorker from torchtitan.experiments.rl.rubrics import RubricOutput @@ -121,8 +126,11 @@ async def run() -> None: ) worker = _CustomWorker(worker_config) await worker.setup_async( - renderer_config=_Config("renderer"), - hf_assets_path="hf_assets_path", + tokenizer_config=HuggingFaceTokenizer.Config(), + renderer_config=RenderersLibraryConfig( + renderers_config=Qwen3RendererConfig(enable_thinking=False) + ), + hf_assets_path="tests/assets/tokenizer", ) group = await worker.run_group( generate_fn=generate_fn, @@ -142,7 +150,10 @@ async def run() -> None: assert [rollout.reward for rollout in group.rollouts] == [1.0, 2.0] assert [rollout.advantage for rollout in group.rollouts] == [10.0, 20.0] assert all(env.closed for env in token_env_config.envs) - assert token_env_config.renderers == ["renderer", "renderer"] + assert [type(r).__name__ for r in token_env_config.renderers] == [ + "Qwen3Renderer", + "Qwen3Renderer", + ] assert [call[1]["request_id"] for call in generate_fn.calls] == [ "group=7/rollout=0/turn=0", "group=7/rollout=1/turn=0", diff --git a/torchtitan/experiments/rl/tests/test_rollouter_pool.py b/torchtitan/experiments/rl/tests/test_rollouter_pool.py index a4dad20494..8bd55f5bbb 100644 --- a/torchtitan/experiments/rl/tests/test_rollouter_pool.py +++ b/torchtitan/experiments/rl/tests/test_rollouter_pool.py @@ -96,6 +96,7 @@ async def _setup( host = _ControllerHost(worker_mesh) monkeypatch.setattr(rollouter_module, "this_host", lambda: host) await rollouter.setup_async( + tokenizer_config="tokenizer_config", renderer_config="renderer_config", hf_assets_path="hf_assets_path", ) @@ -110,6 +111,7 @@ async def run() -> None: rollouter = _rollouter_without_datasets() await rollouter.setup_async( + tokenizer_config="tokenizer_config", renderer_config="renderer_config", hf_assets_path="hf_assets_path", ) @@ -125,6 +127,7 @@ async def run() -> None: } assert worker_mesh.actor_mesh.setup_async.calls == [ { + "tokenizer_config": "tokenizer_config", "renderer_config": "renderer_config", "hf_assets_path": "hf_assets_path", } diff --git a/torchtitan/experiments/rl/tests/test_shutdown.py b/torchtitan/experiments/rl/tests/test_shutdown.py index 90d3db1fa2..034729a5e0 100644 --- a/torchtitan/experiments/rl/tests/test_shutdown.py +++ b/torchtitan/experiments/rl/tests/test_shutdown.py @@ -175,11 +175,11 @@ class _StubConfig: rollout_recorder = RolloutSampleRecorder.Config() hf_assets_path = "./tests/assets/tokenizer" # __init__ builds these too; stub them so construction does no real work. + tokenizer = SimpleNamespace( + build=lambda *, tokenizer_path: SimpleNamespace(eos_id=0) + ) renderer = SimpleNamespace( - build=lambda *, tokenizer_path: SimpleNamespace( - get_stop_token_ids=lambda: [], - _tokenizer=SimpleNamespace(eos_token_id=0), - ) + build=lambda *, tokenizer: SimpleNamespace(get_stop_token_ids=lambda: []) ) # __init__ reads generator.sampling (a dataclass, for replace) + generator.debug.seed. generator = SimpleNamespace( diff --git a/torchtitan/experiments/rl/tests/test_token_env.py b/torchtitan/experiments/rl/tests/test_token_env.py new file mode 100644 index 0000000000..bf36db9802 --- /dev/null +++ b/torchtitan/experiments/rl/tests/test_token_env.py @@ -0,0 +1,96 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import asyncio + +from renderers import ParsedResponse + +from torchtitan.experiments.rl.environment.message import ( + MessageEnvInitOutput, + MessageEnvStepOutput, +) +from torchtitan.experiments.rl.environment.token import TokenEnv +from torchtitan.experiments.rl.rollout import RolloutStatus +from torchtitan.experiments.rl.types import Completion + +_TOOLS = [ + { + "type": "function", + "function": { + "name": "search", + "parameters": { + "type": "object", + "properties": {"query": {"type": "string"}}, + }, + }, + } +] + + +class _RecordingRenderer: + """Records the kwargs of each renderer call the env makes.""" + + def __init__(self) -> None: + self.calls: dict[str, dict] = {} + + def render_ids(self, messages, *, tools=None, add_generation_prompt=False): + self.calls["render_ids"] = {"tools": tools} + return [1, 2, 3] + + def parse_response(self, token_ids, *, tools=None): + self.calls["parse_response"] = {"token_ids": token_ids, "tools": tools} + return ParsedResponse(content="answer", reasoning_content=None, tool_calls=[]) + + def bridge_to_next_turn( + self, previous_prompt_ids, previous_completion_ids, new_messages, *, tools=None + ): + self.calls["bridge_to_next_turn"] = {"tools": tools} + return None + + def get_stop_token_ids(self): + return [] + + +class _ToolMessageEnv: + async def init(self) -> MessageEnvInitOutput: + return MessageEnvInitOutput( + init_prompt_messages=[{"role": "user", "content": "find bob"}], + tools=_TOOLS, + ) + + async def step(self, completion_message) -> MessageEnvStepOutput: + return MessageEnvStepOutput( + env_messages=[{"role": "tool", "name": "search", "content": "bob: found"}] + ) + + async def close(self) -> None: + pass + + +def test_env_passes_tools_to_render_parse_and_bridge() -> None: + # Tool schemas are part of the chat template, and XML-style tool parsers need them + # to type the arguments; every renderer call must see the same list. + renderer = _RecordingRenderer() + env = TokenEnv.Config().build(message_env=_ToolMessageEnv(), renderer=renderer) + + async def run(): + await env.init() + return await env.step( + Completion( + min_policy_version=0, + max_policy_version=0, + request_id="r0", + token_ids=[7, 8], + token_logprobs=[-0.1, -0.2], + finish_reason="stop", + ) + ) + + env_output = asyncio.run(run()) + assert env_output.status == RolloutStatus.ONGOING + assert renderer.calls["render_ids"]["tools"] == _TOOLS + assert renderer.calls["parse_response"] == {"token_ids": [7, 8], "tools": _TOOLS} + assert renderer.calls["bridge_to_next_turn"]["tools"] == _TOOLS