Skip to content
Closed
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
65 changes: 65 additions & 0 deletions packages/prime-rl-configs/src/prime_rl/configs/sft.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@

from prime_rl.configs.shared import (
HeartbeatConfig,
RendererConfig,
SlurmConfig,
TrainerLogConfig,
WandbConfig,
Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -357,6 +373,55 @@ 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_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

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 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:
Expand Down
39 changes: 31 additions & 8 deletions src/prime_rl/trainer/sft/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -127,6 +128,7 @@ def __init__(
loss_mask_config: LossMaskConfig = LossMaskConfig(),
max_examples: int | None = None,
max_epochs: int | None = None,
renderer: Renderer | None = None,
):
super().__init__()
self.logger = get_logger()
Expand All @@ -139,6 +141,8 @@ def __init__(
self.loss_mask_config = loss_mask_config
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")
Expand Down Expand Up @@ -206,18 +210,35 @@ 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:
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,
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:
Expand Down Expand Up @@ -537,6 +558,7 @@ def setup_dataset(
*,
max_epochs: int | None = None,
raw_dataset: Dataset | None = None,
renderer: Renderer | None = None,
) -> StatefulIterableDataset:
if config.type == "fake":
return FakeDataset(
Expand All @@ -554,6 +576,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}")
Expand Down
29 changes: 27 additions & 2 deletions src/prime_rl/trainer/sft/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -157,6 +159,24 @@ def train(config: SFTConfig):
logger.info(f"Initializing tokenizer ({config.tokenizer})")
tokenizer = setup_tokenizer(config.tokenizer)

renderer = None
if config.use_renderer:
renderer = create_renderer(
tokenizer,
renderer=config.renderer.name,
tool_parser=config.renderer.tool_parser,
reasoning_parser=config.renderer.reasoning_parser,
)
Comment thread
cursor[bot] marked this conversation as resolved.
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=<hand-coded renderer> explicitly, or set use_renderer=false."
)
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(
Expand All @@ -175,7 +195,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)

Expand Down Expand Up @@ -283,7 +303,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)

Expand Down
Loading