diff --git a/examples/tinker/README.md b/examples/tinker/README.md index 19cb0eb9bd3..3e49ac988c9 100644 --- a/examples/tinker/README.md +++ b/examples/tinker/README.md @@ -61,6 +61,48 @@ trinity run --config tinker.yaml # Replace with your actual config file path > 💡 A complete example configuration file is available at [`tinker.yaml`](tinker.yaml). +## Optional server-side PPO loss + +A Tinker-compatible server implementing `trinity_ppo` can compute the PPO loss +and backward pass together, avoiding the client-side custom-loss forward/backward +round trip. This is a server extension, not a built-in loss on every Tinker service. + +```yaml +model: + tinker: + enable: true + rank: 32 + server_loss_fn: trinity_ppo +algorithm: + algorithm_type: grpo + policy_loss_fn: ppo + policy_loss_fn_args: + clip_range: 0.2 + clip_ratio_c: 3.0 + loss_agg_mode: token-mean + kl_loss_fn: k2 + kl_loss_fn_args: + kl_coef: 0.001 + entropy_loss_fn: none + loss_agg_mode: token-mean +``` + +The supported objective is dual-clipped PPO with symmetric clipping and optional +K2 KL loss. Set `kl_loss_fn: none` to disable KL and its reference-logprob requests. +Sequence masking, fallback policy gradient, adaptive KL, entropy loss, alternative +reductions, and `fix_actor_microbatch_loss_scale` are rejected for this path. +Omit `server_loss_fn` to keep the existing client-side callback. + +Each datum contributes its mean loss over active response tokens; the sum is +divided by the full training batch's datum count. The action mask excludes both +prompt tokens and masked response tokens. The SDK can split the batch into +requests, but every request receives the same full-batch denominator and the +trainer makes one optimizer step after the batch. Empty masks contribute zero. + +The server reports additive statistics. Trinity derives token-weighted ratio, +clipping, and KL diagnostics from their sums and active-token counts, so uneven +request sizes do not become an unweighted mean of request means. + ## Results on the Llama-3.2-3B Model We trained the **Llama-3.2-3B** model on the **GSM8K** dataset using both the **Tinker** and **veRL** backends. Below are the full configuration files used in our experiments. diff --git a/tests/trainer/tinker_server_loss_test.py b/tests/trainer/tinker_server_loss_test.py new file mode 100644 index 00000000000..5494c3b49f7 --- /dev/null +++ b/tests/trainer/tinker_server_loss_test.py @@ -0,0 +1,257 @@ +"""CPU tests for the optional Tinker server-side PPO loss.""" + +import asyncio +import unittest +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch + +import torch + +from trinity.algorithm.entropy_loss_fn.entropy_loss_fn import DummyEntropyLossFn +from trinity.algorithm.kl_fn.kl_fn import DummyKLFn, K1Fn, K2Fn +from trinity.algorithm.policy_loss_fn.ppo_policy_loss import PPOPolicyLossFn +from trinity.common.config import Config +from trinity.common.experience import Experience +from trinity.trainer.tinker.server_loss import ( + trinity_ppo_metrics, + with_trinity_ppo_inputs, +) +from trinity.trainer.tinker.tinker_trainer import TinkerTrainerWrapper +from trinity.trainer.tinker.utils import to_tinker_input + + +def make_wrapper(): + """Construct a trainer without Ray actors or remote clients.""" + wrapper = TinkerTrainerWrapper.__new__(TinkerTrainerWrapper) + wrapper.config = Config() + wrapper.config.model.tinker.server_loss_fn = "trinity_ppo" + wrapper.algorithm = SimpleNamespace(use_reference=True, compute_advantage_in_trainer=False) + wrapper.policy_loss_fn = PPOPolicyLossFn(backend="tinker", clip_range=0.2) + wrapper.kl_loss_fn = K2Fn(kl_coef=0.001) + wrapper.entropy_loss_fn = DummyEntropyLossFn(entropy_coef=0) + wrapper.loss_agg_mode = "token-mean" + wrapper.do_fix_actor_microbatch_loss_scale = False + return wrapper + + +def make_experience(mask=(1, 0, 1)): + return Experience( + tokens=torch.tensor([10, 11, 12, 13, 14]), + prompt_length=2, + action_mask=torch.tensor(mask, dtype=torch.bool), + logprobs=torch.tensor([-0.2, -0.3, -0.4]), + advantages=torch.tensor([2.0, 3.0, -1.0]), + reward=1.0, + ) + + +class ServerLossInputTest(unittest.TestCase): + def setUp(self): + self.batch, _, self.inputs = to_tinker_input([make_experience()], MagicMock()) + self.inputs[0]["ref_logprob"] = torch.tensor([-0.1, -0.5, -0.2]) + + def test_alignment_preserves_masked_response_tokens(self): + result = with_trinity_ppo_inputs(self.batch, self.inputs, kl_coef=0.001)[0] + expected = { + "target_tokens": [11, 12, 13, 14], + "weights": [0, 1, 0, 1], + "logprobs": [0, -0.2, 0, -0.4], + "advantages": [0, 2, 0, -1], + "ref_logprobs": [0, -0.1, 0, -0.2], + } + for key, values in expected.items(): + torch.testing.assert_close( + result.loss_fn_inputs[key].to_torch().float(), torch.tensor(values).float() + ) + self.assertNotIn("logprobs", self.batch[0].loss_fn_inputs) + + def test_scalar_advantage_and_zero_kl_without_reference(self): + self.inputs[0]["advantages"] = torch.tensor(2.0) + del self.inputs[0]["ref_logprob"] + result = with_trinity_ppo_inputs(self.batch, self.inputs, kl_coef=0)[0] + self.assertNotIn("ref_logprobs", result.loss_fn_inputs) + self.assertEqual(result.loss_fn_inputs["advantages"].data, [0, 2, 0, 2]) + + def test_empty_mask_contributes_zero_and_ignores_masked_nonfinite_values(self): + self.inputs[0]["action_mask"] = torch.zeros(3) + self.inputs[0]["old_logprob"] = torch.tensor([float("-inf")] * 3) + result = with_trinity_ppo_inputs(self.batch, self.inputs, kl_coef=0.001)[0] + for key in ("weights", "logprobs", "advantages", "ref_logprobs"): + self.assertEqual(result.loss_fn_inputs[key].data, [0, 0, 0, 0]) + + def test_missing_required_values_and_invalid_shapes_are_rejected(self): + for key in ("old_logprob", "advantages", "ref_logprob"): + with self.subTest(key=key): + inputs = [dict(self.inputs[0])] + del inputs[0][key] + with self.assertRaisesRegex(ValueError, key): + with_trinity_ppo_inputs(self.batch, inputs, kl_coef=0.001) + self.inputs[0]["old_logprob"] = torch.tensor([1.0, 2.0]) + with self.assertRaisesRegex(ValueError, "per response token"): + with_trinity_ppo_inputs(self.batch, self.inputs, kl_coef=0) + + def test_nonbinary_masks_and_active_nonfinite_values_are_rejected(self): + self.inputs[0]["action_mask"] = torch.tensor([1, 0.5, 1]) + with self.assertRaisesRegex(ValueError, "binary"): + with_trinity_ppo_inputs(self.batch, self.inputs, kl_coef=0) + self.inputs[0]["action_mask"] = torch.ones(3) + self.inputs[0]["advantages"] = torch.tensor([1.0, float("nan"), 2.0]) + with self.assertRaisesRegex(ValueError, "finite"): + with_trinity_ppo_inputs(self.batch, self.inputs, kl_coef=0) + + def test_metrics_use_additive_token_counts(self): + metrics = { + "loss:sum": 0.3, + "trinity/response_tokens:sum": 4, + "trinity/ratio_sum:sum": 6, + "trinity/ratio_squared_sum:sum": 10, + "trinity/clipped_tokens:sum": 1, + "trinity/kl_sum:sum": 0.8, + } + self.assertEqual( + trinity_ppo_metrics(metrics), + { + "actor/final_loss": 0.3, + "actor/ratio_mean": 1.5, + "actor/ratio_var": 0.25, + "actor/ratio_clip_fraction": 0.25, + "actor/kl_token_mean": 0.2, + }, + ) + self.assertEqual(trinity_ppo_metrics({"trinity/response_tokens:sum": 0}), {}) + + +class ServerLossConfigTest(unittest.TestCase): + def test_supported_objective_and_disabled_kl(self): + wrapper = make_wrapper() + self.assertEqual( + wrapper._server_loss_fn_config(), + {"clip_range": 0.2, "clip_ratio_c": 3.0, "kl_coef": 0.001}, + ) + wrapper.kl_loss_fn = DummyKLFn() + self.assertEqual(wrapper._server_loss_fn_config()["kl_coef"], 0) + + def test_unsupported_policy_options_fail_before_training(self): + for attribute, value in ( + ("clip_range_high", 0.3), + ("enable_sequence_masking", True), + ("fallback_to_policy_gradient", True), + ("loss_agg_mode", "seq-mean-token-sum"), + ): + with self.subTest(attribute=attribute): + wrapper = make_wrapper() + setattr(wrapper.policy_loss_fn, attribute, value) + with self.assertRaisesRegex(ValueError, "does not support"): + wrapper._server_loss_fn_config() + + def test_missing_nonfinite_and_out_of_range_clipping_is_rejected(self): + for clip_range in (None, float("nan"), float("inf"), float("-inf"), -0.1, 1.0): + with self.subTest(clip_range=clip_range): + wrapper = make_wrapper() + wrapper.policy_loss_fn.clip_range_low = clip_range + with self.assertRaisesRegex(ValueError, "clip_range in"): + wrapper._server_loss_fn_config() + + def test_zero_clipping_is_supported(self): + wrapper = make_wrapper() + wrapper.policy_loss_fn = PPOPolicyLossFn(backend="tinker", clip_range=0.0) + self.assertEqual(wrapper._server_loss_fn_config()["clip_range"], 0.0) + + def test_unsupported_loss_options_fail_before_training(self): + for attribute, value in ( + ("kl_loss_fn", K1Fn()), + ("entropy_loss_fn", object()), + ("policy_loss_fn", object()), + ("loss_agg_mode", "seq-mean-token-sum"), + ("do_fix_actor_microbatch_loss_scale", True), + ): + with self.subTest(attribute=attribute): + wrapper = make_wrapper() + setattr(wrapper, attribute, value) + with self.assertRaises(ValueError): + wrapper._server_loss_fn_config() + wrapper = make_wrapper() + wrapper.algorithm.use_reference = False + with self.assertRaisesRegex(ValueError, "reference"): + wrapper._server_loss_fn_config() + + +class ServerLossTrainStepTest(unittest.IsolatedAsyncioTestCase): + def setUp(self): + wrapper = make_wrapper() + wrapper.logger = MagicMock() + wrapper.server_loss_config = wrapper._server_loss_fn_config() + wrapper._train_step_num = 0 + wrapper.algorithm_config = wrapper.config.algorithm + wrapper.lr_scheduler_type = "constant" + wrapper.num_warmup_steps = 0 + wrapper.total_steps = 10 + wrapper.min_lr_ratio = 0 + wrapper.ref_client = SimpleNamespace( + compute_logprobs_async=AsyncMock(return_value=[None, -0.1, -0.2, -0.3, -0.4]) + ) + calls = [] + + async def forward_backward(batch, loss_fn, loss_config): + calls.append(("forward_backward", len(batch), loss_fn, loss_config)) + result = asyncio_future(SimpleNamespace(metrics={"loss:sum": 0.3})) + return result + + async def optim_step(params): + calls.append(("optim_step",)) + return asyncio_future(SimpleNamespace(metrics={})) + + wrapper.actor_client = SimpleNamespace( + forward_backward_async=AsyncMock(side_effect=forward_backward), + forward_backward_custom_async=AsyncMock(), + optim_step_async=AsyncMock(side_effect=optim_step), + ) + self.wrapper = wrapper + self.calls = calls + + async def train(self): + with patch( + "trinity.trainer.tinker.tinker_trainer.compute_throughout_metrics", return_value={} + ): + return await self.wrapper.train_step([make_experience(), make_experience()]) + + async def test_one_optimizer_step_and_full_batch_denominator(self): + metrics = await self.train() + wrapper, calls = self.wrapper, self.calls + self.assertEqual(calls[0][1:3], (2, "trinity_ppo")) + self.assertEqual(calls[0][3]["num_total_datums"], 2) + self.assertEqual(calls[1], ("optim_step",)) + self.assertEqual(len(calls), 2) + wrapper.actor_client.forward_backward_custom_async.assert_not_called() + self.assertEqual(metrics["actor/final_loss"], 0.3) + + async def test_zero_kl_skips_reference_requests(self): + self.wrapper.server_loss_config["kl_coef"] = 0 + await self.train() + self.wrapper.ref_client.compute_logprobs_async.assert_not_called() + + async def test_default_path_keeps_client_loss_callback(self): + self.wrapper.server_loss_config = None + future = asyncio_future(SimpleNamespace(metrics={"custom_loss": 0.5})) + self.wrapper.actor_client.forward_backward_custom_async.return_value = future + metrics = await self.train() + self.wrapper.actor_client.forward_backward_async.assert_not_called() + self.wrapper.actor_client.forward_backward_custom_async.assert_awaited_once() + self.assertEqual(metrics["custom_loss"], 0.5) + + async def test_failed_server_request_does_not_step_optimizer(self): + # A failed asynchronous result models server-side validation failures. + future = asyncio.get_running_loop().create_future() + future.set_exception(ValueError("unsupported loss")) + self.wrapper.actor_client.forward_backward_async.side_effect = None + self.wrapper.actor_client.forward_backward_async.return_value = future + with self.assertRaisesRegex(ValueError, "unsupported loss"): + await self.train() + self.wrapper.actor_client.optim_step_async.assert_not_called() + + +def asyncio_future(result): + """Return an awaitable result, like a completed SDK API future.""" + future = asyncio.get_running_loop().create_future() + future.set_result(result) + return future diff --git a/trinity/common/config.py b/trinity/common/config.py index 710d9d242d2..291e37c7de5 100644 --- a/trinity/common/config.py +++ b/trinity/common/config.py @@ -446,6 +446,8 @@ class TinkerConfig: train_attn: bool = True train_unembed: bool = True base_url: Optional[str] = None + # Optional loss implemented by the server; None uses the client-side callback. + server_loss_fn: Optional[str] = None @dataclass diff --git a/trinity/trainer/tinker/server_loss.py b/trinity/trainer/tinker/server_loss.py new file mode 100644 index 00000000000..0dd049a8ad2 --- /dev/null +++ b/trinity/trainer/tinker/server_loss.py @@ -0,0 +1,86 @@ +"""Build inputs and metrics for the optional server-side Trinity PPO loss.""" + +from typing import Dict, List + +import torch +from tinker import types + + +def with_trinity_ppo_inputs( + batch: List[types.Datum], model_inputs: List[dict], kl_coef: float +) -> List[types.Datum]: + """Align response-only tensors with the shifted target tokens. + + Args: + batch: Datums produced by `to_tinker_input`. + model_inputs: Corresponding response masks, old logprobs and advantages. + kl_coef: Nonzero when reference logprobs must be included. + + Returns: + New datums with PPO inputs; the input datums are not mutated. + + Raises: + ValueError: Required response data is absent or inconsistent. + """ + if len(batch) != len(model_inputs): + raise ValueError("Server loss needs one set of model inputs per datum.") + output = [] + for datum, inputs in zip(batch, model_inputs): + mask = inputs["action_mask"].to(dtype=torch.float32, device="cpu") + target_length = len(datum.loss_fn_inputs["target_tokens"].data) + response_length = mask.numel() + if mask.ndim != 1 or not 0 < response_length <= target_length: + raise ValueError("Server loss requires a nonempty response mask matching the targets.") + if not torch.all((mask == 0) | (mask == 1)): + raise ValueError("trinity_ppo requires a binary action mask.") + + def padded(key: str, allow_scalar: bool = False) -> torch.Tensor: + if key not in inputs: + raise ValueError(f"trinity_ppo requires {key}.") + values = torch.as_tensor(inputs[key], dtype=torch.float32, device="cpu").reshape(-1) + if allow_scalar and values.numel() == 1: + values = values.expand(response_length) + if values.numel() != response_length: + raise ValueError(f"{key} must have one value per response token.") + values = values.masked_fill(~mask.bool(), 0) + if not torch.isfinite(values).all(): + raise ValueError(f"{key} must be finite on active response tokens.") + result = torch.zeros(target_length, dtype=torch.float32) + result[-response_length:] = values + return result + + loss_inputs = dict(datum.loss_fn_inputs) + # Use the actual action mask, including masked response tokens, rather + # than inferring it from sequence lengths or nonzero advantages. + weights = torch.zeros(target_length, dtype=torch.float32) + weights[-response_length:] = mask + loss_inputs["weights"] = weights + loss_inputs["logprobs"] = padded("old_logprob") + loss_inputs["advantages"] = padded("advantages", allow_scalar=True) + if kl_coef > 0: + loss_inputs["ref_logprobs"] = padded("ref_logprob") + output.append(types.Datum(model_input=datum.model_input, loss_fn_inputs=loss_inputs)) + return output + + +def trinity_ppo_metrics(metrics: Dict[str, float]) -> Dict[str, float]: + """Derive token-weighted diagnostics from additive server statistics.""" + output = {} + if "loss:sum" in metrics: + output["actor/final_loss"] = metrics["loss:sum"] + count = metrics.get("trinity/response_tokens:sum", 0) + if count <= 0: + return output + for source, destination in ( + ("trinity/ratio_sum:sum", "actor/ratio_mean"), + ("trinity/clipped_tokens:sum", "actor/ratio_clip_fraction"), + ("trinity/kl_sum:sum", "actor/kl_token_mean"), + ): + if source in metrics: + output[destination] = metrics[source] / count + if "trinity/ratio_squared_sum:sum" in metrics and "actor/ratio_mean" in output: + output["actor/ratio_var"] = max( + 0.0, + metrics["trinity/ratio_squared_sum:sum"] / count - output["actor/ratio_mean"] ** 2, + ) + return output diff --git a/trinity/trainer/tinker/tinker_trainer.py b/trinity/trainer/tinker/tinker_trainer.py index 7f6769fb80e..53add6ae7a5 100644 --- a/trinity/trainer/tinker/tinker_trainer.py +++ b/trinity/trainer/tinker/tinker_trainer.py @@ -13,11 +13,17 @@ from trinity.algorithm.entropy_loss_fn import ENTROPY_LOSS_FN from trinity.algorithm.entropy_loss_fn.entropy_loss_fn import DummyEntropyLossFn from trinity.algorithm.kl_fn import KL_FN +from trinity.algorithm.kl_fn.kl_fn import DummyKLFn, K2Fn from trinity.algorithm.policy_loss_fn import POLICY_LOSS_FN +from trinity.algorithm.policy_loss_fn.ppo_policy_loss import PPOPolicyLossFn from trinity.algorithm.utils import prefix_metrics from trinity.common.config import Config from trinity.common.experience import Experience from trinity.manager.synchronizer import Synchronizer +from trinity.trainer.tinker.server_loss import ( + trinity_ppo_metrics, + with_trinity_ppo_inputs, +) from trinity.trainer.tinker.utils import ( compute_data_metrics, compute_throughout_metrics, @@ -64,6 +70,9 @@ def _init_algorithm(self): self.config.trainer.fix_actor_microbatch_loss_scale and (self.loss_agg_mode == "token-mean") ) + self.server_loss_config = ( + self._server_loss_fn_config() if self.config.model.tinker.server_loss_fn else None + ) self.lr_scheduler_type = algorithm_config.optimizer.lr_scheduler_type self.total_steps = self.config.trainer.total_steps or sys.maxsize @@ -108,6 +117,43 @@ def _current_lr_factor(self): factor = self.min_lr_ratio + (1.0 - self.min_lr_ratio) * factor return max(self.min_lr_ratio, factor) + def _server_loss_fn_config(self) -> dict: + """Validate that the server loss implements the configured objective.""" + if self.config.model.tinker.server_loss_fn != "trinity_ppo": + raise ValueError("The only supported server_loss_fn is 'trinity_ppo'.") + policy = self.policy_loss_fn + if type(policy) is not PPOPolicyLossFn: + raise ValueError("trinity_ppo requires the PPO policy loss.") + if type(self.kl_loss_fn) not in (DummyKLFn, K2Fn): + raise ValueError("trinity_ppo supports only K2 or disabled KL loss.") + if type(self.entropy_loss_fn) is not DummyEntropyLossFn: + raise ValueError("trinity_ppo requires entropy_loss_fn='none'.") + clip_range = policy.clip_range_low + if clip_range is None or not math.isfinite(clip_range) or not 0 <= clip_range < 1: + raise ValueError("trinity_ppo requires clip_range in [0, 1).") + unsupported = { + "asymmetric clipping": policy.clip_range_low != policy.clip_range_high, + "sequence masking": policy.enable_sequence_masking, + "fallback policy gradient": policy.fallback_to_policy_gradient, + "policy reduction": policy.loss_agg_mode != "token-mean", + "loss reduction": self.loss_agg_mode != "token-mean", + "microbatch loss rescaling": self.do_fix_actor_microbatch_loss_scale, + "adaptive KL": self.kl_loss_fn.adaptive, + } + if any(unsupported.values()): + names = ", ".join(name for name, enabled in unsupported.items() if enabled) + raise ValueError(f"trinity_ppo does not support: {names}.") + kl_coef = 0.0 if isinstance(self.kl_loss_fn, DummyKLFn) else self.kl_loss_fn.kl_coef + if kl_coef < 0 or not math.isfinite(kl_coef): + raise ValueError("trinity_ppo requires a finite nonnegative KL coefficient.") + if kl_coef > 0 and not self.algorithm.use_reference: + raise ValueError("trinity_ppo with KL loss requires a reference policy.") + return { + "clip_range": clip_range, + "clip_ratio_c": policy.clip_ratio_c, + "kl_coef": kl_coef, + } + @property def current_learning_rate(self): return self._current_lr_factor * self.algorithm_config.optimizer.lr @@ -273,7 +319,9 @@ async def train_step(self, batch_exps: List[Experience]) -> Dict: self._train_step_num += 1 with Timer(timing_raw, "step"): - if self.algorithm.use_reference: # ref_logprob may not be used + if self.algorithm.use_reference and ( + self.server_loss_config is None or self.server_loss_config["kl_coef"] > 0 + ): import asyncio ref_logprobs = await asyncio.gather( @@ -300,13 +348,30 @@ async def train_step(self, batch_exps: List[Experience]) -> Dict: # update actor with Timer(timing_raw, "update_actor"): - fwdbwd_future = await self.actor_client.forward_backward_custom_async( - batch, self._loss_func - ) + if self.server_loss_config is None: + fwdbwd_future = await self.actor_client.forward_backward_custom_async( + batch, self._loss_func + ) + else: + server_batch = with_trinity_ppo_inputs( + batch, model_inputs_list, self.server_loss_config["kl_coef"] + ) + # The SDK may split this call into requests. Every request + # must use the full batch denominator before one optim step. + loss_config = dict(self.server_loss_config, num_total_datums=len(batch)) + fwdbwd_future = await self.actor_client.forward_backward_async( + server_batch, "trinity_ppo", loss_config + ) + # Do not apply a partial optimizer update if a server rejects + # the extension or one of its accumulated requests fails. + fwdbwd_result = await fwdbwd_future optim_future = await self.actor_client.optim_step_async(self.adam_params) - fwdbwd_result = await fwdbwd_future + if self.server_loss_config is None: + fwdbwd_result = await fwdbwd_future optim_result = await optim_future metrics.update(fwdbwd_result.metrics) + if self.server_loss_config is not None: + metrics.update(trinity_ppo_metrics(fwdbwd_result.metrics)) if optim_result.metrics: metrics.update(optim_result.metrics)