From 78f4c7d5ba9b9467b4a6497b6072faffb20fd443 Mon Sep 17 00:00:00 2001 From: "Gavin.Zhu" Date: Mon, 13 Jul 2026 07:32:36 +0000 Subject: [PATCH 1/6] tinker: decoupled train-step seam on upstream (fwd/bwd-only + optimizer step) Tinker's API pipelines N forward_backward calls then one optim_step; upstream couples fwd+bwd+step inside train_one_step. This adds the minimal seam (specs/005 design in tinker-nemorl), reusing upstream data/loss/fanout: - model.py: extract _build_train_forward_step + _configure_train_step from train_one_step/train (pure moves); add forward_backward_pass (no step, grads accumulate across calls) and optimizer_step (per-call LR override; when the client sets LR the Megatron scheduler is not stepped). - loss: per-batch _loss_type_override (per-request loss selection) and _loss_norm_total (=1 -> pure-sum gradients, invariant to how a batch is split across forward_backward calls; fixes the 2x grad inflation measured in the G1 spike). Same rollout-key pattern as dynamic_global_batch_size. - data: forward the two scalar keys in get_batch; keep _partition_indices in process_rollout_data so per-sample outputs reassemble into client order. - actor: forward_backward_only / apply_optimizer_step / forward_logprobs / load_checkpoint (full resume incl. optimizer+scheduler). - ray/tinker_group.py (new): TinkerTrainGroup(RayTrainGroup) async fanout + DP-order merge. Frozen v1 RayTrainGroup untouched. Co-Authored-By: Claude Fable 5 --- miles/backends/megatron_utils/actor.py | 93 ++++++- miles/backends/megatron_utils/model.py | 247 +++++++++++++----- miles/backends/training_utils/data.py | 6 + miles/backends/training_utils/loss.py | 6 +- .../training_utils/loss_hub/losses.py | 9 +- miles/ray/tinker_group.py | 58 ++++ miles/utils/data.py | 4 + 7 files changed, 353 insertions(+), 70 deletions(-) create mode 100644 miles/ray/tinker_group.py diff --git a/miles/backends/megatron_utils/actor.py b/miles/backends/megatron_utils/actor.py index b922417598e..b4e5818147d 100644 --- a/miles/backends/megatron_utils/actor.py +++ b/miles/backends/megatron_utils/actor.py @@ -50,7 +50,15 @@ from .ft.indep_dp import reconfigure_indep_dp_group from .initialize import init, is_first_replica_megatron_main_rank from .lora_utils import is_lora_enabled -from .model import TrainStepOutcome, forward_only, initialize_model_and_optimizer, save, train +from .model import ( + TrainStepOutcome, + forward_backward_pass, + forward_only, + initialize_model_and_optimizer, + optimizer_step, + save, + train, +) from .parallel import verify_megatron_parallel_state from .replay_utils import register_replay_list_moe from .update_weight.common import named_params_and_buffers @@ -661,6 +669,89 @@ def load_other_checkpoint(self, model_tag: str, path: str) -> None: self.weights_backuper.backup(model_tag) self._active_model_tag = model_tag + # --- Tinker seam ------------------------------------------------------- + # Decoupled train-step primitives for the Tinker API (specs/005 in + # tinker-nemorl): N x forward_backward_only accumulate gradients, one + # apply_optimizer_step applies them with a per-call LR. + + @with_logs + def forward_backward_only(self, rollout_id: int, rollout_data_ref: Box) -> dict: + """Forward + backward WITHOUT the optimizer step; grads accumulate.""" + self._heartbeat.bump() + if self.args.offload_train: + self.wake_up() + + with timer("data_preprocess"): + rollout_data = get_rollout_data(self.args, rollout_data_ref, witness_info=None) + + data_iterator, num_microbatches = get_data_iterator(self.args, self.model, rollout_data) + + with timer("forward_backward_only"): + loss_dict = forward_backward_pass( + rollout_id, data_iterator, self.model, self.optimizer, num_microbatches + ) + + return { + "loss": loss_dict, + "partition_indices": rollout_data.get("_partition_indices", []), + } + + @with_logs + def apply_optimizer_step(self, learning_rate: float | None = None) -> dict: + """Apply the optimizer over accumulated grads; optional per-call LR.""" + with timer("apply_optimizer_step"): + result = optimizer_step( + self.model, self.optimizer, self.opt_param_scheduler, learning_rate=learning_rate + ) + + if result["success"] and self._enable_weight_backup: + self.weights_backuper.backup("actor") + self._heartbeat.bump() + return result + + @with_logs + def forward_logprobs(self, rollout_id: int, rollout_data_ref: Box) -> dict: + """Forward-only per-sample log-probs (no grad); client order restored + via partition_indices by the fanout layer.""" + self._heartbeat.bump() + if self.args.offload_train: + self.wake_up() + + with timer("data_preprocess"): + rollout_data = get_rollout_data(self.args, rollout_data_ref, witness_info=None) + + data_iterator, num_microbatches = get_data_iterator(self.args, self.model, rollout_data) + result = self.compute_log_prob(data_iterator, num_microbatches, rollout_id=rollout_id) + + log_probs = result.get("log_probs") or [] + return { + "log_probs": [t.cpu() for t in log_probs], + "partition_indices": rollout_data.get("_partition_indices", []), + } + + @with_logs + def load_checkpoint(self, checkpoint_path: str) -> dict: + """Full resume (model + optimizer + scheduler) from checkpoint_path.""" + old_args = self.args.load, self.args.no_load_optim, self.args.no_load_rng, self.args.finetune + self.args.load = checkpoint_path + self.args.no_load_optim = False + self.args.no_load_rng = False + self.args.finetune = False + try: + iteration, _ = load_checkpoint( + self.model, + self.optimizer, + self.opt_param_scheduler, + checkpointing_context={}, + skip_load_to_model_and_opt=False, + ) + finally: + self.args.load, self.args.no_load_optim, self.args.no_load_rng, self.args.finetune = old_args + + if self._enable_weight_backup: + self.weights_backuper.backup("actor") + return {"success": True, "iteration": iteration} + @with_logs def connect_actor_critic( self, diff --git a/miles/backends/megatron_utils/model.py b/miles/backends/megatron_utils/model.py index 9910df915fc..635b3f2fe76 100644 --- a/miles/backends/megatron_utils/model.py +++ b/miles/backends/megatron_utils/model.py @@ -355,53 +355,11 @@ def forward_step( return rollout_data -def train_one_step( - args: Namespace, - rollout_id: int, - step_id: int, - data_iterator: Sequence[DataIterator], - model: Sequence[DDP], - optimizer: MegatronOptimizer | None, - opt_param_scheduler: OptimizerParamScheduler | None, - num_microbatches: int, - witness_info: WitnessInfo | None, - attempt: int, - ft_test_action_executor: FTTestActionActorExecutor | None = None, -) -> tuple[dict[str, float], float, TrainStepOutcome]: - """Execute a single pipeline-parallel training step. - - Runs forward/backward over ``num_microbatches``, applies optimizer step and - one scheduler step when gradients are valid. +def _build_train_forward_step(args: Namespace, num_microbatches: int, dumper_phase_util: DumperMegatronUtil): + """Build the forward_step closure used by Megatron's pipeline engine. - Args: - args: Runtime arguments. - rollout_id: Rollout identifier. - step_id: Step index within the current rollout. - data_iterator: Iterable(s) yielding training batches. - model: Sequence of DDP-wrapped model chunks. - optimizer: Optimizer instance. - opt_param_scheduler: LR/WD scheduler. - num_microbatches: Number of microbatches to process. - - Returns: - Tuple of (reduced loss dict, gradient norm, step outcome). + Shared by ``train_one_step`` and the Tinker seam's ``forward_backward_pass``. """ - args = get_args() - parallel_state = get_parallel_state() - dumper_phase_util = DumperMegatronUtil(args, model, DumperPhase.FWD_BWD, rollout_id=rollout_id) - disable_optimizer = args.debug_disable_optimizer or optimizer is None - - # Set grad to zero. - for model_chunk in model: - model_chunk.zero_grad_buffer() - if not disable_optimizer: - optimizer.zero_grad() - - if args.custom_megatron_before_train_step_hook_path: - from miles.utils.misc import load_function - - custom_before_train_step_hook = load_function(args.custom_megatron_before_train_step_hook_path) - custom_before_train_step_hook(args, rollout_id, step_id, model, optimizer, opt_param_scheduler) @dumper_phase_util.wrap_forward_step def forward_step(data_iterator: DataIterator, model: GPTModel, return_schedule_plan: bool = False) -> tuple[ @@ -486,6 +444,59 @@ def forward_step(data_iterator: DataIterator, model: GPTModel, return_schedule_p return output_tensor, partial(loss_function, args, batch, num_microbatches, apply_megatron_loss_scaling=True) + return forward_step + + +def train_one_step( + args: Namespace, + rollout_id: int, + step_id: int, + data_iterator: Sequence[DataIterator], + model: Sequence[DDP], + optimizer: MegatronOptimizer | None, + opt_param_scheduler: OptimizerParamScheduler | None, + num_microbatches: int, + witness_info: WitnessInfo | None, + attempt: int, + ft_test_action_executor: FTTestActionActorExecutor | None = None, +) -> tuple[dict[str, float], float, TrainStepOutcome]: + """Execute a single pipeline-parallel training step. + + Runs forward/backward over ``num_microbatches``, applies optimizer step and + one scheduler step when gradients are valid. + + Args: + args: Runtime arguments. + rollout_id: Rollout identifier. + step_id: Step index within the current rollout. + data_iterator: Iterable(s) yielding training batches. + model: Sequence of DDP-wrapped model chunks. + optimizer: Optimizer instance. + opt_param_scheduler: LR/WD scheduler. + num_microbatches: Number of microbatches to process. + + Returns: + Tuple of (reduced loss dict, gradient norm, step outcome). + """ + args = get_args() + parallel_state = get_parallel_state() + dumper_phase_util = DumperMegatronUtil(args, model, DumperPhase.FWD_BWD, rollout_id=rollout_id) + disable_optimizer = args.debug_disable_optimizer or optimizer is None + + # Set grad to zero. + for model_chunk in model: + model_chunk.zero_grad_buffer() + if not disable_optimizer: + optimizer.zero_grad() + + if args.custom_megatron_before_train_step_hook_path: + from miles.utils.misc import load_function + + custom_before_train_step_hook = load_function(args.custom_megatron_before_train_step_hook_path) + custom_before_train_step_hook(args, rollout_id, step_id, model, optimizer, opt_param_scheduler) + + forward_step = _build_train_forward_step(args, num_microbatches, dumper_phase_util) + # Forward pass. forward_backward_func = get_forward_backward_func() losses_reduced = forward_backward_func( @@ -588,6 +599,131 @@ def finalize_model_grads_with_empty_cache(*args, **kwargs): return finalize_model_grads(*args, **kwargs) +def _configure_train_step( + model: Sequence[DDP], + optimizer: MegatronOptimizer | None, + args: Namespace, + disable_optimizer: bool, +): + """Configure the model config for a training fwd/bwd pass. + + Extracted from ``train`` so the Tinker seam's ``forward_backward_pass`` can + reuse it. Safe to call repeatedly: the no_sync guard only asserts on the + first configuration (the seam reconfigures on every forward_backward call). + """ + config = get_model_config(model[0]) + config.grad_scale_func = None if disable_optimizer else optimizer.scale_loss + config.timers = None + if isinstance(model[0], DDP) and args.overlap_grad_reduce: + if not getattr(config, "_miles_no_sync_configured", False): + assert config.no_sync_func is None, ( + "When overlap_grad_reduce is True, config.no_sync_func must be None; " + "a custom no_sync_func is not supported when overlapping grad-reduce" + ) + config._miles_no_sync_configured = True + config.no_sync_func = [model_chunk.no_sync for model_chunk in model] + if len(model) == 1: + config.no_sync_func = config.no_sync_func[0] + if args.align_grad_reduce: + config.grad_sync_func = [model_chunk.start_grad_sync for model_chunk in model] + if len(model) == 1: + config.grad_sync_func = config.grad_sync_func[0] + if args.overlap_param_gather and args.align_param_gather: + config.param_sync_func = [model_chunk.start_param_sync for model_chunk in model] + if len(model) == 1: + config.param_sync_func = config.param_sync_func[0] + config.finalize_model_grads_func = finalize_model_grads_with_empty_cache + return config + + +def forward_backward_pass( + rollout_id: int, + data_iterator: Sequence[DataIterator], + model: Sequence[DDP], + optimizer: MegatronOptimizer | None, + num_microbatches: Sequence[int], + zero_grads: bool = False, +) -> dict[str, float]: + """Tinker seam: forward + backward WITHOUT the optimizer step. + + Gradients accumulate in the DDP grad buffers across successive calls until + ``optimizer_step`` runs (which zeroes them). Split-invariance of the + accumulated gradient across calls relies on the loss normalization being + call-independent — the Tinker path pins it via the ``_loss_norm_total`` + rollout key (see training_utils/loss.py); see specs/005 design.md (G1). + """ + args = get_args() + disable_optimizer = args.debug_disable_optimizer or optimizer is None + + for iterator in data_iterator: + iterator.reset() + for model_module in model: + model_module.train() + + _configure_train_step(model, optimizer, args, disable_optimizer) + + if zero_grads: + for model_chunk in model: + model_chunk.zero_grad_buffer() + if not disable_optimizer: + optimizer.zero_grad() + + forward_backward_func = get_forward_backward_func() + losses_reduced = [] + for num_mbs in num_microbatches: + dumper_phase_util = DumperMegatronUtil(args, model, DumperPhase.FWD_BWD, rollout_id=rollout_id) + forward_step = _build_train_forward_step(args, num_mbs, dumper_phase_util) + losses_reduced += forward_backward_func( + forward_step_func=forward_step, + data_iterator=data_iterator, + model=model, + num_microbatches=num_mbs, + seq_length=args.seq_length, + micro_batch_size=args.micro_batch_size, + decoder_seq_length=args.decoder_seq_length, + forward_only=False, + ) + + if mpu.is_pipeline_last_stage(ignore_virtual=True): + return aggregate_train_losses(losses_reduced) + return {} + + +def optimizer_step( + model: Sequence[DDP], + optimizer: MegatronOptimizer | None, + opt_param_scheduler: OptimizerParamScheduler | None, + learning_rate: float | None = None, +) -> dict[str, float | bool]: + """Tinker seam: apply the optimizer over gradients accumulated by + ``forward_backward_pass``, then release them. + + When ``learning_rate`` is given the client owns the LR schedule: it is set + directly on the param groups and the Megatron scheduler is NOT stepped + (single-owner rule, specs/005 design.md). + """ + args = get_args() + if optimizer is None: + return {"success": False, "grad_norm": 0.0} + + if learning_rate is not None: + for param_group in optimizer.param_groups: + param_group["lr"] = learning_rate * param_group.get("lr_mult", 1.0) + + update_successful, grad_norm, _ = optimizer.step() + + if update_successful and learning_rate is None and opt_param_scheduler is not None: + opt_param_scheduler.step(increment=args.global_batch_size) + + for model_chunk in model: + model_chunk.zero_grad_buffer() + optimizer.zero_grad() + + if isinstance(grad_norm, torch.Tensor): + grad_norm = grad_norm.item() + return {"success": bool(update_successful), "grad_norm": float(grad_norm) if grad_norm is not None else 0.0} + + def train( rollout_id: int, model: Sequence[DDP], @@ -624,26 +760,7 @@ def train( model_module.train() # Setup some training config params. - config = get_model_config(model[0]) - config.grad_scale_func = None if disable_optimizer else optimizer.scale_loss - config.timers = None - if isinstance(model[0], DDP) and args.overlap_grad_reduce: - assert config.no_sync_func is None, ( - "When overlap_grad_reduce is True, config.no_sync_func must be None; " - "a custom no_sync_func is not supported when overlapping grad-reduce" - ) - config.no_sync_func = [model_chunk.no_sync for model_chunk in model] - if len(model) == 1: - config.no_sync_func = config.no_sync_func[0] - if args.align_grad_reduce: - config.grad_sync_func = [model_chunk.start_grad_sync for model_chunk in model] - if len(model) == 1: - config.grad_sync_func = config.grad_sync_func[0] - if args.overlap_param_gather and args.align_param_gather: - config.param_sync_func = [model_chunk.start_param_sync for model_chunk in model] - if len(model) == 1: - config.param_sync_func = config.param_sync_func[0] - config.finalize_model_grads_func = finalize_model_grads_with_empty_cache + config = _configure_train_step(model, optimizer, args, disable_optimizer) pre_hook_enabled = False diff --git a/miles/backends/training_utils/data.py b/miles/backends/training_utils/data.py index b5dccb80610..dfe7c6c613f 100644 --- a/miles/backends/training_utils/data.py +++ b/miles/backends/training_utils/data.py @@ -151,6 +151,12 @@ def get_batch( if "dynamic_global_batch_size" in data_iterator.rollout_data: batch["dynamic_global_batch_size"] = data_iterator.rollout_data["dynamic_global_batch_size"] + # Tinker seam: forward per-request scalar metadata to the loss layer + # (same pattern as dynamic_global_batch_size above). + for tinker_key in ("_loss_type_override", "_loss_norm_total"): + if tinker_key in data_iterator.rollout_data: + batch[tinker_key] = data_iterator.rollout_data[tinker_key] + # No-op safety net if batches reach get_batch without rollout-level preprocessing. expand_multimodal_rollout_data_in_place(batch, qkv_format=qkv_format) diff --git a/miles/backends/training_utils/loss.py b/miles/backends/training_utils/loss.py index 270fcde3e44..45543542834 100644 --- a/miles/backends/training_utils/loss.py +++ b/miles/backends/training_utils/loss.py @@ -134,7 +134,7 @@ def loss_function( batch.get("max_seq_lens", None), ) - func = get_loss_function(args) + func = get_loss_function(args, batch.get("_loss_type_override")) if args.recompute_loss_function: loss, log = checkpoint( @@ -154,6 +154,10 @@ def loss_function( # Here we need to divide by cp_size because to cancel the multiply in Megatron. assert args.use_dynamic_global_batch_size == ("dynamic_global_batch_size" in batch) global_batch_size = batch.get("dynamic_global_batch_size", args.global_batch_size) + # Tinker seam: explicit normalization override. _loss_norm_total=1 gives + # pure-sum gradients, which are invariant to how a logical batch is split + # across forward_backward calls (specs/005 design.md, G1). + global_batch_size = batch.get("_loss_norm_total", global_batch_size) if not args.calculate_per_token_loss: if apply_megatron_loss_scaling: loss_parallel_size = ( diff --git a/miles/backends/training_utils/loss_hub/losses.py b/miles/backends/training_utils/loss_hub/losses.py index 6c52c261e11..38626ac2627 100644 --- a/miles/backends/training_utils/loss_hub/losses.py +++ b/miles/backends/training_utils/loss_hub/losses.py @@ -474,8 +474,11 @@ def sft_loss_function( ) -def get_loss_function(args: Namespace) -> LossFunction: - match args.loss_type: +def get_loss_function(args: Namespace, loss_type: str | None = None) -> LossFunction: + # loss_type overrides args.loss_type when given (the Tinker seam selects + # the loss per request via the _loss_type_override rollout key). + loss_type = loss_type or args.loss_type + match loss_type: case "policy_loss": return policy_loss_function case "value_loss": @@ -485,4 +488,4 @@ def get_loss_function(args: Namespace) -> LossFunction: case "custom_loss": return load_function(args.custom_loss_function_path) case _: - raise ValueError(f"Unknown loss type: {args.loss_type}") + raise ValueError(f"Unknown loss type: {loss_type}") diff --git a/miles/ray/tinker_group.py b/miles/ray/tinker_group.py new file mode 100644 index 00000000000..dae22de9cb7 --- /dev/null +++ b/miles/ray/tinker_group.py @@ -0,0 +1,58 @@ +# Tinker seam: RayTrainGroup extension exposing decoupled train-step +# primitives for the Tinker API (specs/005 in tinker-nemorl). Additive only — +# the frozen v1 RayTrainGroup is untouched; this subclass rides its +# _broadcast fanout. Pattern: N x forward_backward_only accumulate gradients +# on the actors, then one apply_optimizer_step applies them (per-call LR). + +from miles.ray.actor_group import RayTrainGroup + + +def merge_dp_sample_outputs(results: list[dict], key: str = "log_probs") -> list: + """Reassemble per-sample outputs from DP-sharded actor results into the + client's original submission order. + + Each actor returns its shard's outputs plus ``partition_indices`` (the + original indices assigned to its DP rank). TP/PP replicas of the same DP + rank return duplicates or empty lists; last writer wins, which is safe + because duplicates carry identical values. + """ + merged = {} + for result in results: + if not result: + continue + outputs = result.get(key) or [] + indices = result.get("partition_indices") or [] + if outputs and len(outputs) == len(indices): + for original_idx, output in zip(indices, outputs, strict=True): + merged[original_idx] = output + return [merged[i] for i in sorted(merged)] + + +class TinkerTrainGroup(RayTrainGroup): + """RayTrainGroup + the Tinker orchestration surface.""" + + async def forward_backward_only(self, rollout_id, rollout_data_ref): + """Fan out fwd+bwd (no optimizer step) to all actors.""" + return await self._broadcast("forward_backward_only", rollout_id, rollout_data_ref) + + async def apply_optimizer_step(self, learning_rate: float | None = None): + """Apply the optimizer over accumulated grads on all actors. + + Returns one {"success", "grad_norm"} dict per actor. + """ + return await self._broadcast("apply_optimizer_step", learning_rate=learning_rate) + + async def apply_optimizer_step_and_sync(self, learning_rate: float | None = None, rollout_id=None): + """Optimizer step + push updated weights to the inference engines.""" + results = await self.apply_optimizer_step(learning_rate=learning_rate) + await self.update_weights(rollout_id) + return results + + async def forward_logprobs(self, rollout_id, rollout_data_ref): + """Forward-only per-sample log-probs, reassembled into client order.""" + results = await self._broadcast("forward_logprobs", rollout_id, rollout_data_ref) + return merge_dp_sample_outputs(results, key="log_probs") + + async def load_checkpoint(self, checkpoint_path: str): + """Full resume (model + optimizer + scheduler) on all actors.""" + return await self._broadcast("load_checkpoint", checkpoint_path) diff --git a/miles/utils/data.py b/miles/utils/data.py index fd79cc5eaaa..6f03a81c685 100644 --- a/miles/utils/data.py +++ b/miles/utils/data.py @@ -300,4 +300,8 @@ def process_rollout_data( Timer().seq_lens = total_lengths rollout_data["total_lengths"] = [total_lengths[i] for i in partition] + # Tinker seam: keep this rank's original sample indices so per-sample + # outputs can be reassembled into the client's submission order. + rollout_data["_partition_indices"] = list(partition) + return rollout_data From ac4746bc53112d77de76da6a6decb13abee7b2bb Mon Sep 17 00:00:00 2001 From: "Gavin.Zhu" Date: Mon, 13 Jul 2026 07:38:20 +0000 Subject: [PATCH 2/6] tinker: DP-split passthrough for client-supplied loss inputs Tinker clients send advantages/log_probs/returns/values per sample (no actor-side compute_advantages_and_returns), and the seam's scalar keys (_loss_type_override, _loss_norm_total) must reach every rank. Add them to split_train_data_by_dp_raw's whitelists; absent keys are skipped, so non-Tinker paths are unaffected. Co-Authored-By: Claude Fable 5 --- miles/ray/rollout/train_data_conversion.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/miles/ray/rollout/train_data_conversion.py b/miles/ray/rollout/train_data_conversion.py index cf80d7890a3..3880e38963b 100644 --- a/miles/ray/rollout/train_data_conversion.py +++ b/miles/ray/rollout/train_data_conversion.py @@ -160,6 +160,13 @@ def split_train_data_by_dp_raw(args, data: dict[str, Any], *, dp_size: int) -> l "opd_reverse_kl", "seq_witness_ids", "weight_versions", + # Tinker seam: loss inputs supplied per-sample by the client + # (normally computed actor-side by compute_advantages_and_returns). + "log_probs", + "ref_log_probs", + "advantages", + "returns", + "values", ]: if key not in data: continue @@ -170,6 +177,9 @@ def split_train_data_by_dp_raw(args, data: dict[str, Any], *, dp_size: int) -> l "raw_reward", "total_lengths", "dynamic_global_batch_size", + # Tinker seam: per-request scalar metadata for the loss layer. + "_loss_type_override", + "_loss_norm_total", ]: if key not in data: continue From 220faa6e2570bdada924d15788b7862dcb64b3e1 Mon Sep 17 00:00:00 2001 From: "Gavin.Zhu" Date: Mon, 13 Jul 2026 07:46:23 +0000 Subject: [PATCH 3/6] docker: MILES_REPO build arg (build from a fork) Co-Authored-By: Claude Fable 5 --- docker/Dockerfile | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/docker/Dockerfile b/docker/Dockerfile index 945c5839f4c..26272388bfd 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -156,8 +156,9 @@ RUN cd /sgl-workspace/sglang && \ # ====================================== Install main package ============================================ +ARG MILES_REPO=https://github.com/radixark/miles.git ARG MILES_COMMIT=main -RUN git clone https://github.com/radixark/miles.git /root/miles && \ +RUN git clone ${MILES_REPO} /root/miles && \ cd /root/miles && \ git checkout ${MILES_COMMIT} && \ pip install -e . --no-deps From 66b4b1f018d6357870e181bd0c77f1aa62e746bd Mon Sep 17 00:00:00 2001 From: "Gavin.Zhu" Date: Mon, 13 Jul 2026 12:52:50 +0000 Subject: [PATCH 4/6] tinker: forward_backward_only returns per-sample response log-probs Tinker's fb contract includes per-datum logprobs (the SDK weights its chunk reduction by len(loss_fn_outputs) and the cookbook computes NLL from them). Computed via a forward-only pass on the same pre-step weights; extracting them from the loss pass itself is a follow-up optimization. Co-Authored-By: Claude Fable 5 --- miles/backends/megatron_utils/actor.py | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) diff --git a/miles/backends/megatron_utils/actor.py b/miles/backends/megatron_utils/actor.py index b4e5818147d..6cc61b8cd29 100644 --- a/miles/backends/megatron_utils/actor.py +++ b/miles/backends/megatron_utils/actor.py @@ -676,7 +676,12 @@ def load_other_checkpoint(self, model_tag: str, path: str) -> None: @with_logs def forward_backward_only(self, rollout_id: int, rollout_data_ref: Box) -> dict: - """Forward + backward WITHOUT the optimizer step; grads accumulate.""" + """Forward + backward WITHOUT the optimizer step; grads accumulate. + + Also returns per-sample response log-probs (Tinker's fb contract). + Computed via a forward-only pass on the same (pre-step) weights; can + later be extracted from the loss pass itself to save a forward. + """ self._heartbeat.bump() if self.args.offload_train: self.wake_up() @@ -686,12 +691,19 @@ def forward_backward_only(self, rollout_id: int, rollout_data_ref: Box) -> dict: data_iterator, num_microbatches = get_data_iterator(self.args, self.model, rollout_data) + log_probs_result = self.compute_log_prob(data_iterator, num_microbatches, rollout_id=rollout_id) + log_probs = log_probs_result.get("log_probs") or [] + + for iterator in data_iterator: + iterator.reset() + with timer("forward_backward_only"): loss_dict = forward_backward_pass( rollout_id, data_iterator, self.model, self.optimizer, num_microbatches ) return { + "log_probs": [t.cpu() for t in log_probs], "loss": loss_dict, "partition_indices": rollout_data.get("_partition_indices", []), } From 25cf7176442c893884fd312019f44aca01748700 Mon Sep 17 00:00:00 2001 From: "Gavin.Zhu" Date: Tue, 14 Jul 2026 00:19:06 +0000 Subject: [PATCH 5/6] tinker: defer DP grad finalization to optimizer_step MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fixes multi-fb grad accumulation on bridge-LoRA (G1 ratio 0.698 ~ 1/sqrt(2)). Probe evidence: fb2-entry buffer norms exactly equal fb1-exit (nothing zeroes), but fb2-exit barely moves / can DECREASE (219.23->219.35; 199.85->173.85) instead of matching the fresh fb(8) norms — the per-fb finalize (reduce-scatter) rewrites local grads with reduced shards, so the next backward accumulates onto corrupted contents. The raw full-FT path survived only by ddp_config luck (exact 1.0). Fix mirrors megatron's own microbatch accumulation: fb passes run with a no-op finalize_model_grads_func (restored after), and optimizer_step calls finalize_model_grads once over the locally-accumulated buffers before optimizer.step(). Runtime-verified 2026-07-14 (ns.config 4xH200, DP=4, Qwen2.5-0.5B bridge-LoRA): accumulation probe grad_norm bit-identical for 1xfb(8) vs 2xfb(4) + step (ratio exactly 1.0, was 0.698); full G1/G2 gates + sl_basic recipe rerun green. Co-Authored-By: Claude Fable 5 --- miles/backends/megatron_utils/model.py | 53 +++++++++++++++++--------- 1 file changed, 34 insertions(+), 19 deletions(-) diff --git a/miles/backends/megatron_utils/model.py b/miles/backends/megatron_utils/model.py index 635b3f2fe76..49727788419 100644 --- a/miles/backends/megatron_utils/model.py +++ b/miles/backends/megatron_utils/model.py @@ -646,11 +646,13 @@ def forward_backward_pass( ) -> dict[str, float]: """Tinker seam: forward + backward WITHOUT the optimizer step. - Gradients accumulate in the DDP grad buffers across successive calls until - ``optimizer_step`` runs (which zeroes them). Split-invariance of the - accumulated gradient across calls relies on the loss normalization being - call-independent — the Tinker path pins it via the ``_loss_norm_total`` - rollout key (see training_utils/loss.py); see specs/005 design.md (G1). + Gradients accumulate LOCALLY in the DDP grad buffers across successive + calls until ``optimizer_step`` runs. DP grad finalization (reduce-scatter/ + all-reduce) is DEFERRED to ``optimizer_step``: reducing per fb call + rewrites local grads with reduced shards, corrupting the buffer as an + accumulation substrate (measured as a ~1/sqrt(2) grad-norm regression on + bridge-LoRA; specs/005 design.md R1). This mirrors megatron's own + microbatch accumulation, which syncs only on the final pass. """ args = get_args() disable_optimizer = args.debug_disable_optimizer or optimizer is None @@ -660,7 +662,7 @@ def forward_backward_pass( for model_module in model: model_module.train() - _configure_train_step(model, optimizer, args, disable_optimizer) + config = _configure_train_step(model, optimizer, args, disable_optimizer) if zero_grads: for model_chunk in model: @@ -670,25 +672,34 @@ def forward_backward_pass( forward_backward_func = get_forward_backward_func() losses_reduced = [] - for num_mbs in num_microbatches: - dumper_phase_util = DumperMegatronUtil(args, model, DumperPhase.FWD_BWD, rollout_id=rollout_id) - forward_step = _build_train_forward_step(args, num_mbs, dumper_phase_util) - losses_reduced += forward_backward_func( - forward_step_func=forward_step, - data_iterator=data_iterator, - model=model, - num_microbatches=num_mbs, - seq_length=args.seq_length, - micro_batch_size=args.micro_batch_size, - decoder_seq_length=args.decoder_seq_length, - forward_only=False, - ) + config.finalize_model_grads_func = _skip_finalize_model_grads + try: + for num_mbs in num_microbatches: + dumper_phase_util = DumperMegatronUtil(args, model, DumperPhase.FWD_BWD, rollout_id=rollout_id) + forward_step = _build_train_forward_step(args, num_mbs, dumper_phase_util) + losses_reduced += forward_backward_func( + forward_step_func=forward_step, + data_iterator=data_iterator, + model=model, + num_microbatches=num_mbs, + seq_length=args.seq_length, + micro_batch_size=args.micro_batch_size, + decoder_seq_length=args.decoder_seq_length, + forward_only=False, + ) + finally: + config.finalize_model_grads_func = finalize_model_grads_with_empty_cache if mpu.is_pipeline_last_stage(ignore_virtual=True): return aggregate_train_losses(losses_reduced) return {} +def _skip_finalize_model_grads(*args, **kwargs): + """No-op finalize for the Tinker seam's fb passes (deferred to step).""" + return None + + def optimizer_step( model: Sequence[DDP], optimizer: MegatronOptimizer | None, @@ -710,6 +721,10 @@ def optimizer_step( for param_group in optimizer.param_groups: param_group["lr"] = learning_rate * param_group.get("lr_mult", 1.0) + # Deferred DP grad finalization for everything forward_backward_pass + # accumulated locally (see its docstring). + finalize_model_grads_with_empty_cache(list(model), None) + update_successful, grad_norm, _ = optimizer.step() if update_successful and learning_rate is None and opt_param_scheduler is not None: From 926f499337cbcdfe5641975d8926d00df9917e90 Mon Sep 17 00:00:00 2001 From: "Gavin.Zhu" Date: Tue, 14 Jul 2026 00:26:56 +0000 Subject: [PATCH 6/6] tinker seam: self-contained comments (drop external spec references) Co-Authored-By: Claude Fable 5 --- miles/backends/megatron_utils/actor.py | 4 ++-- miles/backends/megatron_utils/model.py | 4 ++-- miles/backends/training_utils/loss.py | 2 +- miles/ray/tinker_group.py | 2 +- 4 files changed, 6 insertions(+), 6 deletions(-) diff --git a/miles/backends/megatron_utils/actor.py b/miles/backends/megatron_utils/actor.py index 6cc61b8cd29..ff9e86f9263 100644 --- a/miles/backends/megatron_utils/actor.py +++ b/miles/backends/megatron_utils/actor.py @@ -670,8 +670,8 @@ def load_other_checkpoint(self, model_tag: str, path: str) -> None: self._active_model_tag = model_tag # --- Tinker seam ------------------------------------------------------- - # Decoupled train-step primitives for the Tinker API (specs/005 in - # tinker-nemorl): N x forward_backward_only accumulate gradients, one + # Decoupled train-step primitives for the Tinker API: + # N x forward_backward_only accumulate gradients, one # apply_optimizer_step applies them with a per-call LR. @with_logs diff --git a/miles/backends/megatron_utils/model.py b/miles/backends/megatron_utils/model.py index 49727788419..28ccfa0c812 100644 --- a/miles/backends/megatron_utils/model.py +++ b/miles/backends/megatron_utils/model.py @@ -651,7 +651,7 @@ def forward_backward_pass( all-reduce) is DEFERRED to ``optimizer_step``: reducing per fb call rewrites local grads with reduced shards, corrupting the buffer as an accumulation substrate (measured as a ~1/sqrt(2) grad-norm regression on - bridge-LoRA; specs/005 design.md R1). This mirrors megatron's own + bridge-LoRA). This mirrors megatron's own microbatch accumulation, which syncs only on the final pass. """ args = get_args() @@ -711,7 +711,7 @@ def optimizer_step( When ``learning_rate`` is given the client owns the LR schedule: it is set directly on the param groups and the Megatron scheduler is NOT stepped - (single-owner rule, specs/005 design.md). + (single-owner rule: exactly one of client/scheduler drives the LR). """ args = get_args() if optimizer is None: diff --git a/miles/backends/training_utils/loss.py b/miles/backends/training_utils/loss.py index 45543542834..3249d031387 100644 --- a/miles/backends/training_utils/loss.py +++ b/miles/backends/training_utils/loss.py @@ -156,7 +156,7 @@ def loss_function( global_batch_size = batch.get("dynamic_global_batch_size", args.global_batch_size) # Tinker seam: explicit normalization override. _loss_norm_total=1 gives # pure-sum gradients, which are invariant to how a logical batch is split - # across forward_backward calls (specs/005 design.md, G1). + # across forward_backward calls. global_batch_size = batch.get("_loss_norm_total", global_batch_size) if not args.calculate_per_token_loss: if apply_megatron_loss_scaling: diff --git a/miles/ray/tinker_group.py b/miles/ray/tinker_group.py index dae22de9cb7..bdef469162f 100644 --- a/miles/ray/tinker_group.py +++ b/miles/ray/tinker_group.py @@ -1,5 +1,5 @@ # Tinker seam: RayTrainGroup extension exposing decoupled train-step -# primitives for the Tinker API (specs/005 in tinker-nemorl). Additive only — +# primitives for the Tinker API. Additive only — # the frozen v1 RayTrainGroup is untouched; this subclass rides its # _broadcast fanout. Pattern: N x forward_backward_only accumulate gradients # on the actors, then one apply_optimizer_step applies them (per-call LR).