From 9f14b241a72c017ce476d01e6f6d7b76de9eee87 Mon Sep 17 00:00:00 2001 From: Eric Botti Date: Sun, 10 May 2026 14:11:45 -0400 Subject: [PATCH 1/5] Feat: opt-in renderer-based tokenization for SFT Adds `use_renderer` flag to SFTConfig, mirroring the RL path added in #2278. When enabled, SFTDataset tokenizes via `renderers.base.build_training_sample` (single render() + message_indices mask) instead of the incremental Jinja template path. Fixes silent multiturn drops and the Qwen3.5 system-only TemplateError crash for chat templates that render position-dependently. Default path is unchanged. Co-Authored-By: Claude Opus 4.7 --- .../src/prime_rl/configs/sft.py | 16 ++++++++++ src/prime_rl/trainer/sft/data.py | 30 ++++++++++++++----- src/prime_rl/trainer/sft/train.py | 21 +++++++++++-- 3 files changed, 57 insertions(+), 10 deletions(-) diff --git a/packages/prime-rl-configs/src/prime_rl/configs/sft.py b/packages/prime-rl-configs/src/prime_rl/configs/sft.py index 868ae2d973..8ce48745ae 100644 --- a/packages/prime-rl-configs/src/prime_rl/configs/sft.py +++ b/packages/prime-rl-configs/src/prime_rl/configs/sft.py @@ -6,6 +6,7 @@ from prime_rl.configs.shared import ( HeartbeatConfig, + RendererConfig, SlurmConfig, TrainerLogConfig, WandbConfig, @@ -172,6 +173,21 @@ class SFTConfig(BaseConfig): # The tokenizer configuration tokenizer: TokenizerConfig = TokenizerConfig() + # The renderer configuration (only used when use_renderer=True) + renderer: RendererConfig = RendererConfig() + + use_renderer: Annotated[ + bool, + Field( + description=( + "If True, tokenize SFT samples through the `renderers` library " + "(single render() + message_indices mask) instead of the default " + "`build_incremental_token_mask` path. Required for chat templates " + "that render position-dependently (e.g. Qwen3, Qwen3.5)." + ), + ), + ] = False + # The data configuration data: DataConfig = SFTDataConfig() diff --git a/src/prime_rl/trainer/sft/data.py b/src/prime_rl/trainer/sft/data.py index 6aa93cc0e3..93d1e20a77 100644 --- a/src/prime_rl/trainer/sft/data.py +++ b/src/prime_rl/trainer/sft/data.py @@ -127,6 +127,7 @@ def __init__( loss_mask_config: LossMaskConfig = LossMaskConfig(), max_examples: int | None = None, max_epochs: int | None = None, + renderer=None, ): super().__init__() self.logger = get_logger() @@ -139,6 +140,7 @@ def __init__( self.loss_mask_config = loss_mask_config self.max_examples = max_examples self.max_epochs = max_epochs + self.renderer = renderer if self.tokenizer is None: self.logger.warning("No tokenizer provided, will not process examples") @@ -206,18 +208,28 @@ def should_mask(message: dict) -> bool: case _: raise ValueError(f"Invalid message role: {message['role']}") - try: - input_ids, loss_mask = build_incremental_token_mask( - self.tokenizer, + if self.renderer is not None: + from renderers.base import build_training_sample + + input_ids, loss_mask = build_training_sample( + self.renderer, messages, role_to_mask=should_mask, tools=tools, - chat_template_kwargs=example.get("chat_template_kwargs", {}), - collapse_consecutive_tool_messages=True, ) - except IncrementalTokenizationError as e: - self.logger.warning(f"Skipping example {example.get('__index', '')}: {e}") - return None + else: + try: + input_ids, loss_mask = build_incremental_token_mask( + self.tokenizer, + messages, + role_to_mask=should_mask, + tools=tools, + chat_template_kwargs=example.get("chat_template_kwargs", {}), + collapse_consecutive_tool_messages=True, + ) + except IncrementalTokenizationError as e: + self.logger.warning(f"Skipping example {example.get('__index', '')}: {e}") + return None # If EOS token is not found, manually append it if not self.tokenizer.eos_token_id in input_ids: @@ -537,6 +549,7 @@ def setup_dataset( *, max_epochs: int | None = None, raw_dataset: Dataset | None = None, + renderer=None, ) -> StatefulIterableDataset: if config.type == "fake": return FakeDataset( @@ -554,6 +567,7 @@ def setup_dataset( loss_mask_config=config.loss_mask, non_dp_size=non_dp_size, max_epochs=max_epochs, + renderer=renderer, ) else: raise ValueError(f"Invalid dataset type: {config.type}") diff --git a/src/prime_rl/trainer/sft/train.py b/src/prime_rl/trainer/sft/train.py index 032f3b136b..e1109174e8 100644 --- a/src/prime_rl/trainer/sft/train.py +++ b/src/prime_rl/trainer/sft/train.py @@ -157,6 +157,18 @@ def train(config: SFTConfig): logger.info(f"Initializing tokenizer ({config.tokenizer})") tokenizer = setup_tokenizer(config.tokenizer) + renderer = None + if config.use_renderer: + from renderers.base import create_renderer + + renderer = create_renderer( + tokenizer, + renderer=config.renderer.name, + tool_parser=config.renderer.tool_parser, + reasoning_parser=config.renderer.reasoning_parser, + ) + logger.info(f"Initialized {type(renderer).__name__} for {config.tokenizer.name}") + # Set up the optimizer logger.info(f"Initializing optimizer ({config.optim})") optimizer = setup_optimizer( @@ -175,7 +187,7 @@ def train(config: SFTConfig): # Set up the dataset and dataloader logger.info(f"Initializing data ({config.data})") - dataset = setup_dataset(tokenizer, config.data, config.model.cp) + dataset = setup_dataset(tokenizer, config.data, config.model.cp, renderer=renderer) dataloader = setup_dataloader(dataset, config.data) dataiter = iter(dataloader) @@ -283,7 +295,12 @@ def run_eval_loop(data_iter): def run_validation(step: int) -> None: val_dataset = setup_dataset( - tokenizer, config.val.data, config.model.cp, max_epochs=1, raw_dataset=val_raw_dataset + tokenizer, + config.val.data, + config.model.cp, + max_epochs=1, + raw_dataset=val_raw_dataset, + renderer=renderer, ) val_dataloader = setup_dataloader(val_dataset, config.val.data) From 4bcde5de3f27f69027e959fe56bf1b5a288b1dd2 Mon Sep 17 00:00:00 2001 From: Eric Botti Date: Sun, 10 May 2026 14:40:19 -0400 Subject: [PATCH 2/5] validate use_renderer is not set with VLM config in SFT --- packages/prime-rl-configs/src/prime_rl/configs/sft.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/packages/prime-rl-configs/src/prime_rl/configs/sft.py b/packages/prime-rl-configs/src/prime_rl/configs/sft.py index 8ce48745ae..542edc1534 100644 --- a/packages/prime-rl-configs/src/prime_rl/configs/sft.py +++ b/packages/prime-rl-configs/src/prime_rl/configs/sft.py @@ -373,6 +373,15 @@ def dont_do_massive_traces(self): ) return self + @model_validator(mode="after") + def validate_renderer_vs_vlm(self): + if self.use_renderer and self.model.vlm is not None: + raise ValueError( + "use_renderer is not supported for VLMs. The renderer tokenizes " + "text-only message dicts client-side and cannot handle image inputs." + ) + return self + @model_validator(mode="after") def validate_lora_adapter_saving(self): if self.ckpt and self.ckpt.weights and self.ckpt.weights.save_adapter_separately: From e61b6dba452726bb07a7df75690a39a7b997d0a5 Mon Sep 17 00:00:00 2001 From: Eric Botti Date: Sun, 10 May 2026 14:51:59 -0400 Subject: [PATCH 3/5] validate no renderer specific args are set when not using renderer. Copied from orch config --- .../src/prime_rl/configs/sft.py | 22 +++++++++++++++++++ 1 file changed, 22 insertions(+) diff --git a/packages/prime-rl-configs/src/prime_rl/configs/sft.py b/packages/prime-rl-configs/src/prime_rl/configs/sft.py index 542edc1534..6ec99380a6 100644 --- a/packages/prime-rl-configs/src/prime_rl/configs/sft.py +++ b/packages/prime-rl-configs/src/prime_rl/configs/sft.py @@ -382,6 +382,28 @@ def validate_renderer_vs_vlm(self): ) return self + @model_validator(mode="after") + def validate_renderer_args(self): + if self.use_renderer: + return self + + renderer_args_set = [] + if self.renderer.name != "auto": + renderer_args_set.append(f"renderer.name={self.renderer.name!r}") + if self.renderer.tool_parser is not None: + renderer_args_set.append(f"renderer.tool_parser={self.renderer.tool_parser!r}") + if self.renderer.reasoning_parser is not None: + renderer_args_set.append(f"renderer.reasoning_parser={self.renderer.reasoning_parser!r}") + if self.renderer.pool_size is not None: + renderer_args_set.append(f"renderer.pool_size={self.renderer.pool_size!r}") + + if renderer_args_set: + raise ValueError( + "Renderer-specific args set without use_renderer=True: " + f"{', '.join(renderer_args_set)}. Either enable the renderer or remove these knobs." + ) + return self + @model_validator(mode="after") def validate_lora_adapter_saving(self): if self.ckpt and self.ckpt.weights and self.ckpt.weights.save_adapter_separately: From 26ecb0610bd53a7db18168e7a58490587fb376df Mon Sep 17 00:00:00 2001 From: Eric Botti Date: Sun, 10 May 2026 18:03:17 -0400 Subject: [PATCH 4/5] Test: add renderer path test for multiturn loss mask with Qwen3 Co-Authored-By: Claude Sonnet 4.6 --- tests/unit/train/sft/test_sft_dataset.py | 24 ++++++++++++++++++++++++ 1 file changed, 24 insertions(+) diff --git a/tests/unit/train/sft/test_sft_dataset.py b/tests/unit/train/sft/test_sft_dataset.py index b8465e59d7..ca38d12d3e 100644 --- a/tests/unit/train/sft/test_sft_dataset.py +++ b/tests/unit/train/sft/test_sft_dataset.py @@ -199,6 +199,30 @@ def test_multiturn_loss_mask(): print_sample(sample["input_ids"], sample["loss_mask"], tokenizer) +def test_multiturn_loss_mask_use_renderer(): + """Renderer path handles position-dependent chat templates that break build_incremental_token_mask.""" + from renderers.base import create_renderer + + dataset = Dataset.from_list( + [ + { + "prompt": [{"role": "system", "content": "System"}, {"role": "user", "content": "Prompt 0"}], + "completion": [ + {"role": "assistant", "content": "Reply 0"}, + {"role": "user", "content": "Prompt 1"}, + {"role": "assistant", "content": "Reply 1"}, + ], + } + ] + ) + tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen3-0.6B") # Directly tests non PrimeIntellect Qwen3 + renderer = create_renderer(tokenizer) + sft_dataset = SFTDataset(dataset, tokenizer=tokenizer, renderer=renderer, max_examples=1) + sample = next(iter(sft_dataset)) + assert sum(sample["loss_mask"]) > 0 + print_sample(sample["input_ids"], sample["loss_mask"], tokenizer) + + def test_multiturn_loss_mask_with_tools(): tool_example = { "prompt": [ From 9bb1180bf11ccefee9d25f5eaed3dd69b6fdf672 Mon Sep 17 00:00:00 2001 From: hallerite Date: Thu, 14 May 2026 15:00:07 +0000 Subject: [PATCH 5/5] =?UTF-8?q?fix(sft):=20tighten=20renderer=20integratio?= =?UTF-8?q?n=20=E2=80=94=20typing,=20no=20DefaultRenderer=20fallback,=20re?= =?UTF-8?q?ject=20pool=5Fsize?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Type `renderer: Renderer | None` in `SFTDataset` / `setup_dataset`; lift `renderers` imports to module scope. - Raise in `train()` when `use_renderer=True` resolves to `DefaultRenderer`: the fallback uses incremental `apply_chat_template` and does NOT fix the position-dependent template bug `use_renderer` is meant to solve. - Warn once per dataset when an example carries `chat_template_kwargs` while the renderer path is on (renderers don't forward template kwargs). - Reject `renderer.pool_size` unconditionally in `SFTConfig`. SFT tokenizes synchronously (`num_workers=0`) and already gets one renderer per DP rank — an in-process pool adds nothing. Comment explains the implicit-pool framing. - Drop the new renderer test (`assert sum(loss_mask) > 0` was too weak to earn its keep). Co-Authored-By: Claude Opus 4.7 (1M context) --- .../src/prime_rl/configs/sft.py | 22 +++++++++++++++-- src/prime_rl/trainer/sft/data.py | 15 +++++++++--- src/prime_rl/trainer/sft/train.py | 12 ++++++++-- tests/unit/train/sft/test_sft_dataset.py | 24 ------------------- 4 files changed, 42 insertions(+), 31 deletions(-) diff --git a/packages/prime-rl-configs/src/prime_rl/configs/sft.py b/packages/prime-rl-configs/src/prime_rl/configs/sft.py index 6ec99380a6..84ee13018d 100644 --- a/packages/prime-rl-configs/src/prime_rl/configs/sft.py +++ b/packages/prime-rl-configs/src/prime_rl/configs/sft.py @@ -384,6 +384,26 @@ def validate_renderer_vs_vlm(self): @model_validator(mode="after") def validate_renderer_args(self): + # pool_size is orchestrator-only. An in-process renderer pool exists + # to amortize tokenization across concurrent rollouts in the + # orchestrator (many async requests render at once, HF fast + # tokenizers release the GIL during Rust encoding, so a pool of N + # tokenizer copies parallelizes well). SFT has no such concurrency: + # the StatefulDataLoader is constructed with num_workers=0, so the + # main process tokenizes one example at a time, between training + # steps. Across DP, each rank already owns its own renderer — an + # implicit pool of size world_size. Pooling within a rank gives + # nothing on top of that. Reject so callers don't silently set a + # knob that does nothing; if SFT tokenization ever becomes a + # bottleneck the fix is num_workers on the dataloader, not a pool. + if self.renderer.pool_size is not None: + raise ValueError( + f"renderer.pool_size={self.renderer.pool_size!r} is only used by the orchestrator. " + "SFT tokenizes synchronously (num_workers=0) and already gets one renderer per DP " + "rank — an in-process pool adds nothing. If tokenization is a bottleneck, raise " + "num_workers on the dataloader instead." + ) + if self.use_renderer: return self @@ -394,8 +414,6 @@ def validate_renderer_args(self): renderer_args_set.append(f"renderer.tool_parser={self.renderer.tool_parser!r}") if self.renderer.reasoning_parser is not None: renderer_args_set.append(f"renderer.reasoning_parser={self.renderer.reasoning_parser!r}") - if self.renderer.pool_size is not None: - renderer_args_set.append(f"renderer.pool_size={self.renderer.pool_size!r}") if renderer_args_set: raise ValueError( diff --git a/src/prime_rl/trainer/sft/data.py b/src/prime_rl/trainer/sft/data.py index 93d1e20a77..78bea8f0da 100644 --- a/src/prime_rl/trainer/sft/data.py +++ b/src/prime_rl/trainer/sft/data.py @@ -6,6 +6,7 @@ import torch from datasets import Dataset, interleave_datasets, load_dataset from jaxtyping import Bool, Int +from renderers.base import Renderer, build_training_sample from torch import Tensor from torch.distributed.checkpoint.stateful import Stateful from torch.utils.data import IterableDataset, get_worker_info @@ -127,7 +128,7 @@ def __init__( loss_mask_config: LossMaskConfig = LossMaskConfig(), max_examples: int | None = None, max_epochs: int | None = None, - renderer=None, + renderer: Renderer | None = None, ): super().__init__() self.logger = get_logger() @@ -141,6 +142,7 @@ def __init__( self.max_examples = max_examples self.max_epochs = max_epochs self.renderer = renderer + self._warned_chat_template_kwargs = False if self.tokenizer is None: self.logger.warning("No tokenizer provided, will not process examples") @@ -209,7 +211,14 @@ def should_mask(message: dict) -> bool: raise ValueError(f"Invalid message role: {message['role']}") if self.renderer is not None: - from renderers.base import build_training_sample + if example.get("chat_template_kwargs") and not self._warned_chat_template_kwargs: + self.logger.warning( + "Example carries chat_template_kwargs but use_renderer=True; " + "renderers don't forward chat_template_kwargs (model-specific " + "renderers bake their template behavior in). These kwargs will " + "be ignored. Further warnings suppressed for this dataset." + ) + self._warned_chat_template_kwargs = True input_ids, loss_mask = build_training_sample( self.renderer, @@ -549,7 +558,7 @@ def setup_dataset( *, max_epochs: int | None = None, raw_dataset: Dataset | None = None, - renderer=None, + renderer: Renderer | None = None, ) -> StatefulIterableDataset: if config.type == "fake": return FakeDataset( diff --git a/src/prime_rl/trainer/sft/train.py b/src/prime_rl/trainer/sft/train.py index e1109174e8..ace4158f2e 100644 --- a/src/prime_rl/trainer/sft/train.py +++ b/src/prime_rl/trainer/sft/train.py @@ -4,6 +4,8 @@ from contextlib import nullcontext from datetime import timedelta +from renderers.base import create_renderer +from renderers.default import DefaultRenderer from ring_flash_attn import substitute_hf_flash_attn from torch.nn import CrossEntropyLoss @@ -159,14 +161,20 @@ def train(config: SFTConfig): renderer = None if config.use_renderer: - from renderers.base import create_renderer - renderer = create_renderer( tokenizer, renderer=config.renderer.name, tool_parser=config.renderer.tool_parser, reasoning_parser=config.renderer.reasoning_parser, ) + if isinstance(renderer, DefaultRenderer): + raise ValueError( + f"use_renderer=True for {config.tokenizer.name!r} resolved to DefaultRenderer. " + "DefaultRenderer falls back to incremental apply_chat_template and does NOT " + "fix position-dependent chat templates — the bug use_renderer is meant to solve. " + "Either use a model with a hand-coded renderer (see renderers.base.MODEL_RENDERER_MAP), " + "set [renderer] name= explicitly, or set use_renderer=false." + ) logger.info(f"Initialized {type(renderer).__name__} for {config.tokenizer.name}") # Set up the optimizer diff --git a/tests/unit/train/sft/test_sft_dataset.py b/tests/unit/train/sft/test_sft_dataset.py index ca38d12d3e..b8465e59d7 100644 --- a/tests/unit/train/sft/test_sft_dataset.py +++ b/tests/unit/train/sft/test_sft_dataset.py @@ -199,30 +199,6 @@ def test_multiturn_loss_mask(): print_sample(sample["input_ids"], sample["loss_mask"], tokenizer) -def test_multiturn_loss_mask_use_renderer(): - """Renderer path handles position-dependent chat templates that break build_incremental_token_mask.""" - from renderers.base import create_renderer - - dataset = Dataset.from_list( - [ - { - "prompt": [{"role": "system", "content": "System"}, {"role": "user", "content": "Prompt 0"}], - "completion": [ - {"role": "assistant", "content": "Reply 0"}, - {"role": "user", "content": "Prompt 1"}, - {"role": "assistant", "content": "Reply 1"}, - ], - } - ] - ) - tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen3-0.6B") # Directly tests non PrimeIntellect Qwen3 - renderer = create_renderer(tokenizer) - sft_dataset = SFTDataset(dataset, tokenizer=tokenizer, renderer=renderer, max_examples=1) - sample = next(iter(sft_dataset)) - assert sum(sample["loss_mask"]) > 0 - print_sample(sample["input_ids"], sample["loss_mask"], tokenizer) - - def test_multiturn_loss_mask_with_tools(): tool_example = { "prompt": [