-
Notifications
You must be signed in to change notification settings - Fork 930
[RL] - MessageEnv, Rollout types, Rubric, Renderer #3453
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from 33 commits
6913120
7c7c9f7
8e4fc60
dc690a1
40ef4c5
767022f
50b5572
86ed252
31dcd8b
7dad6f5
d440e46
1523db5
312c292
14aaeca
6292f1d
e9af27a
85880db
9243ba1
701b4d1
11003ba
5044e86
2dcd0f9
7a72a41
0cf8460
df376bb
d3ba2df
de7c887
c2efd2b
8b597ba
79c3903
621e3cf
3c69ae6
c77d47f
75e69d4
80104f5
4bf9b33
9d91d4c
d33903e
89f8f85
04c47c2
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -151,9 +151,6 @@ def get_vllm_compilation_config( | |
| class SamplingConfig: | ||
| """Sampling parameters passed to vLLM's SamplingParams.""" | ||
|
|
||
| n: int = 8 | ||
| """Number of completions to generate per prompt (vLLM SamplingParams.n).""" | ||
|
|
||
| temperature: float = 0.8 | ||
| """Sampling temperature. 0.0 = greedy, higher = more random.""" | ||
|
|
||
|
|
@@ -164,6 +161,7 @@ class SamplingConfig: | |
| """Maximum number of tokens to generate per completion.""" | ||
|
|
||
|
|
||
| # TODO(naming): rename VLLMGenerator -> InferenceEngine. | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I think after we name Rollouter, this is less necessary. |
||
| class VLLMGenerator(Actor, Configurable): | ||
| """ | ||
| Generates rollouts using vLLM engine. | ||
|
|
@@ -257,6 +255,7 @@ def __init__( | |
| compile_config: CompileConfig, | ||
| max_num_seqs: int, | ||
| output_dir: str, | ||
| stop_token_ids: list[int], | ||
| ): | ||
| init_logger() | ||
| sl.init_structured_logger( | ||
|
|
@@ -274,9 +273,13 @@ def __init__( | |
| # the upper bound for concurrent sequences, determines KV-cache | ||
| # block allocation (and therefore GPU memory usage), and bounds | ||
| # the CUDA graph capture sizes. Always computed by the caller | ||
| # (RLTrainer) as num_prompts_per_step * sampling.n. | ||
| # (RLTrainer) as num_groups_per_rollout_batch * group_size. | ||
| self._max_num_seqs = max_num_seqs | ||
|
|
||
| # Renderer role-boundary stop tokens (e.g. Qwen3 `<|im_end|>`), injected by the | ||
| # controller; | ||
| self._stop_token_ids = stop_token_ids | ||
|
|
||
| # Register TorchTitan model + parser with vLLM | ||
| registry_to_vllm( | ||
| model_spec, | ||
|
|
@@ -379,17 +382,19 @@ async def generate( | |
| self, | ||
| tokenized_prompts: list[list[int]], | ||
| *, | ||
| request_ids: list[str], | ||
| sampling_config: SamplingConfig | None = None, | ||
| metrics_prefix: str = "generator", | ||
| ) -> tuple[list[Completion], list[m.Metric]]: | ||
| """Generate completions and generator metrics for tokenized prompts. | ||
|
|
||
| Takes ``tokenized_prompts`` as ``[num_prompts][prompt_tokens]``. | ||
| Returns completions as ``[num_prompts * n]`` plus generator metrics, | ||
| where ``n`` is the resolved ``SamplingConfig.n`` completions per prompt. | ||
| Returns completions in the same order as ``request_ids`` plus generator | ||
| metrics. | ||
|
|
||
| Args: | ||
| tokenized_prompts: Tokenized prompts shaped ``[num_prompts][prompt_tokens]``. | ||
| request_ids: One id per prompt echoed on each ``Completion.request_id``. | ||
| sampling_config: Optional per-call override for the generator's | ||
| default SamplingConfig. ``seed`` always comes from | ||
| ``config.debug.seed`` (not part of SamplingConfig). | ||
|
|
@@ -410,22 +415,32 @@ async def generate( | |
| temperature=_sampling_config.temperature, | ||
| top_p=_sampling_config.top_p, | ||
| max_tokens=_sampling_config.max_tokens, | ||
| n=_sampling_config.n, | ||
| n=1, # group_size pre-expands prompts; the RL loop always samples n=1 | ||
| stop_token_ids=self._stop_token_ids or None, | ||
| seed=self.config.debug.seed, | ||
| logprobs=1, | ||
| output_kind=RequestOutputKind.FINAL_ONLY, | ||
| ) | ||
|
|
||
| if len(request_ids) != len(tokenized_prompts): | ||
| raise ValueError( | ||
| f"got {len(request_ids)} request_ids for {len(tokenized_prompts)} prompts" | ||
| ) | ||
| if len(set(request_ids)) != len(request_ids): | ||
| raise ValueError(f"request_ids must be unique; got {request_ids}") | ||
|
|
||
| # render_cmpl is vLLM's input-pipeline entry. | ||
| # The tokenize step is a no-op for already-tokenized prompts. The | ||
| # lower-level alternative is vllm.inputs.tokens_input; we use the | ||
| # high-level API to stay resilient to vLLM internal changes. | ||
| engine_inputs = self._engine.renderer.render_cmpl( | ||
| [{"prompt_token_ids": ids} for ids in tokenized_prompts] | ||
| ) | ||
| for i, engine_input in enumerate(engine_inputs): | ||
| for engine_input, request_id in zip( | ||
| engine_inputs, request_ids, strict=True | ||
| ): | ||
| self._engine.add_request( | ||
| request_id=str(i), | ||
| request_id=request_id, | ||
| prompt=engine_input, | ||
| params=sampling_params, | ||
| ) | ||
|
|
@@ -435,15 +450,14 @@ async def generate( | |
| while self._engine.has_unfinished_requests(): | ||
| all_outputs.extend(self._engine.step()) | ||
|
|
||
| # vLLM may return requests out of order; sort by the integer | ||
| # request_id we assigned so prompt_idx lines up with the input. | ||
| all_outputs.sort(key=lambda o: int(o.request_id)) | ||
| # Return completions in caller input order; the controller maps positionally. | ||
| request_order = {request_id: i for i, request_id in enumerate(request_ids)} | ||
| all_outputs.sort(key=lambda output: request_order[output.request_id]) | ||
|
|
||
| completions: list[Completion] = [] | ||
| generation_metrics: list[m.Metric] = [] | ||
| output_token_counts: list[int] = [] | ||
| for output in all_outputs: | ||
| prompt_idx = int(output.request_id) | ||
| generation_metrics.extend( | ||
| _prepare_generation_request_metrics(output, prefix=metrics_prefix) | ||
| ) | ||
|
|
@@ -456,8 +470,7 @@ async def generate( | |
| completions.append( | ||
| Completion( | ||
| policy_version=self.policy_version, | ||
| prompt_idx=prompt_idx, | ||
| text=sample.text, | ||
| request_id=output.request_id, | ||
| token_ids=sample.token_ids, | ||
| token_logprobs=per_token_logprobs, | ||
| finish_reason=sample.finish_reason, | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -214,11 +214,11 @@ def _pack_episodes(self, episodes: list[Episode]) -> Iterator[dict]: | |
| def _iterate_samples() -> Iterator[dict]: | ||
| for ep in episodes: | ||
| prompt_len = len(ep.prompt_token_ids) | ||
| response_len = len(ep.token_ids) | ||
| raw_ids = ep.prompt_token_ids + ep.token_ids | ||
| gen_lp = [0.0] * prompt_len + ep.token_logprobs | ||
| loss_mask = [False] * prompt_len + [True] * response_len | ||
| advantages = [0.0] * prompt_len + [ep.advantage] * response_len | ||
| completion_len = len(ep.completion_token_ids) | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Is "completion" more established than "response"? My nit is "completion" doesn't correspond to "prompt" better than "response".
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Both are common in vllm. Completion is a class.
When I hear "completion", i know its talking about the InferenceEngine output. There is no ambiguity. I know its not tool results or env outputs or something else. It is about generated tokens. Response should work well as well. It is what we used in Forge, but there is more room for ambiguity IMO.
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. sounds ok for now |
||
| raw_ids = ep.prompt_token_ids + ep.completion_token_ids | ||
| gen_lp = [0.0] * prompt_len + ep.completion_logprobs | ||
| loss_mask = [False] * prompt_len + [True] * completion_len | ||
| advantages = [0.0] * prompt_len + [ep.advantage] * completion_len | ||
| yield { | ||
| "input_ids": raw_ids[:-1], | ||
| "labels": raw_ids[1:], | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -23,9 +23,10 @@ | |
| from torchtitan.experiments.rl.actors.generator import SamplingConfig, VLLMGenerator | ||
| from torchtitan.experiments.rl.actors.trainer import PolicyTrainer | ||
| from torchtitan.experiments.rl.batcher import BatchConfig, Batcher | ||
| from torchtitan.experiments.rl.examples.sum_digits import SumDigitsRollouter | ||
| 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.models.qwen3 import model_registry | ||
|
|
||
|
|
||
|
|
@@ -36,13 +37,12 @@ def rl_grpo_qwen3_0_6b() -> RLTrainer.Config: | |
| model_spec=model_registry("0.6B", attn_backend="varlen"), | ||
| hf_assets_path="torchtitan/experiments/rl/example_checkpoint/Qwen3-0.6B", | ||
| num_steps=10, | ||
| num_prompts_per_step=5, | ||
| num_groups_per_rollout_batch=5, | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I feel this should be in Rollouter config.
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. wouldnt that break the symmetry with trainer? i.e. the controller is the know that knows batchsize, microbatchsize, max_seq_len to pack to, etc. It sounds to me that the
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. oh sorry this was a comment during early review that I wanted to delete later as I read more. But I forgot |
||
| 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 | ||
| ), | ||
| rollouter=SumDigitsRollouter.Config(), | ||
| group_size=group_size, | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. similar |
||
| renderer=RendererConfig(name="qwen3", enable_thinking=True), | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
| metrics=MetricsProcessor.Config(enable_wandb=True), | ||
| batcher=Batcher.Config( | ||
| batch=BatchConfig(local_batch_size=2, global_batch_size=8, seq_len=2048), | ||
|
|
@@ -77,10 +77,9 @@ def rl_grpo_qwen3_0_6b() -> RLTrainer.Config: | |
| ), | ||
| checkpoint=CheckpointManager.Config(enable=False), | ||
| sampling=SamplingConfig( | ||
| n=group_size, | ||
| temperature=0.8, | ||
| top_p=0.95, | ||
| max_tokens=100, | ||
| max_tokens=700, | ||
| ), | ||
| ), | ||
| ) | ||
|
|
@@ -93,13 +92,12 @@ def rl_grpo_qwen3_1_7b() -> RLTrainer.Config: | |
| model_spec=model_registry("1.7B", attn_backend="varlen"), | ||
| hf_assets_path="torchtitan/experiments/rl/example_checkpoint/Qwen3-1.7B", | ||
| num_steps=10, | ||
| num_prompts_per_step=5, | ||
| num_groups_per_rollout_batch=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 | ||
| ), | ||
| rollouter=SumDigitsRollouter.Config(), | ||
| group_size=group_size, | ||
| 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), | ||
|
|
@@ -135,10 +133,9 @@ def rl_grpo_qwen3_1_7b() -> RLTrainer.Config: | |
| ), | ||
| checkpoint=CheckpointManager.Config(enable=False), | ||
| sampling=SamplingConfig( | ||
| n=group_size, | ||
| temperature=0.8, | ||
| top_p=0.95, | ||
| max_tokens=100, | ||
| max_tokens=700, | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 |
||
| ), | ||
| ), | ||
| ) | ||
|
|
@@ -151,13 +148,12 @@ def rl_grpo_qwen3_14b() -> RLTrainer.Config: | |
| model_spec=model_registry("14B", attn_backend="varlen"), | ||
| hf_assets_path="torchtitan/experiments/rl/example_checkpoint/Qwen3-14B", | ||
| num_steps=10, | ||
| num_prompts_per_step=5, | ||
| num_groups_per_rollout_batch=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 | ||
| ), | ||
| rollouter=SumDigitsRollouter.Config(), | ||
| group_size=group_size, | ||
| 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), | ||
|
|
@@ -192,17 +188,16 @@ def rl_grpo_qwen3_14b() -> RLTrainer.Config: | |
| ), | ||
| checkpoint=CheckpointManager.Config(enable=False), | ||
| sampling=SamplingConfig( | ||
| 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. | ||
| """ | ||
|
|
@@ -212,13 +207,12 @@ def rl_grpo_qwen3_0_6b_batch_invariant() -> RLTrainer.Config: | |
| model_spec=model_registry("0.6B", attn_backend="varlen"), | ||
| hf_assets_path="torchtitan/experiments/rl/example_checkpoint/Qwen3-0.6B", | ||
| num_steps=10, | ||
| num_prompts_per_step=5, | ||
| num_groups_per_rollout_batch=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 | ||
| ), | ||
| rollouter=SumDigitsRollouter.Config(), | ||
| group_size=group_size, | ||
| 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), | ||
|
|
@@ -257,10 +251,9 @@ def rl_grpo_qwen3_0_6b_batch_invariant() -> RLTrainer.Config: | |
| ), | ||
| checkpoint=CheckpointManager.Config(enable=False), | ||
| sampling=SamplingConfig( | ||
| n=group_size, | ||
| temperature=0.8, | ||
| top_p=0.95, | ||
| max_tokens=100, | ||
| max_tokens=700, | ||
| ), | ||
| debug=batch_invariant_config, | ||
| ), | ||
|
|
||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. In general I would prefer full spelling over shorthand. Could we do
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I wanted to signal somehow that those are "base" envs, not something like a SumDigitsEnvs, thats why i put
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I see. If we follow strictly the sft counterpart, for classes / interfaces you would expect users to implement, it probably should be
Meanwhile, |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,20 @@ | ||
| # 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.environment.message import ( | ||
| MessageEnv, | ||
| MessageInitOutput, | ||
| MessageStepOutput, | ||
| ) | ||
| from torchtitan.experiments.rl.environment.token import TokenEnv, TokenEnvOutput | ||
|
|
||
| __all__ = [ | ||
| "MessageEnv", | ||
| "MessageInitOutput", | ||
| "MessageStepOutput", | ||
| "TokenEnv", | ||
| "TokenEnvOutput", | ||
| ] |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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?