diff --git a/skills/configure/references/path-intake.md b/skills/configure/references/path-intake.md index 458d9d90..88b6f9c1 100644 --- a/skills/configure/references/path-intake.md +++ b/skills/configure/references/path-intake.md @@ -145,3 +145,16 @@ Every confirmed plan includes a **Path** section: | GA vs preview | Managed GA; Training API preview if applicable | Never show only `firectl` commands when the user chose SDK or Training API. + +### DPO recipe on serverless trainers + +`training.recipes.dpo_loop.Config(serverless=True)` uses one positive-rank LoRA +policy run and requires an explicit `max_seq_len`. It saves `dpo-reference` +before the first policy update and scores that fixed snapshot through a sampling +client. The reference is inference-only; no second training run is created. +Reference token logprobs are aligned to the trainer targets before applying the +existing DPO loss. Warm starts and resume are rejected on this initial path so +the reference cannot accidentally be reanchored to an already-trained policy. +The snapshot is internal reference state, not the final trained output. Policy +training and sampler scoring have separate usage paths; do not quote the +reference as free or included in training-token usage. ORPO remains unchanged. diff --git a/skills/fireworks-training/references/rl-concurrency.md b/skills/fireworks-training/references/rl-concurrency.md index a100d884..c2d146e3 100644 --- a/skills/fireworks-training/references/rl-concurrency.md +++ b/skills/fireworks-training/references/rl-concurrency.md @@ -90,16 +90,12 @@ Request uploads and future-result downloads have separate contracts: client fetches the full body. The client can then reserve download capacity and avoid concurrent multi-megabyte response fan-out. -For a large batch that produces large future results, prefer metadata-only -retrieval after confirming that the selected trainer and client both support -the protocol. Do not describe it as a generic large-batch or upload-timeout -mode. - -The current staged cookbook intentionally disables metadata-only retrieval -while the trainer fleet is being upgraded. It is not a cookbook default, and -the cookbook does not currently expose a supported public toggle. If Fireworks -has approved the staged rollout for a compatible trainer/client pair, enable it -manually using the rollout-specific instructions. Do not monkeypatch private -`tinker` internals. Otherwise, keep the compatibility behavior until the -backend rollout is complete and the cookbook publishes a supported opt-in or -changes the default. +Metadata-only retrieval is the default for supported Tinker clients and +trainers. The client reserves the advertised response size before fetching the +full body, bounding concurrent result downloads without changing forward +execution or replaying the request. The cookbook does not expose a separate +toggle and must not monkeypatch private `tinker` internals. + +Do not describe metadata-only retrieval as a generic large-batch or +upload-timeout mode. It controls completed future-result downloads, not request +uploads or trainer execution. diff --git a/training/recipes/dpo_loop.py b/training/recipes/dpo_loop.py index 6e4298a6..166d576a 100644 --- a/training/recipes/dpo_loop.py +++ b/training/recipes/dpo_loop.py @@ -43,6 +43,7 @@ import time from contextlib import ExitStack from dataclasses import dataclass, field +from types import SimpleNamespace from typing import Any, Callable import tinker @@ -85,6 +86,7 @@ ) from training.utils.checkpoints import TrainingCheckpoints, validate_warm_start_config from training.utils.resource_autosizing import select_render_worker_count +from training.utils.serverless import setup_serverless_training from training.utils.runner_state import write_completed, write_running_step from training.utils.timer import flush_timing, timer @@ -135,6 +137,8 @@ class Config: max_seq_len: int | None = None max_pairs: int | None = None """Cap on *valid rendered pairs* after schema/length filtering.""" + serverless: bool = False + """Use one pooled policy run and a frozen snapshot sampler for reference scoring.""" lora_rank: int = 0 lora_alpha: int | None = 32 """LoRA alpha scaling factor. Ignored when ``lora_rank == 0``. @@ -272,9 +276,40 @@ def _render_pair_worker(row: dict[str, Any]) -> dict[str, Any] | None: # --------------------------------------------------------------------------- +class SamplerReference: + """Score a frozen sampler snapshot in the trainer-forward token layout. + + Calls run inside the existing reference semaphore. Submit/wait per sequence + so ref_cache_concurrency bounds sampler requests, independent of batch size. + """ + + def __init__(self, sampler: Any, *, timeout: int) -> None: + self._sampler = sampler + self._timeout = timeout + + def forward(self, datums: list[tinker.Datum], loss_fn: str) -> SimpleNamespace: + if loss_fn != "cross_entropy": + raise ValueError("sampler reference supports only cross_entropy scoring") + outputs = [] + for datum in datums: + targets = datum.loss_fn_inputs["target_tokens"].data + if not targets: + raise ValueError("reference datum must contain target tokens") + sequence = datum.model_input.append_int(int(targets[-1])) + values = self._sampler.compute_logprobs(sequence).result(timeout=self._timeout) + if len(values) != len(targets) + 1: + raise ValueError("reference logprob length does not match target tokens") + # Position 0 has no preceding token. Every target position must be scored. + aligned = values[1:] + if any(value is None for value in aligned): + raise ValueError("reference sampler returned an unscored target token") + outputs.append({"logprobs": SimpleNamespace(data=aligned)}) + return SimpleNamespace(loss_fn_outputs=outputs) + + async def _ref_forward_batch( pairs: list[dict[str, Any]], - reference: ReconnectableClient, + reference: ReconnectableClient | SamplerReference, semaphore: asyncio.Semaphore, ref_batch_size: int, ) -> list[dict[str, Any]]: @@ -354,7 +389,7 @@ def _forward_backward_pairs( async def _train_loop( pair_dataset: JsonlRenderDataset, ref_cache_log: AppendOnlyPickleLog | None, - reference: ReconnectableClient, + reference: ReconnectableClient | SamplerReference, policy: ReconnectableClient, adam_params: tinker.AdamParams, cfg: Config, @@ -672,6 +707,8 @@ def main( ): cfg = config _validate_dpo_beta(cfg.beta) + if cfg.serverless and (cfg.init_from_checkpoint or cfg.warm_start_from_adapter): + raise ValueError("serverless DPO does not support warm starts or resume") # Internal shape validation can request a resumable-only live handoff # without expanding the public cookbook/control-plane Config contract. final_checkpoint_promotable = getattr(cfg, "_save_final_checkpoint_promotable", True) @@ -737,70 +774,96 @@ def _signal_handler(signum, frame): runner.write_status(RunStatus.PENDING, message="provisioning") with runner, ExitStack() as stack: - service = build_service_client( - api_key=api_key, - base_url=base_url, - additional_headers=additional_headers, - base_model=cfg.base_model, - tokenizer_model=cfg.tokenizer_model, - max_lora_rank=cfg.lora_rank, - max_context_length=cfg.max_seq_len, - learning_rate=cfg.learning_rate, - trainer=cfg.trainer, - reference_required=True, - cleanup_trainer_on_close=cfg.cleanup_on_exit, - ) - stack.callback(service.close) - training_client = service.create_training_client( - cfg.base_model, - lora_rank=cfg.lora_rank, - lora_alpha=cfg.lora_alpha, - ) - runner.set_accelerator_info( - service.accelerator_type, - service.accelerator_count, - profile=service.training_profile, - ) - policy_job_id = service.trainer_job_id - max_seq_len = service.max_context_length - - policy = ReconnectableClient.from_training_client( - training_client, - base_model=cfg.base_model, - lora_rank=cfg.lora_rank, - job_id=policy_job_id, - default_timeout=cfg.step_timeout or 3600, - service=service, - ) - # DPO always needs a reference. The SDK owns the shared-vs-separate - # decision: LoRA without an explicit reference shape reuses the policy - # session; full-param (or an explicit reference_training_shape_id) - # provisions a separate frozen reference trainer that `service` owns. - # Backend trainer creation selects a LoRA-capable shape unless - # cfg.trainer.reference_training_shape_id pins a LoRA-capable shape. - reference = ReconnectableClient.from_training_client( - service.create_reference_client(policy_client=training_client), - base_model=cfg.base_model, - lora_rank=0, - job_id=service.reference_client_job_id, - default_timeout=cfg.step_timeout or 3600, - service=service, - base_only=True, - ) - reference_job_id = service.reference_trainer_job_id - - ckpt = TrainingCheckpoints( - policy, - service, - trainer_id=policy_job_id, - log_path=cfg.log_path, - lora_rank=cfg.lora_rank, - ) + if cfg.serverless: + service, policy, ckpt, policy_job_id, max_seq_len = setup_serverless_training( + cfg, + api_key=api_key, + base_url=base_url, + additional_headers=additional_headers, + stack=stack, + ) + stack.callback(service.close) + runner.set_accelerator_info(None, None, profile=None) + runner.mark_serverless() + # Snapshot before the first update: fresh LoRA has zero adapter effect. + # Keep this exact identity for the entire job; never resnapshot the policy. + reference_path = policy.save_weights_for_sampler("dpo-reference").path + if not reference_path: + raise RuntimeError("reference snapshot returned no path") + sampler = service.create_sampling_client(model_path=reference_path) + stack.callback(sampler.close) + reference = SamplerReference(sampler, timeout=cfg.step_timeout or 3600) + reference_job_id = None + runner.set_reference_checkpoint(reference_path) + runner.write_metadata() + logger.info("Frozen DPO reference snapshot: %s", reference_path) + else: + service = build_service_client( + api_key=api_key, + base_url=base_url, + additional_headers=additional_headers, + base_model=cfg.base_model, + tokenizer_model=cfg.tokenizer_model, + max_lora_rank=cfg.lora_rank, + max_context_length=cfg.max_seq_len, + learning_rate=cfg.learning_rate, + trainer=cfg.trainer, + reference_required=True, + cleanup_trainer_on_close=cfg.cleanup_on_exit, + ) + stack.callback(service.close) + training_client = service.create_training_client( + cfg.base_model, + lora_rank=cfg.lora_rank, + lora_alpha=cfg.lora_alpha, + ) + runner.set_accelerator_info( + service.accelerator_type, + service.accelerator_count, + profile=service.training_profile, + ) + policy_job_id = service.trainer_job_id + max_seq_len = service.max_context_length + + policy = ReconnectableClient.from_training_client( + training_client, + base_model=cfg.base_model, + lora_rank=cfg.lora_rank, + job_id=policy_job_id, + default_timeout=cfg.step_timeout or 3600, + service=service, + ) + # DPO always needs a reference. The SDK owns the shared-vs-separate + # decision: LoRA without an explicit reference shape reuses the policy + # session; full-param (or an explicit reference_training_shape_id) + # provisions a separate frozen reference trainer that `service` owns. + # Backend trainer creation selects a LoRA-capable shape unless + # cfg.trainer.reference_training_shape_id pins a LoRA-capable shape. + reference = ReconnectableClient.from_training_client( + service.create_reference_client(policy_client=training_client), + base_model=cfg.base_model, + lora_rank=0, + job_id=service.reference_client_job_id, + default_timeout=cfg.step_timeout or 3600, + service=service, + base_only=True, + ) + reference_job_id = service.reference_trainer_job_id + + ckpt = TrainingCheckpoints( + policy, + service, + trainer_id=policy_job_id, + log_path=cfg.log_path, + lora_rank=cfg.lora_rank, + ) - resume_info = ckpt.resume( - init_from_checkpoint=cfg.init_from_checkpoint, - warm_start_from_adapter=cfg.warm_start_from_adapter, - ) + resume_info = None + if not cfg.serverless: + resume_info = ckpt.resume( + init_from_checkpoint=cfg.init_from_checkpoint, + warm_start_from_adapter=cfg.warm_start_from_adapter, + ) step_offset = resume_info.step if resume_info else 0 wandb_log({"train/step": step_offset}, step_offset) adam_kwargs = dict(DEFAULT_ADAM) @@ -867,7 +930,10 @@ def _on_ref_done() -> None: nonlocal reference_job_id if not cfg.release_reference_after_cache: return - service.release_references() + if cfg.serverless: + sampler.close() + else: + service.release_references() reference_job_id = None cursor = RawRowCursor(max_rows=len(pair_dataset) * cfg.epochs) @@ -898,7 +964,10 @@ def _on_ref_done() -> None: ) if getattr(cfg, "output_model_id", None): - ckpt.promote_latest(cfg.output_model_id, cfg.base_model) + if cfg.serverless: + ckpt.promote_latest(cfg.output_model_id, cfg.base_model, checkpoint_name=cp_name) + else: + ckpt.promote_latest(cfg.output_model_id, cfg.base_model) runner.write_output_model( model_id=cfg.output_model_id, checkpoint=cp_name, job_id=policy_job_id, ) diff --git a/training/tests/unit/test_checkpoints.py b/training/tests/unit/test_checkpoints.py index 757f065f..004c3476 100644 --- a/training/tests/unit/test_checkpoints.py +++ b/training/tests/unit/test_checkpoints.py @@ -1003,3 +1003,28 @@ def test_bounded_history_keeps_new_checkpoint_after_recipe_step_reset( source_job_id="job-1", ) client.load_state_with_optimizer.assert_called_once_with("path://self/step-0") + + +def test_exact_serverless_promotion_never_uses_reference(log_dir): + rows = [ + _row( + "run-current-dpo-reference", + ctype="CHECKPOINT_TYPE_INFERENCE_LORA", + promotable=True, + create_time="2099-01-01T00:00:00Z", + ) + ] + ckpt, _, fw = _make(log_dir, fw_rows=rows, lora_rank=8, serverless=True, current_run_id="run-current") + with pytest.raises(RuntimeError, match="No promotable checkpoints"): + ckpt.promote_latest("output", "base", checkpoint_name="step-10") + fw.promote_checkpoint.assert_not_called() + fw._rows.append( + _row( + "run-current-step-10", + ctype="CHECKPOINT_TYPE_INFERENCE_LORA", + promotable=True, + create_time="2026-01-01T00:00:00Z", + ) + ) + ckpt.promote_latest("output", "base", checkpoint_name="step-10") + assert fw.promote_checkpoint.call_args.kwargs["name"].endswith("/run-current-step-10") diff --git a/training/tests/unit/test_client.py b/training/tests/unit/test_client.py index b7eff043..77d19b0a 100644 --- a/training/tests/unit/test_client.py +++ b/training/tests/unit/test_client.py @@ -9,12 +9,12 @@ from training.utils.client import ReconnectableClient -def test_client_disables_tinker_metadata_only_future_retrieval_by_default(): +def test_client_preserves_tinker_metadata_only_future_retrieval_by_default(): request = tinker_api_future_impl.FutureRetrieveRequest( request_id="req-1", allow_metadata_only=True, ) - assert request.allow_metadata_only is False + assert request.allow_metadata_only is True class _FakeFuture: diff --git a/training/tests/unit/test_dpo_loop.py b/training/tests/unit/test_dpo_loop.py index 13f05522..0c34e62e 100644 --- a/training/tests/unit/test_dpo_loop.py +++ b/training/tests/unit/test_dpo_loop.py @@ -25,6 +25,10 @@ import json import time from types import SimpleNamespace +from concurrent.futures import Future +from unittest.mock import Mock + +import tinker import pytest @@ -1327,3 +1331,136 @@ def test_rejects_non_object_jsonl_row(self, tmp_path) -> None: with pytest.raises(module.DatasetError, match="row must be an object"): load_preference_dataset(str(path)) + + +def test_sampler_reference_aligns_targets_and_preserves_pair_order(): + + calls = [] + + class Sampler: + def compute_logprobs(self, sequence): + tokens = sequence.to_ints() + calls.append(tokens) + future = Future() + future.set_result([None] + [-float(token) for token in tokens[1:]]) + return future + + def datum(tokens): + return tinker.Datum( + model_input=tinker.ModelInput.from_ints(tokens[:-1]), + loss_fn_inputs={ + "target_tokens": tinker.TensorData(data=tokens[1:], dtype="int64", shape=[len(tokens) - 1]) + }, + ) + + chosen, rejected = datum([1, 2, 3]), datum([1, 4, 5]) + reference = module.SamplerReference(Sampler(), timeout=10) + pairs = [ + { + "chosen_datum": chosen, + "rejected_datum": rejected, + "chosen_tokens_len": 3, + "rejected_tokens_len": 3, + "response_start": 1, + } + ] + scored = asyncio.run(module._ref_forward_batch(pairs, reference, asyncio.Semaphore(1), 1)) + assert calls == [[1, 2, 3], [1, 4, 5]] + assert list(scored[0]["ref_chosen"]) == [-2, -3] + assert list(scored[0]["ref_rejected"]) == [-4, -5] + assert scored[0]["chosen_datum"] is chosen + assert scored[0]["response_start"] == 1 + + +@pytest.mark.parametrize("values", [[None, -1.0], [None, -1.0, None]]) +def test_sampler_reference_rejects_missing_scores(values): + + future = Future() + future.set_result(values) + sampler = Mock() + sampler.compute_logprobs.return_value = future + datum = tinker.Datum( + model_input=tinker.ModelInput.from_ints([1, 2]), + loss_fn_inputs={"target_tokens": tinker.TensorData(data=[2, 3], dtype="int64", shape=[2])}, + ) + with pytest.raises(ValueError, match="logprob length|unscored target"): + module.SamplerReference(sampler, timeout=10).forward([datum], "cross_entropy") + + +@pytest.mark.parametrize("field", ["init_from_checkpoint", "warm_start_from_adapter"]) +def test_serverless_dpo_rejects_resume_before_setup(field): + cfg = module.Config(log_path="unused", serverless=True) + setattr(cfg, field, "prior-policy") + with pytest.raises(ValueError, match="does not support warm starts or resume"): + module.main(cfg) + + +@pytest.mark.parametrize("failure", [None, "snapshot", "sampler", "training"]) +def test_serverless_main_freezes_reference_and_finalizes_policy(monkeypatch, tmp_path, failure): + + monkeypatch.setenv("FIREWORKS_API_KEY", "test-key") + for name in [ + "setup_wandb", + "wandb_log", + "wandb_finish", + "validate_config", + "validate_warm_start_config", + "_init_pair_worker", + ]: + monkeypatch.setattr(module, name, Mock()) + monkeypatch.setattr(module, "resolve_renderer_snapshot", lambda **kwargs: "qwen3") + service, policy, ckpt, sampler = Mock(), Mock(), Mock(), Mock() + policy.save_weights_for_sampler.return_value = SimpleNamespace(path="snapshot://initial-reference") + service.create_sampling_client.return_value = sampler + if failure == "snapshot": + policy.save_weights_for_sampler.side_effect = RuntimeError("snapshot failed") + if failure == "sampler": + service.create_sampling_client.side_effect = RuntimeError("sampler failed") + monkeypatch.setattr( + module, "setup_serverless_training", lambda *args, **kwargs: (service, policy, ckpt, "ts-policy", 512) + ) + monkeypatch.setattr(module, "build_service_client", Mock(side_effect=AssertionError("dedicated provisioning"))) + monkeypatch.setattr(module, "JsonlRenderDataset", lambda *args: ["pair"]) + runner = Mock() + runner.__enter__ = Mock(return_value=runner) + runner.__exit__ = Mock(return_value=False) + monkeypatch.setattr(module, "RunnerIO", lambda *args: runner) + + async def train(*args, **kwargs): + policy.save_weights_for_sampler.assert_called_once_with("dpo-reference") + service.create_sampling_client.assert_called_once_with(model_path="snapshot://initial-reference") + assert isinstance(args[2], module.SamplerReference) + assert args[3] is policy + ckpt.resume.assert_not_called() + if failure == "training": + raise RuntimeError("training failed") + kwargs["on_ref_done"]() + return 1 + + monkeypatch.setattr(module, "_train_loop", train) + cfg = module.Config( + log_path=str(tmp_path), + dataset="preferences.jsonl", + tokenizer_model="Qwen/Qwen3-8B", + serverless=True, + lora_rank=8, + max_seq_len=512, + render_workers=0, + output_model_id="trained", + ) + if failure: + with pytest.raises(RuntimeError, match=f"{failure} failed"): + module.main(cfg) + service.close.assert_called_once() + ckpt.promote_latest.assert_not_called() + if failure == "training": + sampler.close.assert_called_once() + return + result = module.main(cfg) + assert result["steps"] == 1 + runner.mark_serverless.assert_called_once() + ckpt.save.assert_called_once() + ckpt.promote_latest.assert_called_once_with("trained", cfg.base_model, checkpoint_name="step-1") + service.create_reference_client.assert_not_called() + service.release_references.assert_not_called() + assert sampler.close.called and service.close.called diff --git a/training/tests/unit/test_runner.py b/training/tests/unit/test_runner.py index 919d649f..f2fac250 100644 --- a/training/tests/unit/test_runner.py +++ b/training/tests/unit/test_runner.py @@ -619,3 +619,13 @@ def test_default_runner_is_noop(self): runner.write_output_model(model_id="m") runner.start_training() runner.set_accelerator_info("H100", 8) + + +def test_fixed_reference_checkpoint_metadata(tmp_path): + path = str(tmp_path / "metadata.json") + runner = RunnerIO(RunnerConfig(metadata_file=path)) + runner.set_reference_checkpoint("snapshot://fixed") + runner.write_metadata() + assert json.loads(open(path).read())["metadata"]["reference_checkpoint"] == "snapshot://fixed" + with pytest.raises(ValueError, match="remain fixed"): + runner.set_reference_checkpoint("snapshot://updated-policy") diff --git a/training/tests/unit/test_serverless.py b/training/tests/unit/test_serverless.py index 12403201..887bc591 100644 --- a/training/tests/unit/test_serverless.py +++ b/training/tests/unit/test_serverless.py @@ -14,6 +14,9 @@ def __init__(self, *, base_url, api_key, default_headers): self.default_headers = default_headers self.training_session_id = "ts-1234" + def close(self): + pass + def create_lora_training_client(self, base_model, rank, alpha): assert base_model == "accounts/fireworks/models/qwen3-4b" assert rank == 8 diff --git a/training/utils/checkpoints.py b/training/utils/checkpoints.py index 3d124ff0..2f3e7f8c 100644 --- a/training/utils/checkpoints.py +++ b/training/utils/checkpoints.py @@ -634,9 +634,14 @@ def promote_latest( base_model: str, *, hot_load_deployment_id: str | None = None, + checkpoint_name: str | None = None, ) -> dict: """Promote the newest promotable row on the control plane. + When checkpoint_name is supplied, only that logical checkpoint in the + current run is eligible; missing final exports must not fall back to an + earlier snapshot (for example a frozen DPO reference). + No local lookup. Works identically for full and LoRA runs; in LoRA runs this transparently picks up the most recent weight-sync sampler row without requiring an explicit final @@ -649,9 +654,13 @@ def promote_latest( if r.get("promotable") and self._row_matches_current_run(r) ] ) + if checkpoint_name is not None: + rows = [r for r in rows if self._trainer_logical_name(_short_name(r["name"])) == checkpoint_name] if not rows: raise RuntimeError( - f"No promotable checkpoints found for trainer job '{self._trainer_id}'. " + f"No promotable checkpoints found for trainer job '{self._trainer_id}'" + + (f" matching {checkpoint_name!r}" if checkpoint_name is not None else "") + + ". " "Call save(promotable=True) or run a promotable weight sync first." ) # Use the 4-segment resource name end-to-end: the SDK accepts diff --git a/training/utils/client.py b/training/utils/client.py index 841328d0..aa449fb0 100644 --- a/training/utils/client.py +++ b/training/utils/client.py @@ -35,8 +35,6 @@ class GradAccNormalization(str, Enum): NUM_SEQUENCES = "num_sequences" NUM_LOSS_TOKENS = "num_loss_tokens" from fireworks.training.sdk.trainer import TrainerJobManager, TrainerServiceEndpoint -import tinker.lib.api_future_impl as tinker_api_future_impl -from tinker.types.future_retrieve_request import FutureRetrieveRequest as _FutureRetrieveRequest logger = logging.getLogger(__name__) @@ -55,23 +53,6 @@ class GradAccNormalization(str, Enum): in an importable cookbook (a77 has the method but lacks other symbols utils imports).""" -def _install_tinker_future_retrieve_compat() -> None: - """Keep metadata-only disabled until all trainer versions support it.""" - current = getattr(tinker_api_future_impl, "FutureRetrieveRequest", None) - if current is None or getattr(current, "_fw_cookbook_compat", False): - return - - def _compat_future_retrieve_request(*args, **kwargs): - kwargs.pop("allow_metadata_only", None) - return _FutureRetrieveRequest(*args, **kwargs) - - _compat_future_retrieve_request._fw_cookbook_compat = True - tinker_api_future_impl.FutureRetrieveRequest = _compat_future_retrieve_request - - -_install_tinker_future_retrieve_compat() - - class ReconnectableClient: """Training client wrapper: dispatch + wait with timeout. diff --git a/training/utils/runner.py b/training/utils/runner.py index bc0e70d6..9c90c5b5 100644 --- a/training/utils/runner.py +++ b/training/utils/runner.py @@ -219,6 +219,7 @@ def __init__(self, config: RunnerConfig | None = None): self._last_step: int = 0 self._last_total_steps: int = 0 self._serverless: bool = False + self._reference_checkpoint: str | None = None # -- context manager ------------------------------------------------------- @@ -344,6 +345,12 @@ def start_training(self) -> None: """Mark training start for accelerator-seconds calculation.""" self._training_start = time.monotonic() + def set_reference_checkpoint(self, path: str) -> None: + """Record the fixed DPO reference identity in managed metadata.""" + if self._reference_checkpoint is not None and self._reference_checkpoint != path: + raise ValueError("reference checkpoint must remain fixed during a run") + self._reference_checkpoint = path + def write_metadata(self) -> None: if not self._metadata_file: return @@ -364,6 +371,8 @@ def write_metadata(self) -> None: "serverless": self._serverless, } } + if self._reference_checkpoint is not None: + payload["metadata"]["reference_checkpoint"] = self._reference_checkpoint self._write_json(self._metadata_file, payload) def set_tokens_processed(self, tokens: int) -> None: diff --git a/training/utils/serverless.py b/training/utils/serverless.py index 1601976b..14e73671 100644 --- a/training/utils/serverless.py +++ b/training/utils/serverless.py @@ -1,4 +1,4 @@ -"""Serverless SFT setup helpers. +"""Serverless managed training setup helpers. The serverless counterpart to ``build_service_client`` in ``training/utils/service.py``: connects to a shared, already-running pooled @@ -64,12 +64,12 @@ def promote_checkpoint( def setup_serverless_training(cfg, *, api_key, base_url, additional_headers, stack): - """Build the training + checkpoint handles for a serverless SFT run. + """Build the training + checkpoint handles for a serverless LoRA run. Returns ``(service, client, ckpt, session_id, max_seq_len)``. The caller - registers ``service.close`` for teardown; the internal control-plane client - used for checkpoint list/promote is registered on the provided ``stack`` - (an ``ExitStack``) here, so it is closed on teardown too. Requires + may close the returned service; both the service and internal control-plane + client are registered on the provided ``stack`` (an ``ExitStack``), including + when setup fails after the training session is reserved. Requires ``cfg.lora_rank > 0`` and a concrete ``cfg.max_seq_len`` (there is no training shape to resolve sequence length from on this path). """ @@ -87,6 +87,7 @@ def setup_serverless_training(cfg, *, api_key, base_url, additional_headers, sta api_key=api_key, default_headers=additional_headers or None, ) + stack.callback(service.close) training_client = service.create_lora_training_client( cfg.base_model, rank=cfg.lora_rank,