Skip to content
Merged
Show file tree
Hide file tree
Changes from 24 commits
Commits
Show all changes
40 commits
Select commit Hold shift + click to select a range
6913120
[RL] Add Batcher as Configurable, support [B, L] microbatches
wwwjn May 18, 2026
7c7c9f7
rebase
wwwjn May 18, 2026
8e4fc60
remove position in loss calc
wwwjn May 19, 2026
dc690a1
update algorithm
wwwjn May 19, 2026
40ef4c5
update rebase
wwwjn May 19, 2026
767022f
update trainer to using new TrainBatch
wwwjn May 20, 2026
50b5572
update loss
wwwjn May 20, 2026
86ed252
update names and configs
wwwjn May 20, 2026
31dcd8b
update batcher with waste metrics
wwwjn May 20, 2026
7dad6f5
update configs
wwwjn May 22, 2026
d440e46
address comments
wwwjn May 26, 2026
1523db5
v8 datatypes + env protocol
May 27, 2026
312c292
v8 rubric + rollouts refactor
May 27, 2026
14aaeca
v8 Task surface refactor + style pass
May 28, 2026
6292f1d
v8 docstring/style pass + truncation fix + renderer tokenizer
May 29, 2026
e9af27a
Merge upstream/main (landed batcher PR) into v8 branch
May 29, 2026
85880db
Fix NoReduce collision: rollout/group_failures sums across collection…
May 29, 2026
9243ba1
Adopt renderers typed-config API (create_renderer(tokenizer, config))
May 29, 2026
701b4d1
Revert unrelated merge drift to upstream/main
May 29, 2026
11003ba
RL rollout loop cleanup: RolloutGroup, per-rollout step, status enum
May 29, 2026
5044e86
RL example default: thinking on, max_tokens=700; tidy rollout types
May 29, 2026
2dcd0f9
RL env/types review pass (docs 69+70): renames, env reward, renderer …
May 29, 2026
7a72a41
Fix _build_episodes group-skip + rollout_to_episode text
May 29, 2026
0cf8460
Docstring concision pass + rename envs/ -> env_types/
May 29, 2026
df376bb
Make Task an ABC; base Config owns rubric, subclass builds it
May 30, 2026
d3ba2df
[rl] v8: message/token env types, Configurable rubric, Task ABC
Jun 2, 2026
de7c887
Refine RL rollout protocol
Jun 3, 2026
c2efd2b
Rename Task.Config.env_config -> env_wrapper_cfg
Jun 3, 2026
8b597ba
Merge remote-tracking branch 'upstream/main' into v8-datatypes-env-pr…
Jun 3, 2026
79c3903
[rl] address PR review (round 2): env/rollout naming + cleanups
Jun 3, 2026
621e3cf
[rl] PR review round 2: shape hints + Task->Rollouter rename
Jun 4, 2026
3c69ae6
[rl] PR review round 2: env/rollout/examples package layout
Jun 4, 2026
c77d47f
[rl] renderer: None-default knobs; build the typed config via create_…
Jun 4, 2026
75e69d4
Merge remote-tracking branch 'upstream/main' into v8-datatypes-env-pr…
Jun 4, 2026
80104f5
[rl] PR review round 3: naming (completion_message, MessageEnv*Output…
Jun 5, 2026
4bf9b33
[rl] PR review round 3: get_train_sample -> get_training_sample (tian…
Jun 5, 2026
9d91d4c
Merge branch 'main' of https://github.com/pytorch/torchtitan into v8-…
Jun 5, 2026
d33903e
[rl] token: simplify the history-edit TODO comment
Jun 5, 2026
89f8f85
[rl] rollout_to_episode: TODO for branching when multi-turn history d…
Jun 5, 2026
04c47c2
[rl] CI: install renderers in RL integration test workflows
Jun 5, 2026
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
3 changes: 2 additions & 1 deletion torchtitan/experiments/rl/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -27,11 +27,12 @@ uv venv --python 3.12 titan-rl
source titan-rl/bin/activate
```

1. Install Monarch and TorchStore from main:
1. Install Monarch, TorchStore, and Renderers from main:
```bash
uv pip install torchmonarch==0.4.1
uv pip install --no-deps "git+https://github.com/meta-pytorch/torchstore.git@main"
uv pip install pygtrie portpicker
uv pip install "git+https://github.com/PrimeIntellect-ai/renderers.git@main"

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.

would this be used for sft as well?

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.

Do we need to depend on renders lastest main?

```

2. Install Flash Attention 3 kernels:
Expand Down
5 changes: 4 additions & 1 deletion torchtitan/experiments/rl/actors/generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -163,6 +163,9 @@ class SamplingConfig:
max_tokens: int = 100
"""Maximum number of tokens to generate per completion."""

stop_token_ids: list[int] = field(default_factory=list)
"""Role-boundary stop tokens from the renderer (e.g. Qwen3 `<|im_end|>`)."""


class VLLMGenerator(Actor, Configurable):
"""
Expand Down Expand Up @@ -411,6 +414,7 @@ async def generate(
top_p=_sampling_config.top_p,
max_tokens=_sampling_config.max_tokens,
n=_sampling_config.n,
stop_token_ids=list(_sampling_config.stop_token_ids) or None,

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.

Suggested change
stop_token_ids=list(_sampling_config.stop_token_ids) or None,
stop_token_ids=_sampling_config.stop_token_ids or None,

seed=self.config.debug.seed,
logprobs=1,
output_kind=RequestOutputKind.FINAL_ONLY,
Expand Down Expand Up @@ -457,7 +461,6 @@ async def generate(
Completion(
policy_version=self.policy_version,
prompt_idx=prompt_idx,
text=sample.text,
token_ids=sample.token_ids,
token_logprobs=per_token_logprobs,
finish_reason=sample.finish_reason,
Expand Down
2 changes: 1 addition & 1 deletion torchtitan/experiments/rl/actors/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -455,7 +455,7 @@ async def forward_backward(
async def optim_step(self) -> OptimStepOutput:
"""Clip gradients, step optimizer + LR scheduler, return updated state."""
# TODO: Accept optional optimizer params (e.g. learning rate)
# to allow controller-owned schedules (see Tinker API).
# to allow controller-owned schedules.

# capture LR before step
current_lrs = self.lr_schedulers.schedulers[0].get_last_lr()
Expand Down
45 changes: 23 additions & 22 deletions torchtitan/experiments/rl/config_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,8 @@
from torchtitan.experiments.rl.batcher import BatchConfig, Batcher
from torchtitan.experiments.rl.grpo import GRPOLoss, RLTrainer
from torchtitan.experiments.rl.observability.metrics import MetricsProcessor
from torchtitan.experiments.rl.sum_digits import SumDigitsEnv
from torchtitan.experiments.rl.renderer import RendererConfig
from torchtitan.experiments.rl.tasks.sum_digits import SumDigitsDataset, SumDigitsTask
from torchtitan.models.qwen3 import model_registry


Expand All @@ -39,10 +40,10 @@ def rl_grpo_qwen3_0_6b() -> RLTrainer.Config:
num_prompts_per_step=5,
num_validation_samples=20,
compile=CompileConfig(enable=True, backend="aot_eager"),
env=SumDigitsEnv.Config(seed=42, correctness_reward=1.0, format_reward=0.3),
validation_env=SumDigitsEnv.Config(
seed=99, correctness_reward=1.0, format_reward=0.3
),
train_dataset=SumDigitsDataset.Config(seed=42),
validation_dataset=SumDigitsDataset.Config(seed=99),

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 was thinking should we change validation -> eval, I see eval is more widely used in RL

validation sounds more like supervised learning

@felipemello1 felipemello1 Jun 1, 2026

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

no strong opinions. I dont mind changing. But i think that we should use the same terminology across the codebase.

tasks={"sum_digits": SumDigitsTask.Config()},
renderer=RendererConfig(name="qwen3", enable_thinking=True),

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.

Why render related to a model name? Is it because it will use tokenizer path? Currently hf_assets_path is only used to local tokenizer. Can we move the tokenizer and tokenizer path under render?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Why render related to a model name? Is it because it will use tokenizer path?
yes

Currently hf_assets_path is only used to local tokenizer. Can we move the tokenizer and tokenizer path under render?
Not yet. They dont support users providing their own tokenizer yet, but they will soon.

metrics=MetricsProcessor.Config(enable_wandb=True),
batcher=Batcher.Config(
batch=BatchConfig(local_batch_size=2, global_batch_size=8, seq_len=2048),
Expand Down Expand Up @@ -80,7 +81,7 @@ def rl_grpo_qwen3_0_6b() -> RLTrainer.Config:
n=group_size,
temperature=0.8,
top_p=0.95,
max_tokens=100,
max_tokens=700,
),
),
)
Expand All @@ -96,10 +97,10 @@ def rl_grpo_qwen3_1_7b() -> RLTrainer.Config:
num_prompts_per_step=5,
num_validation_samples=20,
compile=CompileConfig(enable=True, backend="aot_eager"),
env=SumDigitsEnv.Config(seed=42, correctness_reward=1.0, format_reward=0.3),
validation_env=SumDigitsEnv.Config(
seed=99, correctness_reward=1.0, format_reward=0.3
),
train_dataset=SumDigitsDataset.Config(seed=42),
validation_dataset=SumDigitsDataset.Config(seed=99),
tasks={"sum_digits": SumDigitsTask.Config()},
renderer=RendererConfig(name="qwen3", enable_thinking=True),
metrics=MetricsProcessor.Config(enable_wandb=True),
batcher=Batcher.Config(
batch=BatchConfig(local_batch_size=2, global_batch_size=8, seq_len=2048),
Expand Down Expand Up @@ -138,7 +139,7 @@ def rl_grpo_qwen3_1_7b() -> RLTrainer.Config:
n=group_size,
temperature=0.8,
top_p=0.95,
max_tokens=100,
max_tokens=700,

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.

Will need to change batcher max_seq_length accordingly, otherwise all the samples will be dropped because they are longer than max_seq_length

),
),
)
Expand All @@ -154,10 +155,10 @@ def rl_grpo_qwen3_14b() -> RLTrainer.Config:
num_prompts_per_step=5,
num_validation_samples=20,
compile=CompileConfig(enable=True, backend="aot_eager"),
env=SumDigitsEnv.Config(seed=42, correctness_reward=1.0, format_reward=0.3),
validation_env=SumDigitsEnv.Config(
seed=99, correctness_reward=1.0, format_reward=0.3
),
train_dataset=SumDigitsDataset.Config(seed=42),
validation_dataset=SumDigitsDataset.Config(seed=99),
tasks={"sum_digits": SumDigitsTask.Config()},
renderer=RendererConfig(name="qwen3", enable_thinking=True),
metrics=MetricsProcessor.Config(enable_wandb=True),
batcher=Batcher.Config(
batch=BatchConfig(local_batch_size=2, global_batch_size=8, seq_len=2048),
Expand Down Expand Up @@ -195,14 +196,14 @@ def rl_grpo_qwen3_14b() -> RLTrainer.Config:
n=group_size,
temperature=0.8,
top_p=0.95,
max_tokens=100,
max_tokens=700,
),
),
)


def rl_grpo_qwen3_0_6b_batch_invariant() -> RLTrainer.Config:
"""On-policy GRPO config for Qwen3-0.6B under same parallelism (4 GPUs: 2 gen + 2 train).
"""On-policy GRPO config for Qwen3-0.6B (4 GPUs: 2 gen + 2 train).

Enables deterministic + batch-invariant mode for true on-policy RL training.
"""
Expand All @@ -215,10 +216,10 @@ def rl_grpo_qwen3_0_6b_batch_invariant() -> RLTrainer.Config:
num_prompts_per_step=5,
num_validation_samples=20,
compile=CompileConfig(enable=True, backend="aot_eager"),
env=SumDigitsEnv.Config(seed=42, correctness_reward=1.0, format_reward=0.3),
validation_env=SumDigitsEnv.Config(
seed=99, correctness_reward=1.0, format_reward=0.3
),
train_dataset=SumDigitsDataset.Config(seed=42),
validation_dataset=SumDigitsDataset.Config(seed=99),
tasks={"sum_digits": SumDigitsTask.Config()},
renderer=RendererConfig(name="qwen3", enable_thinking=True),
metrics=MetricsProcessor.Config(enable_wandb=True),
batcher=Batcher.Config(
batch=BatchConfig(local_batch_size=2, global_batch_size=8, seq_len=2048),
Expand Down Expand Up @@ -260,7 +261,7 @@ def rl_grpo_qwen3_0_6b_batch_invariant() -> RLTrainer.Config:
n=group_size,
temperature=0.8,
top_p=0.95,
max_tokens=100,
max_tokens=700,
),
debug=batch_invariant_config,
),
Expand Down
25 changes: 25 additions & 0 deletions torchtitan/experiments/rl/env_types/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
# 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.

from torchtitan.experiments.rl.env_types.message_env import (
MessageEnv,
ResetOutput,
StepOutput,
)
from torchtitan.experiments.rl.env_types.renderer_env import (
RendererEnv,
RendererEnvConfig,
TokenizedStepOutput,
)

__all__ = [
"RendererEnvConfig",
"MessageEnv",
"ResetOutput",
"StepOutput",
"RendererEnv",
"TokenizedStepOutput",
]
86 changes: 86 additions & 0 deletions torchtitan/experiments/rl/env_types/message_env.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,86 @@
# 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.

from __future__ import annotations

import abc
from dataclasses import dataclass, field

from renderers import Message, ToolSpec


@dataclass(kw_only=True, slots=True)
class ResetOutput:

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.

can this be special subclass of StepOutput?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

I was also a bit unhappy with these classes, but i didnt really want to subclass. Maybe i could have the StepOutput be a field in it. Let me check a few options.

"""Initial messages + tool specs from `MessageEnv.reset`."""

messages: list[Message] # [M_initial]
"""The set of messages that form an initial prompt for the model."""

tools: list[ToolSpec] = field(default_factory=list) # [K_tools]
"""Tool schemas exposed to the model. Empty for tool-less envs."""


@dataclass(kw_only=True, slots=True)
class StepOutput:
"""Env response to one parsed assistant message.

`done` ends the conversation; `RendererEnv` maps it to a `RolloutStatus`
and detects truncation/errors. The env may attach `reward_components`
(e.g. tool-call success); the rubric decides how to use them.
"""

messages: list[Message] = field(default_factory=list) # [M_env]

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.

noob question: why it's a list (instead of this step's single message)?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

The prompt can start from an ongoing conversation, for example.

But, a simple case it: ["system_msg", "user_message"]. Thats already two.

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.

let's docstring this

"""Env-appended messages (tool / user replies). Empty when the rollout
terminates with no follow-up."""

done: bool = False
"""`True` ends the rollout."""

reward_components: dict[str, float] = field(default_factory=dict)

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.

components sounds confusing, does rewards work?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

i think that rewards is also not obvious. Let me think of a few other options.

"""Optional reward signal the env provides for this step; the rubric decides
whether and how to use it. Empty if the env scores nothing."""

def __post_init__(self) -> None:
# env replies are tool/user turns; the assistant turn comes from the model
if any(m.get("role") == "assistant" for m in self.messages):
raise ValueError("StepOutput.messages may not contain assistant messages")


class MessageEnv(abc.ABC):
"""User-written env in message space. Implement `reset` + `step_message`;
`RendererEnv` wraps it with a `Renderer` for tokens.

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.

what does "a Renderer for tokens" mean?

@felipemello1 felipemello1 Jun 1, 2026

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

MessageEnv operate in messages, i.e. strings. No tokens.
RendererEnv takes tokens from generator and transforms into messages
It takes messages from envs and transforms into tokens for generator

I can rewrite to make it more obvious

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.

transforms into images

?

It takes messages from envs and transforms into tokens for generator

sounds confusing when you say "RendererEnv takes messages from envs"

@felipemello1 felipemello1 Jun 1, 2026

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

transforms into messages** haha

sounds confusing

why?

Not real APIs or precise flow -- just sharing the main idea

tokens_completion = generator.generate(token_input)
renderer_token_output = RendererEnv.step(tokens_completion)
token_input = renderer_token_output

class RendererEnv()
	def step(tokens_completion):
		assistant_message = renderer.decode(tokens_completion)
		env_message_output = self.MessageEnv(assistant_message)
		renderer_token_output = renderer.encode(env_message_output)
		return renderer_token_output


Example:

class SumDigitsEnv(MessageEnv):
async def reset(self) -> ResetOutput:
return ResetOutput(messages=[{"role": "user", "content": "sum [1, 2]"}])

async def step_message(self, msg: Message) -> StepOutput:
return StepOutput(done=True) # single-turn; rubric scores later
"""

@abc.abstractmethod
async def reset(self) -> ResetOutput:
"""Return the initial conversation + tools for prompt rendering."""

@abc.abstractmethod
async def step_message(self, msg: Message) -> StepOutput:

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.

function signature could be improved

  • is msg always from assistant?
  • MessageEnv.step_message(msg) sounds confusing

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

"is msg always from assistant?"
thats that expected, but i dont know if there is some crazy setup where this is not true, so i wouldnt assume.

"MessageEnv.step_message(msg) "
I will think of a few other options

"""Return the env's reply to an assistant message.

`RendererEnv` handles finish_reason / length / parse failures before
calling this, so the env only sees a successfully parsed message.

Args:
msg: Parsed assistant message (`content`, optional
`reasoning_content`, optional `tool_calls`).

Returns:
`StepOutput` with env reply messages.
"""

async def close(self) -> None:
"""Release env-owned resources. Default no-op; idempotent."""
Loading
Loading