From 26f69e7c6ee782efbab58a44032a5a72caf76ed8 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 18:00:58 -0700 Subject: [PATCH 01/28] feat: unified resume system with checkpoint_utils Replace dual-path resume logic (ResumeConfig + setup_resume) with a single source of truth: local_checkpoint_state.jsonl via checkpoint_utils. - Add checkpoint_utils.py: ResumeState, resolve_resume(), load_dcp(), save_loop_state(), dataset_fingerprint(), validate_dataset() - Delete resume.py and ResumeConfig dataclass - Track data_consumed in SFT and RL loops for dataset position persistence - Add init_from_checkpoint config for loading pretrained DCP weights on a fresh dataset (supports cross-job "job_id:checkpoint_name" format) - Make ref_logprobs optional in RL losses (grpo, dapo, gspo, cispo) so kl_beta=0 works without a reference model - Log perf/dcp_load_time to wandb on resume - Update all 4 recipe loops (rl, sft, dpo, orpo) and e2e resume tests Made-with: Cursor --- .../examples/frozen_lake/train_frozen_lake.py | 7 +- training/recipes/dpo_loop.py | 14 +- training/recipes/orpo_loop.py | 10 +- training/recipes/rl_loop.py | 50 +++- training/recipes/sft_loop.py | 74 ++++-- training/tests/e2e/test_dpo_resume_e2e.py | 4 +- training/tests/e2e/test_grpo_resume_e2e.py | 7 +- training/tests/e2e/test_sft_resume_e2e.py | 4 +- training/tests/unit/test_rl_loop.py | 4 +- training/utils/__init__.py | 4 - training/utils/checkpoint_utils.py | 214 ++++++++++++++++++ training/utils/client.py | 3 + training/utils/config.py | 7 - training/utils/resume.py | 49 ---- training/utils/rl/cispo.py | 2 +- training/utils/rl/dapo.py | 2 +- training/utils/rl/grpo.py | 2 +- training/utils/rl/gspo.py | 2 +- training/utils/rl/losses.py | 5 +- training/utils/validation.py | 15 +- 20 files changed, 352 insertions(+), 127 deletions(-) create mode 100644 training/utils/checkpoint_utils.py delete mode 100644 training/utils/resume.py diff --git a/training/examples/frozen_lake/train_frozen_lake.py b/training/examples/frozen_lake/train_frozen_lake.py index 51cacdc4..b626ec33 100644 --- a/training/examples/frozen_lake/train_frozen_lake.py +++ b/training/examples/frozen_lake/train_frozen_lake.py @@ -53,12 +53,10 @@ InfraConfig, WandBConfig, DeployConfig, - ResumeConfig, HotloadConfig, ReconnectableClient, wandb_log, setup_wandb, - setup_resume, wandb_finish, log_metrics_json, setup_deployment, @@ -545,7 +543,10 @@ def _make_job(label: str, precreated_id: str | None, **extra_kw): infra_boot_time = time.time() - _infra_start wandb_log({"train/step": 0, "infra/total_boot_time": infra_boot_time}, step=0) - step_offset, _ = setup_resume(policy, ResumeConfig()) + from training.utils.checkpoint_utils import resolve_resume, load_dcp + state = resolve_resume("./frozen_lake_logs") + load_dcp(policy, state) + step_offset = state.step if hotload_cfg.hot_load_before_training and deploy_cfg.deployment_id: name = f"resume-{step_offset}-base" if step_offset > 0 else "step-0-base" weight_syncer.save_and_hotload(name, checkpoint_type="base") diff --git a/training/recipes/dpo_loop.py b/training/recipes/dpo_loop.py index 5d92ee84..7d4d42a1 100644 --- a/training/recipes/dpo_loop.py +++ b/training/recipes/dpo_loop.py @@ -39,12 +39,10 @@ InfraConfig, WandBConfig, DeployConfig, - ResumeConfig, HotloadConfig, ReconnectableClient, wandb_log, setup_wandb, - setup_resume, wandb_finish, validate_config, log_metrics_json, @@ -58,6 +56,7 @@ ) from fireworks.training.sdk.deployment import DEFAULT_DELTA_COMPRESSION from fireworks.training.sdk.weight_syncer import WeightSyncer +from training.utils.checkpoint_utils import resolve_resume, load_dcp, save_loop_state, dataset_fingerprint, validate_dataset from training.utils.timer import timer, flush_timing logger = logging.getLogger(__name__) @@ -120,7 +119,8 @@ class Config: deployment: DeployConfig = field(default_factory=DeployConfig) hotload: HotloadConfig = field(default_factory=lambda: HotloadConfig(hot_load_interval=0)) wandb: WandBConfig = field(default_factory=lambda: WandBConfig(project="dpo-tinker")) - resume: ResumeConfig = field(default_factory=ResumeConfig) + log_path: str = "./dpo_logs" + init_from_checkpoint: str | None = None # --------------------------------------------------------------------------- @@ -389,7 +389,7 @@ def _signal_handler(signum, frame): signal.signal(signal.SIGTERM, _signal_handler) signal.signal(signal.SIGINT, _signal_handler) - validate_config(cfg.base_model, cfg.dataset, cfg.hotload, cfg.deployment, cfg.infra, cfg.resume) + validate_config(cfg.base_model, cfg.dataset, cfg.hotload, cfg.deployment, cfg.infra) if not cfg.tokenizer_model: raise ValueError( "Config.tokenizer_model is required for client-side tokenization. " @@ -485,7 +485,11 @@ def _signal_handler(signum, frame): compression_format=DEFAULT_DELTA_COMPRESSION, ) - step_offset, _ = setup_resume(policy, cfg.resume) + state = resolve_resume(cfg.log_path, cfg.init_from_checkpoint) + dcp_load_time = load_dcp(policy, state) + if dcp_load_time > 0: + wandb_log({"perf/dcp_load_time": dcp_load_time}, state.step) + step_offset = state.step adam_params = tinker.AdamParams(learning_rate=cfg.learning_rate, **DEFAULT_ADAM) # -- Cache reference logprobs concurrently ------------------------------ diff --git a/training/recipes/orpo_loop.py b/training/recipes/orpo_loop.py index 780f6555..8b839ca4 100644 --- a/training/recipes/orpo_loop.py +++ b/training/recipes/orpo_loop.py @@ -49,11 +49,9 @@ DEFAULT_ADAM, InfraConfig, WandBConfig, - ResumeConfig, ReconnectableClient, wandb_log, setup_wandb, - setup_resume, wandb_finish, log_metrics_json, make_orpo_loss_fn, @@ -63,6 +61,7 @@ render_preference_pair, resolve_renderer_name, ) +from training.utils.checkpoint_utils import resolve_resume, load_dcp logger = logging.getLogger(__name__) @@ -94,7 +93,8 @@ class Config: project="dsv3-training", ) ) - resume: ResumeConfig = field(default_factory=ResumeConfig) + log_path: str = "./orpo_logs" + init_from_checkpoint: str | None = None # --------------------------------------------------------------------------- @@ -175,7 +175,9 @@ def _signal_handler(signum, frame): ) job_id = endpoint.job_id - step_offset, _ = setup_resume(client, cfg.resume) + state = resolve_resume(cfg.log_path, cfg.init_from_checkpoint) + dcp_load_time = load_dcp(client, state) + step_offset = state.step adam_params = tinker.AdamParams(learning_rate=cfg.learning_rate, **DEFAULT_ADAM) # -- Data ---------------------------------------------------------------- diff --git a/training/recipes/rl_loop.py b/training/recipes/rl_loop.py index 59a43f07..e3c0f40f 100644 --- a/training/recipes/rl_loop.py +++ b/training/recipes/rl_loop.py @@ -34,12 +34,10 @@ InfraConfig, WandBConfig, DeployConfig, - ResumeConfig, HotloadConfig, ReconnectableClient, wandb_log, setup_wandb, - setup_resume, wandb_finish, validate_config, log_metrics_json, @@ -49,6 +47,13 @@ load_jsonl_dataset, prepare_sampling_messages, ) +from training.utils.checkpoint_utils import ( + resolve_resume, + load_dcp, + save_loop_state, + dataset_fingerprint, + validate_dataset, +) from fireworks.training.sdk.deployment import DeploymentSampler from training.utils.rl import PromptGroup from training.utils.rl.importance_sampling import ISConfig @@ -126,11 +131,15 @@ class Config: reference_base_url: str | None = None """Base URL for the reference trainer (bypass direct route).""" + log_path: str = "./rl_logs" + init_from_checkpoint: str | None = None + """Load pretrained DCP weights on a fresh dataset. Supports cross-job + format ``"job_id:checkpoint_name"``.""" + infra: InfraConfig = field(default_factory=InfraConfig) deployment: DeployConfig = field(default_factory=DeployConfig) hotload: HotloadConfig = field(default_factory=HotloadConfig) wandb: WandBConfig = field(default_factory=lambda: WandBConfig(project="grpo-tinker")) - resume: ResumeConfig = field(default_factory=ResumeConfig) # --------------------------------------------------------------------------- @@ -222,7 +231,7 @@ def _signal_handler(signum, frame): signal.signal(signal.SIGTERM, _signal_handler) signal.signal(signal.SIGINT, _signal_handler) - validate_config(cfg.base_model, cfg.dataset, cfg.hotload, cfg.deployment, cfg.infra, cfg.resume) + validate_config(cfg.base_model, cfg.dataset, cfg.hotload, cfg.deployment, cfg.infra) completions_per_prompt = cfg.completions_per_prompt prompt_groups_per_step = cfg.prompt_groups_per_step if not cfg.deployment.tokenizer_model: @@ -377,7 +386,14 @@ def _signal_handler(signum, frame): boot_metrics["infra/deploy_boot_time"] = deploy_mgr.boot_time_s wandb_log(boot_metrics, step=0) - step_offset, _ = setup_resume(policy, cfg.resume) + # -- Resume --------------------------------------------------------------- + + state = resolve_resume(cfg.log_path, cfg.init_from_checkpoint) + dcp_load_time = load_dcp(policy, state) + if dcp_load_time > 0: + wandb_log({"perf/dcp_load_time": dcp_load_time}, state.step) + step_offset = state.step + if cfg.hotload.hot_load_before_training and cfg.deployment.deployment_id: name = f"resume-{step_offset}-base" if step_offset > 0 else "step-0-base" weight_syncer.save_and_hotload(name, checkpoint_type="base") @@ -385,6 +401,8 @@ def _signal_handler(signum, frame): # -- Prepare sampling and training -------------------------------------- dataset = load_jsonl_dataset(cfg.dataset, cfg.max_rows) + fp = dataset_fingerprint(dataset) + validate_dataset(state.dataset_fingerprint, fp, state.data_consumed) adam_params = tinker.AdamParams(learning_rate=cfg.learning_rate, **DEFAULT_ADAM) loss_builder = build_loss_fn( policy_loss=cfg.policy_loss, kl_beta=cfg.kl_beta, @@ -486,7 +504,7 @@ async def sample_one_prompt(row: dict) -> PromptGroup | None: return PromptGroup( data=policy_data, ref_data=reference_data, - advantages=adv_filtered, ref_logprobs=[], + advantages=adv_filtered, ref_logprobs=None, prompt_len=prompt_len, rewards=rewards, inf_logprobs=inf_logprobs_aligned, completion_lens=comp_lens, truncated=trunc, prompt=input_messages if cfg.trajectory_dir else None, @@ -557,6 +575,15 @@ def finish_step( with timer("dcp_save"): weight_syncer.save_dcp(f"step-{step}") logger.info("[step %d] dcp_save: done (%.1fs)", step, _t.time() - t0) + data_consumed = state.data_consumed + (step - state.step) * prompt_groups_per_step + save_loop_state(cfg.log_path, { + "step": step, + "data_consumed": data_consumed, + "dcp_name": f"step-{step}", + "dataset_fingerprint": fp, + "training_shape_id": getattr(cfg.infra, "training_shape_id", None), + "source_job_id": policy_job_id, + }) metrics = compute_step_metrics( prompt_groups=prompt_groups, @@ -596,7 +623,7 @@ def _loop_metrics_callback(loop_metrics: dict) -> None: finish_step=finish_step, ) - all_rows = dataset * cfg.epochs + all_rows = (dataset * cfg.epochs)[state.data_consumed:] global_step = asyncio.run(run_rl_loop( sample_fns=(sample_one_prompt(row) for row in all_rows), @@ -614,6 +641,15 @@ def _loop_metrics_callback(loop_metrics: dict) -> None: if global_step > step_offset: try: policy.save_state(f"step-{global_step}", timeout=cfg.hotload.dcp_timeout) + data_consumed = state.data_consumed + (global_step - state.step) * prompt_groups_per_step + save_loop_state(cfg.log_path, { + "step": global_step, + "data_consumed": data_consumed, + "dcp_name": f"step-{global_step}", + "dataset_fingerprint": fp, + "training_shape_id": getattr(cfg.infra, "training_shape_id", None), + "source_job_id": policy_job_id, + }) except Exception as e: logger.warning("Failed to save final checkpoint: %s", e) diff --git a/training/recipes/sft_loop.py b/training/recipes/sft_loop.py index b04e7bc8..4621661a 100644 --- a/training/recipes/sft_loop.py +++ b/training/recipes/sft_loop.py @@ -33,11 +33,9 @@ DEFAULT_ADAM, InfraConfig, WandBConfig, - ResumeConfig, ReconnectableClient, wandb_log, setup_wandb, - setup_resume, wandb_finish, validate_config, log_metrics_json, @@ -48,6 +46,13 @@ render_messages_to_datum, resolve_renderer_name, ) +from training.utils.checkpoint_utils import ( + resolve_resume, + load_dcp, + save_loop_state, + dataset_fingerprint, + validate_dataset, +) from training.utils.timer import timer, flush_timing @@ -77,9 +82,13 @@ class Config: dcp_save_interval: int = 0 # save DCP checkpoint every N steps (0 = off) + log_path: str = "./sft_logs" + init_from_checkpoint: str | None = None + """Load pretrained DCP weights on a fresh dataset. Supports cross-job + format ``"job_id:checkpoint_name"``.""" + infra: InfraConfig = field(default_factory=InfraConfig) wandb: WandBConfig = field(default_factory=lambda: WandBConfig(project="sft-tinker")) - resume: ResumeConfig = field(default_factory=ResumeConfig) # --------------------------------------------------------------------------- @@ -101,7 +110,7 @@ def _signal_handler(signum, frame): signal.signal(signal.SIGTERM, _signal_handler) signal.signal(signal.SIGINT, _signal_handler) - validate_config(cfg.base_model, cfg.dataset, infra=cfg.infra, resume=cfg.resume) + validate_config(cfg.base_model, cfg.dataset, infra=cfg.infra) setup_wandb( cfg.wandb, { @@ -204,14 +213,27 @@ def _signal_handler(signum, frame): if not training_data: raise RuntimeError("No valid training examples after tokenization") - step_offset, _ = setup_resume(client, cfg.resume) + # -- Resume --------------------------------------------------------------- + + state = resolve_resume(cfg.log_path, cfg.init_from_checkpoint) + dcp_load_time = load_dcp(client, state) + if dcp_load_time > 0: + wandb_log({"perf/dcp_load_time": dcp_load_time}, state.step) + + fp = dataset_fingerprint(raw_data) + validate_dataset(state.dataset_fingerprint, fp, state.data_consumed) + adam_params = tinker.AdamParams(learning_rate=cfg.learning_rate, **DEFAULT_ADAM) # -- Training loop (batched) ------------------------------------------- batch_size = cfg.batch_size - step = step_offset - total_steps = len(training_data) * cfg.epochs // (cfg.grad_accum * batch_size) + all_examples = training_data * cfg.epochs + examples_to_process = all_examples[state.data_consumed:] + data_consumed = state.data_consumed + + step = state.step + total_steps = len(all_examples) // (cfg.grad_accum * batch_size) accum = 0 agg_loss_sum = 0.0 agg_resp_tokens = 0 @@ -239,6 +261,14 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ with timer("dcp_save"): logger.info("Saving DCP checkpoint at step %d", step) client.inner.save_state(f"step-{step}") + save_loop_state(cfg.log_path, { + "step": step, + "data_consumed": data_consumed, + "dcp_name": f"step-{step}", + "dataset_fingerprint": fp, + "training_shape_id": getattr(cfg.infra, "training_shape_id", None), + "source_job_id": job_id, + }) step_metrics: Dict[str, Any] = flush_timing() @@ -269,16 +299,16 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ return step, accum - for _epoch in range(cfg.epochs): - batch_buffer: list[tinker.Datum] = [] - for ex in training_data: - batch_buffer.append(ex) - if len(batch_buffer) >= batch_size: - step, accum = _flush_batch(batch_buffer, step, accum) - batch_buffer = [] - - if batch_buffer: + batch_buffer: list[tinker.Datum] = [] + for ex in examples_to_process: + batch_buffer.append(ex) + data_consumed += 1 + if len(batch_buffer) >= batch_size: step, accum = _flush_batch(batch_buffer, step, accum) + batch_buffer = [] + + if batch_buffer: + step, accum = _flush_batch(batch_buffer, step, accum) if accum > 0: client.optim_step(adam_params) @@ -286,9 +316,17 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ # -- Final checkpoint -------------------------------------------------- - if step > step_offset: + if step > state.step: logger.info("Saving final DCP checkpoint (step %d)...", step) - client.inner.save_state(f"final-step-{step}") + client.inner.save_state(f"step-{step}") + save_loop_state(cfg.log_path, { + "step": step, + "data_consumed": data_consumed, + "dcp_name": f"step-{step}", + "dataset_fingerprint": fp, + "training_shape_id": getattr(cfg.infra, "training_shape_id", None), + "source_job_id": job_id, + }) logger.info("Saving final base checkpoint (step %d)...", step) result = client.inner.save_weights_for_sampler_ext( diff --git a/training/tests/e2e/test_dpo_resume_e2e.py b/training/tests/e2e/test_dpo_resume_e2e.py index 948fd211..66309ae8 100644 --- a/training/tests/e2e/test_dpo_resume_e2e.py +++ b/training/tests/e2e/test_dpo_resume_e2e.py @@ -17,7 +17,7 @@ import pytest -from training.utils import InfraConfig, DeployConfig, ResumeConfig, HotloadConfig +from training.utils import InfraConfig, DeployConfig, HotloadConfig from training.recipes.dpo_loop import Config, main logger = logging.getLogger(__name__) @@ -112,7 +112,7 @@ def test_dpo_resume_from_checkpoint( infra=shared_infra, deployment=DeployConfig(), hotload=HotloadConfig(hot_load_interval=0), - resume=ResumeConfig(resume_from=dcp_name, resume_job_id=phase1_job_id), + init_from_checkpoint=f"{phase1_job_id}:{dcp_name}", ) phase2_metrics = main(phase2_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) diff --git a/training/tests/e2e/test_grpo_resume_e2e.py b/training/tests/e2e/test_grpo_resume_e2e.py index 0e14e93f..0b4b5cf7 100644 --- a/training/tests/e2e/test_grpo_resume_e2e.py +++ b/training/tests/e2e/test_grpo_resume_e2e.py @@ -21,7 +21,7 @@ import pytest -from training.utils import InfraConfig, DeployConfig, ResumeConfig, HotloadConfig +from training.utils import InfraConfig, DeployConfig, HotloadConfig from training.utils.rl import ISConfig from training.tests.e2e.conftest import GSM8K_SAMPLE_URL from training.recipes.rl_loop import Config, main @@ -132,10 +132,7 @@ def test_grpo_resume_from_checkpoint( hot_load_before_training=True, hot_load_timeout=600, ), - resume=ResumeConfig( - resume_from=dcp_name, - resume_job_id=phase1_policy_job_id, - ), + init_from_checkpoint=f"{phase1_policy_job_id}:{dcp_name}", ) phase2_metrics = main(phase2_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) diff --git a/training/tests/e2e/test_sft_resume_e2e.py b/training/tests/e2e/test_sft_resume_e2e.py index 080080ad..3e466667 100644 --- a/training/tests/e2e/test_sft_resume_e2e.py +++ b/training/tests/e2e/test_sft_resume_e2e.py @@ -17,7 +17,7 @@ import pytest -from training.utils import InfraConfig, DeployConfig, ResumeConfig, HotloadConfig +from training.utils import InfraConfig, DeployConfig, HotloadConfig from training.recipes.sft_loop import Config, main logger = logging.getLogger(__name__) @@ -115,7 +115,7 @@ def test_sft_resume_from_checkpoint( infra=shared_infra, deployment=DeployConfig(), hotload=HotloadConfig(hot_load_interval=0), - resume=ResumeConfig(resume_from=dcp_name, resume_job_id=phase1_job_id), + init_from_checkpoint=f"{phase1_job_id}:{dcp_name}", ) phase2_metrics = main(phase2_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) diff --git a/training/tests/unit/test_rl_loop.py b/training/tests/unit/test_rl_loop.py index 408ae1a8..0704d9ed 100644 --- a/training/tests/unit/test_rl_loop.py +++ b/training/tests/unit/test_rl_loop.py @@ -380,7 +380,9 @@ def _builder(adv, ref_lp, prompt_lens, inf_lp, prox_lp): monkeypatch.setattr(transformers.AutoTokenizer, "from_pretrained", lambda *args, **kwargs: object()) monkeypatch.setattr(module, "DeploymentSampler", FakeSampler) monkeypatch.setattr(module, "WeightSyncer", FakeWeightSyncer) - monkeypatch.setattr(module, "setup_resume", lambda *args, **kwargs: (1, None)) + from training.utils.checkpoint_utils import ResumeState + monkeypatch.setattr(module, "resolve_resume", lambda *args, **kwargs: ResumeState(step=1)) + monkeypatch.setattr(module, "load_dcp", lambda *args, **kwargs: 0.0) monkeypatch.setattr(module, "load_jsonl_dataset", lambda *args, **kwargs: [ { "messages": [ diff --git a/training/utils/__init__.py b/training/utils/__init__.py index a59595d4..cc62cd46 100644 --- a/training/utils/__init__.py +++ b/training/utils/__init__.py @@ -12,7 +12,6 @@ "HotloadConfig", "InfraConfig", "ReconnectableClient", - "ResumeConfig", "RewardFn", "StepCallback", "WandBConfig", @@ -43,7 +42,6 @@ "resolve_renderer_name", "prepare_sampling_messages", "setup_deployment", - "setup_resume", "setup_training_client", "setup_wandb", "flush_timing", @@ -78,7 +76,6 @@ InfraConfig, WandBConfig, DeployConfig, - ResumeConfig, StepCallback, HotloadConfig, ) @@ -102,7 +99,6 @@ render_messages_to_datum, resolve_renderer_name, ) -from training.utils.resume import setup_resume from training.utils.logging import ( wandb_log, setup_wandb, diff --git a/training/utils/checkpoint_utils.py b/training/utils/checkpoint_utils.py new file mode 100644 index 00000000..e34a9f6a --- /dev/null +++ b/training/utils/checkpoint_utils.py @@ -0,0 +1,214 @@ +"""Checkpoint utilities -- single source of truth for resume state. + +All resume logic flows through ``local_checkpoint_state.jsonl``: + +- **Continuing a run**: last entry has ``dcp_name``, ``step``, ``data_consumed``. +- **Fresh start with pretrained weights**: ``init_from_checkpoint`` writes an + initial entry with ``step=0, data_consumed=0``, then loads the DCP. +- **Completely fresh start**: no entries, no DCP load. +""" + +from __future__ import annotations + +import hashlib +import json +import logging +import os +import time +from dataclasses import dataclass +from typing import Any + +logger = logging.getLogger(__name__) + +STATE_FILE = "local_checkpoint_state.jsonl" + + +# -- Resume state -------------------------------------------------------------- + + +@dataclass +class ResumeState: + """Resolved resume state -- the single source of truth.""" + + step: int = 0 + data_consumed: int = 0 + dcp_name: str | None = None + dataset_fingerprint: str | None = None + training_shape_id: str | None = None + source_job_id: str | None = None + + +def resolve_resume( + log_path: str, + init_from_checkpoint: str | None = None, +) -> ResumeState: + """Determine resume state from ``local_checkpoint_state.jsonl``. + + Priority: + 1. If ``local_checkpoint_state.jsonl`` has entries, resume from the last one. + 2. If empty/missing but *init_from_checkpoint* is set, start fresh + with DCP weights (step=0, data_consumed=0). + 3. Otherwise, completely fresh start (no DCP). + + *init_from_checkpoint* supports cross-job format ``"job_id:checkpoint_name"``. + """ + last = _load_last_loop_state(log_path) + if last is not None: + logger.info("Resuming from local_checkpoint_state.jsonl: %s", last) + return ResumeState( + step=last.get("step", 0), + data_consumed=last.get("data_consumed", 0), + dcp_name=last.get("dcp_name"), + dataset_fingerprint=last.get("dataset_fingerprint"), + training_shape_id=last.get("training_shape_id"), + source_job_id=last.get("source_job_id"), + ) + + if init_from_checkpoint: + source_job_id = None + dcp_name = init_from_checkpoint + if ":" in init_from_checkpoint and not init_from_checkpoint.startswith(("gs://", "/")): + source_job_id, dcp_name = init_from_checkpoint.split(":", 1) + + logger.info( + "Fresh start with pretrained weights: dcp=%s source_job=%s", + dcp_name, + source_job_id, + ) + initial = ResumeState( + step=0, + dcp_name=dcp_name, + source_job_id=source_job_id, + ) + save_loop_state(log_path, _state_to_dict(initial)) + return initial + + logger.info("Fresh start (no checkpoint)") + return ResumeState() + + +def load_dcp(client: Any, state: ResumeState) -> float: + """Load DCP checkpoint if *state.dcp_name* is set. + + Resolves cross-job references and loads model weights + optimizer. + Returns the load time in seconds (0.0 if no checkpoint was loaded). + """ + if state.dcp_name is None: + return 0.0 + + checkpoint_ref = client.resolve_checkpoint_path( + state.dcp_name, + source_job_id=state.source_job_id, + ) + logger.info("Loading DCP checkpoint: %s", checkpoint_ref) + t0 = time.time() + client.load_state_with_optimizer(checkpoint_ref) + elapsed = time.time() - t0 + logger.info("DCP checkpoint loaded: %s (%.1fs)", state.dcp_name, elapsed) + return elapsed + + +# -- Loop state persistence ---------------------------------------------------- + + +def save_loop_state(log_path: str, loop_state: dict[str, Any]) -> None: + """Append a loop-state entry to ``local_checkpoint_state.jsonl``.""" + os.makedirs(log_path, exist_ok=True) + path = os.path.join(log_path, STATE_FILE) + with open(path, "a") as f: + f.write(json.dumps(loop_state) + "\n") + logger.info("Saved loop state: %s", loop_state) + + +def _load_last_loop_state(log_path: str) -> dict[str, Any] | None: + """Read the most recent entry from ``local_checkpoint_state.jsonl``.""" + path = os.path.join(log_path, STATE_FILE) + if not os.path.exists(path): + return None + with open(path) as f: + lines = [line.strip() for line in f if line.strip()] + if not lines: + return None + return json.loads(lines[-1]) + + +def _state_to_dict(state: ResumeState) -> dict[str, Any]: + """Serialize a ResumeState to a dict for local_checkpoint_state.jsonl.""" + return { + "step": state.step, + "data_consumed": state.data_consumed, + "dcp_name": state.dcp_name, + "dataset_fingerprint": state.dataset_fingerprint, + "training_shape_id": state.training_shape_id, + "source_job_id": state.source_job_id, + } + + +# -- Dataset fingerprint ------------------------------------------------------- + + +def dataset_fingerprint(rows: list[dict]) -> str: + """Short hash of row count + first/last row content.""" + if not rows: + return "empty" + content = ( + f"{len(rows)}:" + f"{json.dumps(rows[0], sort_keys=True)}:" + f"{json.dumps(rows[-1], sort_keys=True)}" + ) + return hashlib.sha256(content.encode()).hexdigest()[:12] + + +def validate_dataset( + saved_fingerprint: str | None, + current_fingerprint: str, + data_consumed: int, +) -> None: + """Warn if the dataset changed between checkpoint save and resume.""" + if saved_fingerprint and saved_fingerprint != current_fingerprint: + logger.warning( + "Dataset changed since checkpoint! " + "fingerprint: saved=%s current=%s. " + "data_consumed=%d may point to different data.", + saved_fingerprint, + current_fingerprint, + data_consumed, + ) + + +# -- Checkpoint availability --------------------------------------------------- + + +def verify_checkpoint_available(client: Any, dcp_name: str) -> bool: + """Check that a DCP checkpoint exists on the trainer.""" + try: + checkpoints, _ = client.list_checkpoints() + if dcp_name in checkpoints: + logger.info( + "Checkpoint '%s' available (all: %s)", dcp_name, checkpoints, + ) + return True + logger.warning( + "Checkpoint '%s' not found. Available: %s", dcp_name, checkpoints, + ) + return False + except Exception as e: + logger.warning("Could not list checkpoints: %s. Proceeding anyway.", e) + return True + + +# -- Training shape validation ------------------------------------------------- + + +def validate_training_shape( + saved_shape_id: str | None, + current_shape_id: str | None, +) -> None: + """Warn if the training shape changed between checkpoint save and resume.""" + if saved_shape_id and current_shape_id and saved_shape_id != current_shape_id: + logger.warning( + "Training shape changed! saved=%s current=%s. " + "Model parallelism may be incompatible.", + saved_shape_id, + current_shape_id, + ) diff --git a/training/utils/client.py b/training/utils/client.py index 8109e057..5999cf5a 100644 --- a/training/utils/client.py +++ b/training/utils/client.py @@ -115,6 +115,9 @@ def load_state_with_optimizer(self, path: str, timeout: int = DCP_TIMEOUT_S): def resolve_checkpoint_path(self, name: str, source_job_id: str | None = None) -> str: return self.inner.resolve_checkpoint_path(name, source_job_id=source_job_id) + def list_checkpoints(self) -> tuple[list[str], str | None]: + return self.inner.list_checkpoints() + # -- Internal -------------------------------------------------------------- def _use_endpoint(self, ep: TrainerServiceEndpoint) -> None: diff --git a/training/utils/config.py b/training/utils/config.py index bb763e32..1dba89ad 100644 --- a/training/utils/config.py +++ b/training/utils/config.py @@ -124,10 +124,3 @@ class WandBConfig: run_name: str | None = None -@dataclass -class ResumeConfig: - """Checkpoint resume settings.""" - - resume_from: str | None = None - resume_job_id: str | None = None - step_offset: int | None = None diff --git a/training/utils/resume.py b/training/utils/resume.py deleted file mode 100644 index f5cd88bc..00000000 --- a/training/utils/resume.py +++ /dev/null @@ -1,49 +0,0 @@ -"""Checkpoint resume helpers.""" - -from __future__ import annotations - -import re -import logging - -from training.utils.client import ReconnectableClient -from training.utils.config import ResumeConfig - -logger = logging.getLogger(__name__) - - -def setup_resume( - client: ReconnectableClient, - resume: ResumeConfig, -) -> tuple[int, str | None]: - """Load a checkpoint and return (step_offset, checkpoint_name).""" - if not resume.resume_from: - return 0, None - - checkpoint_ref = client.resolve_checkpoint_path( - resume.resume_from, - source_job_id=resume.resume_job_id, - ) - logger.info("Loading checkpoint: %s", checkpoint_ref) - client.load_state_with_optimizer(checkpoint_ref).result(timeout=1800) - logger.info("Checkpoint loaded: %s", resume.resume_from) - - if resume.step_offset is not None: - logger.info("Step offset (explicit): %d", resume.step_offset) - return resume.step_offset, resume.resume_from - - step_offset = 0 - step_match = re.search(r"step-(\d+)", resume.resume_from) - if step_match: - step_offset = int(step_match.group(1)) - logger.warning( - "Inferred step_offset=%d from checkpoint name '%s'. " - "Set ResumeConfig.step_offset explicitly to avoid this heuristic.", - step_offset, - resume.resume_from, - ) - else: - logger.warning( - "Could not infer step offset from '%s'. Starting from step 0.", - resume.resume_from, - ) - return step_offset, resume.resume_from diff --git a/training/utils/rl/cispo.py b/training/utils/rl/cispo.py index bcd4dc40..5ec108d3 100644 --- a/training/utils/rl/cispo.py +++ b/training/utils/rl/cispo.py @@ -81,7 +81,7 @@ def loss_fn( for i, pi_logprobs in enumerate(logprobs_list): adv = advantages[i] - ref_lp = ref_logprobs[i] + ref_lp = ref_logprobs[i] if ref_logprobs else [] inf_lp = inf_logprobs[i] prox_lp = prox_logprobs[i] response_start = max(0, prompt_lens[i] - 1) diff --git a/training/utils/rl/dapo.py b/training/utils/rl/dapo.py index 91a912c9..6459073b 100644 --- a/training/utils/rl/dapo.py +++ b/training/utils/rl/dapo.py @@ -72,7 +72,7 @@ def loss_fn( for i, pi_logprobs in enumerate(logprobs_list): adv = advantages[i] - ref_lp = ref_logprobs[i] + ref_lp = ref_logprobs[i] if ref_logprobs else [] inf_lp = inf_logprobs[i] prox_lp = prox_logprobs[i] response_start = max(0, prompt_lens[i] - 1) diff --git a/training/utils/rl/grpo.py b/training/utils/rl/grpo.py index 07a41756..327ce318 100644 --- a/training/utils/rl/grpo.py +++ b/training/utils/rl/grpo.py @@ -61,7 +61,7 @@ def loss_fn( for i, pi_logprobs in enumerate(logprobs_list): adv = advantages[i] - ref_lp = ref_logprobs[i] + ref_lp = ref_logprobs[i] if ref_logprobs else [] inf_lp = inf_logprobs[i] prox_lp = prox_logprobs[i] response_start = max(0, prompt_lens[i] - 1) diff --git a/training/utils/rl/gspo.py b/training/utils/rl/gspo.py index 51118fa1..334979d3 100644 --- a/training/utils/rl/gspo.py +++ b/training/utils/rl/gspo.py @@ -73,7 +73,7 @@ def loss_fn( for i, pi_logprobs in enumerate(logprobs_list): adv = advantages[i] - ref_lp = ref_logprobs[i] + ref_lp = ref_logprobs[i] if ref_logprobs else [] inf_lp = inf_logprobs[i] prox_lp = prox_logprobs[i] response_start = max(0, prompt_lens[i] - 1) diff --git a/training/utils/rl/losses.py b/training/utils/rl/losses.py index 9c329525..0f13265c 100644 --- a/training/utils/rl/losses.py +++ b/training/utils/rl/losses.py @@ -16,7 +16,7 @@ class PromptGroup: data: List[tinker.Datum] advantages: List[float] - ref_logprobs: List[List[float]] + ref_logprobs: List[List[float]] | None prompt_len: int rewards: List[float] ref_data: List[tinker.Datum] = field(default_factory=list) @@ -50,7 +50,8 @@ def combine_prompt_groups( for pg in groups: data.extend(pg.data) advantages.extend(pg.advantages) - ref_logprobs.extend(pg.ref_logprobs) + if pg.ref_logprobs is not None: + ref_logprobs.extend(pg.ref_logprobs) prompt_lens.extend([pg.prompt_len] * len(pg.data)) inf_logprobs.extend(pg.inf_logprobs) diff --git a/training/utils/validation.py b/training/utils/validation.py index c77a0190..45dd588e 100644 --- a/training/utils/validation.py +++ b/training/utils/validation.py @@ -5,7 +5,7 @@ import logging from fireworks.training.sdk.errors import format_sdk_error, DOCS_SDK -from training.utils.config import InfraConfig, DeployConfig, ResumeConfig, HotloadConfig +from training.utils.config import InfraConfig, DeployConfig, HotloadConfig logger = logging.getLogger(__name__) @@ -16,7 +16,6 @@ def validate_config( hotload: HotloadConfig | None = None, deploy: DeployConfig | None = None, infra: InfraConfig | None = None, - resume: ResumeConfig | None = None, ) -> None: """Pre-flight validation. Catches misconfiguration before provisioning GPUs.""" errors: list[str] = [] @@ -48,18 +47,6 @@ def validate_config( ) ) - if ( - resume - and resume.resume_from - and not resume.resume_from.startswith(("gs://", "/")) - ): - if resume.resume_job_id is None: - logger.warning( - "resume_from='%s' looks like a checkpoint name, not a full path. " - "If resuming from a different job, set resume_job_id.", - resume.resume_from, - ) - if infra and infra.node_count is not None and infra.node_count < 1: errors.append( format_sdk_error( From 1e9c9a212c2a8234ceae180403ac605462d84d75 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 18:03:05 -0700 Subject: [PATCH 02/28] chore: bump fireworks-ai to >=1.0.0a40 Made-with: Cursor --- training/pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/training/pyproject.toml b/training/pyproject.toml index 7adc0a3e..03095c37 100644 --- a/training/pyproject.toml +++ b/training/pyproject.toml @@ -7,7 +7,7 @@ name = "fireworks-training-cookbook" version = "0.1.0" requires-python = ">=3.11" dependencies = [ - "fireworks-ai>=1.0.0a39,<2", + "fireworks-ai>=1.0.0a40,<2", "tinker-cookbook @ git+https://github.com/thinking-machines-lab/tinker-cookbook.git@934f0d9b2f53c3edff02cbf23ec6da8682047fa5", "eval-protocol>=0.3.23", "tqdm", From 8453b077aac0e9cb7598d3006c9ce39909c2037f Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 18:05:28 -0700 Subject: [PATCH 03/28] fix: align SFT resume test with current Config API Made-with: Cursor --- training/tests/e2e/test_sft_resume_e2e.py | 11 ++++------- 1 file changed, 4 insertions(+), 7 deletions(-) diff --git a/training/tests/e2e/test_sft_resume_e2e.py b/training/tests/e2e/test_sft_resume_e2e.py index 3e466667..b1bcf968 100644 --- a/training/tests/e2e/test_sft_resume_e2e.py +++ b/training/tests/e2e/test_sft_resume_e2e.py @@ -17,7 +17,7 @@ import pytest -from training.utils import InfraConfig, DeployConfig, HotloadConfig +from training.utils import InfraConfig from training.recipes.sft_loop import Config, main logger = logging.getLogger(__name__) @@ -85,12 +85,11 @@ def test_sft_resume_from_checkpoint( epochs=2, grad_accum=2, max_examples=10, + dcp_save_interval=4, infra=shared_infra, - deployment=DeployConfig(), - hotload=HotloadConfig(hot_load_interval=0, dcp_save_interval=4), ) - phase1_metrics = main(phase1_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) + phase1_metrics = main(phase1_config, rlor_mgr=rlor_mgr) assert isinstance(phase1_metrics, dict) assert "steps" in phase1_metrics @@ -113,12 +112,10 @@ def test_sft_resume_from_checkpoint( grad_accum=2, max_examples=10, infra=shared_infra, - deployment=DeployConfig(), - hotload=HotloadConfig(hot_load_interval=0), init_from_checkpoint=f"{phase1_job_id}:{dcp_name}", ) - phase2_metrics = main(phase2_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) + phase2_metrics = main(phase2_config, rlor_mgr=rlor_mgr) assert isinstance(phase2_metrics, dict) assert "steps" in phase2_metrics From c5797ad5ee01eb8db963cc48197776910d5f0380 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 18:11:13 -0700 Subject: [PATCH 04/28] fix: add max_seq_len to SFT resume test Made-with: Cursor --- training/tests/e2e/test_sft_resume_e2e.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/training/tests/e2e/test_sft_resume_e2e.py b/training/tests/e2e/test_sft_resume_e2e.py index b1bcf968..70e8ae46 100644 --- a/training/tests/e2e/test_sft_resume_e2e.py +++ b/training/tests/e2e/test_sft_resume_e2e.py @@ -84,6 +84,7 @@ def test_sft_resume_from_checkpoint( learning_rate=1e-4, epochs=2, grad_accum=2, + max_seq_len=4096, max_examples=10, dcp_save_interval=4, infra=shared_infra, @@ -110,6 +111,7 @@ def test_sft_resume_from_checkpoint( learning_rate=1e-4, epochs=2, grad_accum=2, + max_seq_len=4096, max_examples=10, infra=shared_infra, init_from_checkpoint=f"{phase1_job_id}:{dcp_name}", From a89d081fff4b45f401d230c26dc02cdcfe3a9dcc Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 18:16:55 -0700 Subject: [PATCH 05/28] fix: remove skip_validations from SFT resume test (shared dev key) Made-with: Cursor --- training/tests/e2e/test_sft_resume_e2e.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/training/tests/e2e/test_sft_resume_e2e.py b/training/tests/e2e/test_sft_resume_e2e.py index 70e8ae46..c5cb5e3c 100644 --- a/training/tests/e2e/test_sft_resume_e2e.py +++ b/training/tests/e2e/test_sft_resume_e2e.py @@ -69,9 +69,7 @@ def test_sft_resume_from_checkpoint( shared_infra = InfraConfig( region=e2e_region, - skip_validations=True, - accelerator_type=e2e_training_accelerator, - custom_image_tag=custom_image_tag, + custom_image_tag=custom_image_tag or "0.65.4", ) # Phase 1: train, save DCP From fa65b07ee958d52894cfcd718b42d6fb87ae6517 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 18:19:54 -0700 Subject: [PATCH 06/28] fix: use training shape for SFT resume test Made-with: Cursor --- training/tests/e2e/test_sft_resume_e2e.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/training/tests/e2e/test_sft_resume_e2e.py b/training/tests/e2e/test_sft_resume_e2e.py index c5cb5e3c..d37061e7 100644 --- a/training/tests/e2e/test_sft_resume_e2e.py +++ b/training/tests/e2e/test_sft_resume_e2e.py @@ -69,7 +69,8 @@ def test_sft_resume_from_checkpoint( shared_infra = InfraConfig( region=e2e_region, - custom_image_tag=custom_image_tag or "0.65.4", + training_shape_id="ts-qwen3-30b-a3b-policy", + custom_image_tag=custom_image_tag, ) # Phase 1: train, save DCP From 6ebe9c231540e34db76ef1a907d133e14a3462a7 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 18:28:27 -0700 Subject: [PATCH 07/28] fix: pass max_context_length in validated shape path Made-with: Cursor --- training/utils/infra.py | 1 + 1 file changed, 1 insertion(+) diff --git a/training/utils/infra.py b/training/utils/infra.py index 8ef26749..a281f7da 100644 --- a/training/utils/infra.py +++ b/training/utils/infra.py @@ -98,6 +98,7 @@ def create_trainer_job( config = TrainerJobConfig( base_model=base_model, lora_rank=lora_rank, + max_context_length=max_seq_len or profile.max_supported_context_length, learning_rate=learning_rate, gradient_accumulation_steps=grad_accum, display_name=display_name, From af3203fcc522c30514145629d3970a22920963f5 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 18:30:09 -0700 Subject: [PATCH 08/28] fix: use 2-node shape for SFT resume test (has matching trainer image) Made-with: Cursor --- training/tests/e2e/test_sft_resume_e2e.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/training/tests/e2e/test_sft_resume_e2e.py b/training/tests/e2e/test_sft_resume_e2e.py index d37061e7..7567e5b5 100644 --- a/training/tests/e2e/test_sft_resume_e2e.py +++ b/training/tests/e2e/test_sft_resume_e2e.py @@ -69,7 +69,7 @@ def test_sft_resume_from_checkpoint( shared_infra = InfraConfig( region=e2e_region, - training_shape_id="ts-qwen3-30b-a3b-policy", + training_shape_id="ts-qwen3-30b-a3b-128k-2node", custom_image_tag=custom_image_tag, ) From cec58e61d3e86f2e59c3cb3b0ea8fd4936acfc7c Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 18:32:59 -0700 Subject: [PATCH 09/28] fix: use custom_image_tag without shape to avoid server context length override Made-with: Cursor --- training/tests/e2e/test_sft_resume_e2e.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/training/tests/e2e/test_sft_resume_e2e.py b/training/tests/e2e/test_sft_resume_e2e.py index 7567e5b5..1152313c 100644 --- a/training/tests/e2e/test_sft_resume_e2e.py +++ b/training/tests/e2e/test_sft_resume_e2e.py @@ -69,8 +69,7 @@ def test_sft_resume_from_checkpoint( shared_infra = InfraConfig( region=e2e_region, - training_shape_id="ts-qwen3-30b-a3b-128k-2node", - custom_image_tag=custom_image_tag, + custom_image_tag=custom_image_tag or "0.33.0", ) # Phase 1: train, save DCP From 5ce83319ce323fca3d17f78de55f847fcf80a5d5 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 18:45:08 -0700 Subject: [PATCH 10/28] fix: pass fw_api_key to ReconnectableClient for gateway auth Made-with: Cursor --- training/recipes/sft_loop.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/training/recipes/sft_loop.py b/training/recipes/sft_loop.py index 4621661a..431d420e 100644 --- a/training/recipes/sft_loop.py +++ b/training/recipes/sft_loop.py @@ -160,7 +160,7 @@ def _signal_handler(signum, frame): display_name="sft-trainer", ) job_id = endpoint.job_id - client = ReconnectableClient(rlor_mgr, job_id, cfg.base_model, cfg.lora_rank) + client = ReconnectableClient(rlor_mgr, job_id, cfg.base_model, cfg.lora_rank, fw_api_key=api_key) # -- Prepare data ------------------------------------------------------ try: From a0d8cf3d160ce1d1f85c9161910d6d08cb4f492b Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 19:36:25 -0700 Subject: [PATCH 11/28] fix: increase examples and reduce batch_size in SFT resume test for >= 2 steps Made-with: Cursor --- sft_logs/local_checkpoint_state.jsonl | 1 + training/tests/e2e/test_sft_resume_e2e.py | 8 +++++--- 2 files changed, 6 insertions(+), 3 deletions(-) create mode 100644 sft_logs/local_checkpoint_state.jsonl diff --git a/sft_logs/local_checkpoint_state.jsonl b/sft_logs/local_checkpoint_state.jsonl new file mode 100644 index 00000000..67e1a5c5 --- /dev/null +++ b/sft_logs/local_checkpoint_state.jsonl @@ -0,0 +1 @@ +{"step": 1, "data_consumed": 20, "dcp_name": "step-1", "dataset_fingerprint": "e12293449e59", "training_shape_id": null, "source_job_id": "bgyhqcmj03tee38w"} diff --git a/training/tests/e2e/test_sft_resume_e2e.py b/training/tests/e2e/test_sft_resume_e2e.py index 1152313c..91134282 100644 --- a/training/tests/e2e/test_sft_resume_e2e.py +++ b/training/tests/e2e/test_sft_resume_e2e.py @@ -65,7 +65,7 @@ def test_sft_resume_from_checkpoint( dataset_path = f.name try: - _make_chat_dataset(dataset_path, num_examples=10) + _make_chat_dataset(dataset_path, num_examples=50) shared_infra = InfraConfig( region=e2e_region, @@ -81,9 +81,10 @@ def test_sft_resume_from_checkpoint( tokenizer_model=tokenizer_model, learning_rate=1e-4, epochs=2, + batch_size=4, grad_accum=2, max_seq_len=4096, - max_examples=10, + max_examples=50, dcp_save_interval=4, infra=shared_infra, ) @@ -108,9 +109,10 @@ def test_sft_resume_from_checkpoint( tokenizer_model=tokenizer_model, learning_rate=1e-4, epochs=2, + batch_size=4, grad_accum=2, max_seq_len=4096, - max_examples=10, + max_examples=50, infra=shared_infra, init_from_checkpoint=f"{phase1_job_id}:{dcp_name}", ) From 8eb66ba530328b27f78c70a10b63d3fb5c28f94e Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 20:40:21 -0700 Subject: [PATCH 12/28] chore: remove test artifact and add to gitignore Made-with: Cursor --- .gitignore | 1 + sft_logs/local_checkpoint_state.jsonl | 1 - 2 files changed, 1 insertion(+), 1 deletion(-) delete mode 100644 sft_logs/local_checkpoint_state.jsonl diff --git a/.gitignore b/.gitignore index 77aa39a1..1dc657a8 100644 --- a/.gitignore +++ b/.gitignore @@ -295,3 +295,4 @@ cython_debug/ .idea/ wandb/ *_dev.sh +sft_logs/ diff --git a/sft_logs/local_checkpoint_state.jsonl b/sft_logs/local_checkpoint_state.jsonl deleted file mode 100644 index 67e1a5c5..00000000 --- a/sft_logs/local_checkpoint_state.jsonl +++ /dev/null @@ -1 +0,0 @@ -{"step": 1, "data_consumed": 20, "dcp_name": "step-1", "dataset_fingerprint": "e12293449e59", "training_shape_id": null, "source_job_id": "bgyhqcmj03tee38w"} From cb1f62d9fe0de0274e2e897b0aa4e93e1c9ac0f8 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 21:06:38 -0700 Subject: [PATCH 13/28] fix: use separate log_path for phase1 and phase2 in resume test Made-with: Cursor --- training/tests/e2e/test_sft_resume_e2e.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/training/tests/e2e/test_sft_resume_e2e.py b/training/tests/e2e/test_sft_resume_e2e.py index 91134282..2ecba661 100644 --- a/training/tests/e2e/test_sft_resume_e2e.py +++ b/training/tests/e2e/test_sft_resume_e2e.py @@ -75,6 +75,10 @@ def test_sft_resume_from_checkpoint( # Phase 1: train, save DCP logger.info("PHASE 1: initial SFT training") + import tempfile as _tf + phase1_log = _tf.mkdtemp(prefix="sft_resume_p1_") + phase2_log = _tf.mkdtemp(prefix="sft_resume_p2_") + phase1_config = Config( base_model=e2e_model, dataset=dataset_path, @@ -86,6 +90,7 @@ def test_sft_resume_from_checkpoint( max_seq_len=4096, max_examples=50, dcp_save_interval=4, + log_path=phase1_log, infra=shared_infra, ) @@ -113,6 +118,7 @@ def test_sft_resume_from_checkpoint( grad_accum=2, max_seq_len=4096, max_examples=50, + log_path=phase2_log, infra=shared_infra, init_from_checkpoint=f"{phase1_job_id}:{dcp_name}", ) From ce9d3d4e4f60ca0bc511beb6c04596a31bbd37b2 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 21:19:04 -0700 Subject: [PATCH 14/28] fix: init_from_checkpoint overrides existing state file When init_from_checkpoint is explicitly set, it now takes priority over the existing local_checkpoint_state.jsonl. The old state file is cleared so the new run starts fresh from step=0 with the specified DCP weights. Previously, resolve_resume checked the state file first (priority 1), which silently ignored init_from_checkpoint when a state file existed. This caused resumed runs to complete immediately with no new training. Made-with: Cursor --- training/tests/e2e/test_sft_resume_e2e.py | 7 ++-- training/utils/checkpoint_utils.py | 42 +++++++++++++++-------- 2 files changed, 30 insertions(+), 19 deletions(-) diff --git a/training/tests/e2e/test_sft_resume_e2e.py b/training/tests/e2e/test_sft_resume_e2e.py index 2ecba661..7c590dc8 100644 --- a/training/tests/e2e/test_sft_resume_e2e.py +++ b/training/tests/e2e/test_sft_resume_e2e.py @@ -76,8 +76,7 @@ def test_sft_resume_from_checkpoint( logger.info("PHASE 1: initial SFT training") import tempfile as _tf - phase1_log = _tf.mkdtemp(prefix="sft_resume_p1_") - phase2_log = _tf.mkdtemp(prefix="sft_resume_p2_") + log_dir = _tf.mkdtemp(prefix="sft_resume_") phase1_config = Config( base_model=e2e_model, @@ -90,7 +89,7 @@ def test_sft_resume_from_checkpoint( max_seq_len=4096, max_examples=50, dcp_save_interval=4, - log_path=phase1_log, + log_path=log_dir, infra=shared_infra, ) @@ -118,7 +117,7 @@ def test_sft_resume_from_checkpoint( grad_accum=2, max_seq_len=4096, max_examples=50, - log_path=phase2_log, + log_path=log_dir, infra=shared_infra, init_from_checkpoint=f"{phase1_job_id}:{dcp_name}", ) diff --git a/training/utils/checkpoint_utils.py b/training/utils/checkpoint_utils.py index e34a9f6a..94b5a8d7 100644 --- a/training/utils/checkpoint_utils.py +++ b/training/utils/checkpoint_utils.py @@ -45,25 +45,16 @@ def resolve_resume( """Determine resume state from ``local_checkpoint_state.jsonl``. Priority: - 1. If ``local_checkpoint_state.jsonl`` has entries, resume from the last one. - 2. If empty/missing but *init_from_checkpoint* is set, start fresh - with DCP weights (step=0, data_consumed=0). + 1. If *init_from_checkpoint* is set, always start fresh with those + DCP weights (step=0, data_consumed=0). Any existing state file + is cleared — this is an explicit directive to start from specific + weights, not continue a previous run. + 2. If ``local_checkpoint_state.jsonl`` has entries, resume from the + last one. 3. Otherwise, completely fresh start (no DCP). *init_from_checkpoint* supports cross-job format ``"job_id:checkpoint_name"``. """ - last = _load_last_loop_state(log_path) - if last is not None: - logger.info("Resuming from local_checkpoint_state.jsonl: %s", last) - return ResumeState( - step=last.get("step", 0), - data_consumed=last.get("data_consumed", 0), - dcp_name=last.get("dcp_name"), - dataset_fingerprint=last.get("dataset_fingerprint"), - training_shape_id=last.get("training_shape_id"), - source_job_id=last.get("source_job_id"), - ) - if init_from_checkpoint: source_job_id = None dcp_name = init_from_checkpoint @@ -80,9 +71,22 @@ def resolve_resume( dcp_name=dcp_name, source_job_id=source_job_id, ) + _clear_state_file(log_path) save_loop_state(log_path, _state_to_dict(initial)) return initial + last = _load_last_loop_state(log_path) + if last is not None: + logger.info("Resuming from local_checkpoint_state.jsonl: %s", last) + return ResumeState( + step=last.get("step", 0), + data_consumed=last.get("data_consumed", 0), + dcp_name=last.get("dcp_name"), + dataset_fingerprint=last.get("dataset_fingerprint"), + training_shape_id=last.get("training_shape_id"), + source_job_id=last.get("source_job_id"), + ) + logger.info("Fresh start (no checkpoint)") return ResumeState() @@ -120,6 +124,14 @@ def save_loop_state(log_path: str, loop_state: dict[str, Any]) -> None: logger.info("Saved loop state: %s", loop_state) +def _clear_state_file(log_path: str) -> None: + """Remove the state file so a fresh run starts clean.""" + path = os.path.join(log_path, STATE_FILE) + if os.path.exists(path): + os.remove(path) + logger.info("Cleared previous state file: %s", path) + + def _load_last_loop_state(log_path: str) -> dict[str, Any] | None: """Read the most recent entry from ``local_checkpoint_state.jsonl``.""" path = os.path.join(log_path, STATE_FILE) From 0c7c4bc34afed11e7a94f7312f59e902fc202135 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 21:54:16 -0700 Subject: [PATCH 15/28] fix: correct phase 2 assertion for init_from_checkpoint (starts from step 0) Made-with: Cursor --- training/tests/e2e/test_sft_resume_e2e.py | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/training/tests/e2e/test_sft_resume_e2e.py b/training/tests/e2e/test_sft_resume_e2e.py index 7c590dc8..68bf63d4 100644 --- a/training/tests/e2e/test_sft_resume_e2e.py +++ b/training/tests/e2e/test_sft_resume_e2e.py @@ -127,10 +127,14 @@ def test_sft_resume_from_checkpoint( assert isinstance(phase2_metrics, dict) assert "steps" in phase2_metrics phase2_steps = phase2_metrics["steps"] - assert ( - phase2_steps > phase1_steps - ), f"Expected global_step > {phase1_steps} after resume, got {phase2_steps}" + assert phase2_steps >= 2, ( + f"Expected >= 2 steps in phase 2 (init_from_checkpoint), got {phase2_steps}" + ) - logger.info("Resume verified: phase1=%d, phase2=%d", phase1_steps, phase2_steps) + logger.info( + "Resume verified: phase1=%d steps, phase2=%d steps (from init_from_checkpoint)", + phase1_steps, + phase2_steps, + ) finally: os.unlink(dataset_path) From 7e676708d1d8aaa1d7420d9a6bff56c5d60b03d1 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 22:34:05 -0700 Subject: [PATCH 16/28] fix: eliminate client.inner usage, ensure all recipes save DCP - Add save_weights_for_sampler_ext() to ReconnectableClient wrapper so recipes never need to reach into client.inner - Replace all client.inner.save_state() with client.save_state() (SFT loop had 2 occurrences) - Replace all client.inner.save_weights_for_sampler_ext() with client.save_weights_for_sampler_ext() (SFT + ORPO) - Add missing DCP save to ORPO final checkpoint (was only saving sampler/HF format, no optimizer state for resume) - Add missing DCP save to DPO final checkpoint (was conditional on hotload interval, now always saves DCP for resume) Made-with: Cursor --- training/recipes/dpo_loop.py | 6 ++++-- training/recipes/orpo_loop.py | 5 ++++- training/recipes/sft_loop.py | 6 +++--- training/utils/client.py | 3 +++ 4 files changed, 14 insertions(+), 6 deletions(-) diff --git a/training/recipes/dpo_loop.py b/training/recipes/dpo_loop.py index 7d4d42a1..bfd5aa43 100644 --- a/training/recipes/dpo_loop.py +++ b/training/recipes/dpo_loop.py @@ -537,8 +537,10 @@ def _signal_handler(signum, frame): # -- Final checkpoint -------------------------------------------------- hl = cfg.hotload - if step > step_offset and (hl.hot_load_interval > 0 or hl.dcp_save_interval > 0): - weight_syncer.save_and_hotload(f"final-step-{step}") + if step > step_offset: + weight_syncer.save_dcp(f"step-{step}") + if hl.hot_load_interval > 0: + weight_syncer.save_and_hotload(f"final-step-{step}") logger.info("Training complete: %d optimizer steps (%d new)", step, step - step_offset) return {"steps": step, "policy_job_id": policy_job_id, "reference_job_id": reference_job_id} diff --git a/training/recipes/orpo_loop.py b/training/recipes/orpo_loop.py index 8b839ca4..80fed530 100644 --- a/training/recipes/orpo_loop.py +++ b/training/recipes/orpo_loop.py @@ -327,8 +327,11 @@ def _signal_handler(signum, frame): # -- Final checkpoint ------------------------------------------------ if step > step_offset: + logger.info("Saving final DCP checkpoint (step %d)...", step) + client.save_state(f"step-{step}") + logger.info("Saving final base checkpoint (step %d)...", step) - result = client.inner.save_weights_for_sampler_ext( + result = client.save_weights_for_sampler_ext( f"final-step-{step}", checkpoint_type="base" ) logger.info("Final base checkpoint saved: %s", result.path) diff --git a/training/recipes/sft_loop.py b/training/recipes/sft_loop.py index 431d420e..73abfa50 100644 --- a/training/recipes/sft_loop.py +++ b/training/recipes/sft_loop.py @@ -260,7 +260,7 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ if cfg.dcp_save_interval > 0 and step % cfg.dcp_save_interval == 0: with timer("dcp_save"): logger.info("Saving DCP checkpoint at step %d", step) - client.inner.save_state(f"step-{step}") + client.save_state(f"step-{step}") save_loop_state(cfg.log_path, { "step": step, "data_consumed": data_consumed, @@ -318,7 +318,7 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ if step > state.step: logger.info("Saving final DCP checkpoint (step %d)...", step) - client.inner.save_state(f"step-{step}") + client.save_state(f"step-{step}") save_loop_state(cfg.log_path, { "step": step, "data_consumed": data_consumed, @@ -329,7 +329,7 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ }) logger.info("Saving final base checkpoint (step %d)...", step) - result = client.inner.save_weights_for_sampler_ext( + result = client.save_weights_for_sampler_ext( f"final-step-{step}", checkpoint_type="base" ) logger.info("Final base checkpoint saved: %s", result.path) diff --git a/training/utils/client.py b/training/utils/client.py index 5999cf5a..c51fa486 100644 --- a/training/utils/client.py +++ b/training/utils/client.py @@ -112,6 +112,9 @@ def save_state(self, name: str, timeout: int = DCP_TIMEOUT_S): def load_state_with_optimizer(self, path: str, timeout: int = DCP_TIMEOUT_S): return self._client.load_state_with_optimizer(path).result(timeout=timeout) + def save_weights_for_sampler_ext(self, name: str, checkpoint_type: str | None = None, timeout: int = DCP_TIMEOUT_S): + return self.inner.save_weights_for_sampler_ext(name, checkpoint_type=checkpoint_type) + def resolve_checkpoint_path(self, name: str, source_job_id: str | None = None) -> str: return self.inner.resolve_checkpoint_path(name, source_job_id=source_job_id) From 526771de311e2f155ca940462b4943d562d7f81c Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 22:36:00 -0700 Subject: [PATCH 17/28] test: reduce SFT resume test to 20 examples for faster iteration Made-with: Cursor --- training/tests/e2e/test_sft_resume_e2e.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/training/tests/e2e/test_sft_resume_e2e.py b/training/tests/e2e/test_sft_resume_e2e.py index 68bf63d4..5539312b 100644 --- a/training/tests/e2e/test_sft_resume_e2e.py +++ b/training/tests/e2e/test_sft_resume_e2e.py @@ -65,7 +65,7 @@ def test_sft_resume_from_checkpoint( dataset_path = f.name try: - _make_chat_dataset(dataset_path, num_examples=50) + _make_chat_dataset(dataset_path, num_examples=20) shared_infra = InfraConfig( region=e2e_region, @@ -87,8 +87,8 @@ def test_sft_resume_from_checkpoint( batch_size=4, grad_accum=2, max_seq_len=4096, - max_examples=50, - dcp_save_interval=4, + max_examples=20, + dcp_save_interval=2, log_path=log_dir, infra=shared_infra, ) @@ -116,7 +116,7 @@ def test_sft_resume_from_checkpoint( batch_size=4, grad_accum=2, max_seq_len=4096, - max_examples=50, + max_examples=20, log_path=log_dir, infra=shared_infra, init_from_checkpoint=f"{phase1_job_id}:{dcp_name}", From 5e0e641db7cb5131898d83eb13cc0328ce3e9081 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 23:17:55 -0700 Subject: [PATCH 18/28] fix: resolve unit test failures + add checkpoint_utils tests - Fix test_smoke_imports: replace deleted training.utils.resume with training.utils.checkpoint_utils - Fix test_shape_override_paths: validated path now passes max_context_length from profile (not None) - Add 22 unit tests for checkpoint_utils: resolve_resume (fresh start, state file resume, init_from_checkpoint, cross-job, override), data_consumed slicing, dataset_fingerprint, validate_dataset, validate_training_shape, save_loop_state Made-with: Cursor --- training/tests/test_smoke_imports.py | 2 +- training/tests/unit/test_checkpoint_utils.py | 181 ++++++++++++++++++ .../tests/unit/test_shape_override_paths.py | 2 +- 3 files changed, 183 insertions(+), 2 deletions(-) create mode 100644 training/tests/unit/test_checkpoint_utils.py diff --git a/training/tests/test_smoke_imports.py b/training/tests/test_smoke_imports.py index aab1d064..679b9e00 100644 --- a/training/tests/test_smoke_imports.py +++ b/training/tests/test_smoke_imports.py @@ -42,7 +42,7 @@ def test_recipe_imports(module: str): "training.utils.infra", "training.utils.losses", "training.utils.logging", - "training.utils.resume", + "training.utils.checkpoint_utils", "training.utils.timer", "training.utils.validation", "training.utils.rl", diff --git a/training/tests/unit/test_checkpoint_utils.py b/training/tests/unit/test_checkpoint_utils.py new file mode 100644 index 00000000..cf652c04 --- /dev/null +++ b/training/tests/unit/test_checkpoint_utils.py @@ -0,0 +1,181 @@ +"""Unit tests for checkpoint_utils -- resume state, data position, fingerprints.""" + +import json +import os +import tempfile + +import pytest + +from training.utils.checkpoint_utils import ( + ResumeState, + resolve_resume, + save_loop_state, + dataset_fingerprint, + validate_dataset, + validate_training_shape, + _load_last_loop_state, + STATE_FILE, +) + + +@pytest.fixture +def log_dir(): + with tempfile.TemporaryDirectory() as d: + yield d + + +class TestResolveResume: + def test_fresh_start_no_state_no_init(self, log_dir): + state = resolve_resume(log_dir) + assert state.step == 0 + assert state.data_consumed == 0 + assert state.dcp_name is None + + def test_resume_from_state_file(self, log_dir): + save_loop_state(log_dir, { + "step": 5, + "data_consumed": 40, + "dcp_name": "step-5", + "dataset_fingerprint": "abc123", + "training_shape_id": "ts-qwen3-8b-policy", + "source_job_id": "job-abc", + }) + state = resolve_resume(log_dir) + assert state.step == 5 + assert state.data_consumed == 40 + assert state.dcp_name == "step-5" + assert state.dataset_fingerprint == "abc123" + assert state.training_shape_id == "ts-qwen3-8b-policy" + assert state.source_job_id == "job-abc" + + def test_resume_reads_last_entry(self, log_dir): + save_loop_state(log_dir, {"step": 2, "data_consumed": 16, "dcp_name": "step-2"}) + save_loop_state(log_dir, {"step": 4, "data_consumed": 32, "dcp_name": "step-4"}) + state = resolve_resume(log_dir) + assert state.step == 4 + assert state.data_consumed == 32 + + def test_init_from_checkpoint_simple(self, log_dir): + state = resolve_resume(log_dir, init_from_checkpoint="step-10") + assert state.step == 0 + assert state.data_consumed == 0 + assert state.dcp_name == "step-10" + assert state.source_job_id is None + + def test_init_from_checkpoint_cross_job(self, log_dir): + state = resolve_resume(log_dir, init_from_checkpoint="job-xyz:step-5") + assert state.step == 0 + assert state.dcp_name == "step-5" + assert state.source_job_id == "job-xyz" + + def test_init_from_checkpoint_overrides_existing_state(self, log_dir): + save_loop_state(log_dir, { + "step": 10, + "data_consumed": 80, + "dcp_name": "step-10", + }) + state = resolve_resume(log_dir, init_from_checkpoint="other-job:step-3") + assert state.step == 0 + assert state.data_consumed == 0 + assert state.dcp_name == "step-3" + assert state.source_job_id == "other-job" + + last = _load_last_loop_state(log_dir) + assert last["step"] == 0 + assert last["dcp_name"] == "step-3" + + def test_init_from_checkpoint_gcs_path(self, log_dir): + state = resolve_resume(log_dir, init_from_checkpoint="gs://bucket/path/step-5") + assert state.dcp_name == "gs://bucket/path/step-5" + assert state.source_job_id is None + + +class TestDataConsumedSlicing: + def test_slice_from_zero(self): + data = list(range(100)) + state = ResumeState(data_consumed=0) + remaining = data[state.data_consumed:] + assert len(remaining) == 100 + assert remaining[0] == 0 + + def test_slice_from_middle(self): + data = list(range(100)) + state = ResumeState(data_consumed=40) + remaining = data[state.data_consumed:] + assert len(remaining) == 60 + assert remaining[0] == 40 + + def test_slice_past_end(self): + data = list(range(100)) + state = ResumeState(data_consumed=100) + remaining = data[state.data_consumed:] + assert len(remaining) == 0 + + +class TestDatasetFingerprint: + def test_consistent(self): + rows = [{"a": 1}, {"b": 2}, {"c": 3}] + fp1 = dataset_fingerprint(rows) + fp2 = dataset_fingerprint(rows) + assert fp1 == fp2 + assert len(fp1) == 12 + + def test_different_data(self): + fp1 = dataset_fingerprint([{"a": 1}]) + fp2 = dataset_fingerprint([{"a": 2}]) + assert fp1 != fp2 + + def test_different_length(self): + fp1 = dataset_fingerprint([{"a": 1}]) + fp2 = dataset_fingerprint([{"a": 1}, {"b": 2}]) + assert fp1 != fp2 + + def test_empty(self): + assert dataset_fingerprint([]) == "empty" + + +class TestValidateDataset: + def test_no_warning_on_match(self, log_dir, caplog): + validate_dataset("abc123", "abc123", 40) + assert "Dataset changed" not in caplog.text + + def test_warning_on_mismatch(self, log_dir, caplog): + validate_dataset("abc123", "xyz789", 40) + assert "Dataset changed" in caplog.text + + def test_no_warning_when_saved_is_none(self, log_dir, caplog): + validate_dataset(None, "abc123", 0) + assert "Dataset changed" not in caplog.text + + +class TestValidateTrainingShape: + def test_no_warning_on_match(self, caplog): + validate_training_shape("ts-qwen3-8b", "ts-qwen3-8b") + assert "Training shape changed" not in caplog.text + + def test_warning_on_mismatch(self, caplog): + validate_training_shape("ts-qwen3-8b", "ts-qwen3-30b") + assert "Training shape changed" in caplog.text + + def test_no_warning_when_none(self, caplog): + validate_training_shape(None, "ts-qwen3-8b") + assert "Training shape changed" not in caplog.text + + +class TestSaveLoopState: + def test_appends_entries(self, log_dir): + save_loop_state(log_dir, {"step": 1}) + save_loop_state(log_dir, {"step": 2}) + + path = os.path.join(log_dir, STATE_FILE) + with open(path) as f: + lines = [l.strip() for l in f if l.strip()] + assert len(lines) == 2 + assert json.loads(lines[0])["step"] == 1 + assert json.loads(lines[1])["step"] == 2 + + def test_creates_directory(self): + with tempfile.TemporaryDirectory() as parent: + nested = os.path.join(parent, "sub", "dir") + save_loop_state(nested, {"step": 1}) + assert os.path.exists(os.path.join(nested, STATE_FILE)) diff --git a/training/tests/unit/test_shape_override_paths.py b/training/tests/unit/test_shape_override_paths.py index 77eadc63..d7f152ed 100644 --- a/training/tests/unit/test_shape_override_paths.py +++ b/training/tests/unit/test_shape_override_paths.py @@ -102,7 +102,7 @@ def test_config_has_minimal_fields(self): assert c.accelerator_type is None assert c.accelerator_count is None assert c.custom_image_tag is None - assert c.max_context_length is None + assert c.max_context_length == PROFILE.max_supported_context_length assert c.node_count is None def test_payload_omits_shape_derived_fields(self): From d9d4c912912529226cb2fb6d81cd026090709974 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Mon, 9 Mar 2026 23:25:49 -0700 Subject: [PATCH 19/28] fix: always log perf/dcp_load_time (remove redundant guard) Made-with: Cursor --- training/recipes/dpo_loop.py | 3 +-- training/recipes/rl_loop.py | 3 +-- training/recipes/sft_loop.py | 3 +-- 3 files changed, 3 insertions(+), 6 deletions(-) diff --git a/training/recipes/dpo_loop.py b/training/recipes/dpo_loop.py index bfd5aa43..ad103544 100644 --- a/training/recipes/dpo_loop.py +++ b/training/recipes/dpo_loop.py @@ -487,8 +487,7 @@ def _signal_handler(signum, frame): state = resolve_resume(cfg.log_path, cfg.init_from_checkpoint) dcp_load_time = load_dcp(policy, state) - if dcp_load_time > 0: - wandb_log({"perf/dcp_load_time": dcp_load_time}, state.step) + wandb_log({"perf/dcp_load_time": dcp_load_time}, state.step) step_offset = state.step adam_params = tinker.AdamParams(learning_rate=cfg.learning_rate, **DEFAULT_ADAM) diff --git a/training/recipes/rl_loop.py b/training/recipes/rl_loop.py index e3c0f40f..5a420787 100644 --- a/training/recipes/rl_loop.py +++ b/training/recipes/rl_loop.py @@ -390,8 +390,7 @@ def _signal_handler(signum, frame): state = resolve_resume(cfg.log_path, cfg.init_from_checkpoint) dcp_load_time = load_dcp(policy, state) - if dcp_load_time > 0: - wandb_log({"perf/dcp_load_time": dcp_load_time}, state.step) + wandb_log({"perf/dcp_load_time": dcp_load_time}, state.step) step_offset = state.step if cfg.hotload.hot_load_before_training and cfg.deployment.deployment_id: diff --git a/training/recipes/sft_loop.py b/training/recipes/sft_loop.py index 73abfa50..5d17eb90 100644 --- a/training/recipes/sft_loop.py +++ b/training/recipes/sft_loop.py @@ -217,8 +217,7 @@ def _signal_handler(signum, frame): state = resolve_resume(cfg.log_path, cfg.init_from_checkpoint) dcp_load_time = load_dcp(client, state) - if dcp_load_time > 0: - wandb_log({"perf/dcp_load_time": dcp_load_time}, state.step) + wandb_log({"perf/dcp_load_time": dcp_load_time}, state.step) fp = dataset_fingerprint(raw_data) validate_dataset(state.dataset_fingerprint, fp, state.data_consumed) From eda0f23cfc5efb4d1270431d6c59f722275276af Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Tue, 10 Mar 2026 00:14:49 -0700 Subject: [PATCH 20/28] docs: add e2e resume test README with setup instructions Made-with: Cursor --- training/tests/e2e/README.md | 81 ++++++++++++++++++++++++++++++++++++ 1 file changed, 81 insertions(+) create mode 100644 training/tests/e2e/README.md diff --git a/training/tests/e2e/README.md b/training/tests/e2e/README.md new file mode 100644 index 00000000..b92e74b6 --- /dev/null +++ b/training/tests/e2e/README.md @@ -0,0 +1,81 @@ +# E2E Resume Tests + +End-to-end tests for DCP checkpoint save/load across SFT, GRPO, and DPO recipes. + +## Prerequisites + +```bash +pip install -e "training[dev]" +``` + +## Environment Variables + +| Variable | Required | Description | +|----------|----------|-------------| +| `FIREWORKS_API_KEY` | Yes | Shared dev key: `fw_3ZkNBrXgLw1EJ4y77kqSMBU5` | +| `FIREWORKS_ACCOUNT_ID` | Yes | `pyroworks-dev` | +| `FIREWORKS_BASE_URL` | Yes | `https://dev.api.fireworks.ai` | + +## Training Shapes (pyroworks-dev) + +| Shape | Nodes | GPUs | Model | Trainer Tag | Region | +|-------|-------|------|-------|-------------|--------| +| `ts-qwen3-30b-a3b-policy` | 1 | 8xB200 | qwen3-30b-a3b | 0.33.0 | US_OHIO_1 | +| `ts-qwen3-30b-a3b-128k-2node` | 2 | 16xB200 | qwen3-30b-a3b | 0.0.0-dev-chengxili-dcp-ci-v3 | US_OHIO_1 | + +## Running Tests + +### SFT Resume + +```bash +FIREWORKS_API_KEY="fw_3ZkNBrXgLw1EJ4y77kqSMBU5" \ +FIREWORKS_ACCOUNT_ID="pyroworks-dev" \ +FIREWORKS_BASE_URL="https://dev.api.fireworks.ai" \ +python -m pytest \ + "training/tests/e2e/test_sft_resume_e2e.py::TestSFTResumeE2E::test_sft_resume_from_checkpoint" \ + -v -s +``` + +Phase 1 trains SFT on synthetic data (20 examples, 5 steps), saves DCP checkpoints. +Phase 2 uses `init_from_checkpoint` to load DCP from phase 1 and trains 5 more steps. +Verifies model weights are preserved (loss starts low, not from scratch). + +Expected runtime: ~15-25 minutes (2x job creation + model download + training). + +### GRPO Resume + +```bash +FIREWORKS_API_KEY="fw_3ZkNBrXgLw1EJ4y77kqSMBU5" \ +FIREWORKS_ACCOUNT_ID="pyroworks-dev" \ +FIREWORKS_BASE_URL="https://dev.api.fireworks.ai" \ +python -m pytest \ + "training/tests/e2e/test_grpo_resume_e2e.py::TestGRPOResumeE2E::test_grpo_resume_from_checkpoint" \ + -v -s +``` + +### DPO Resume + +```bash +FIREWORKS_API_KEY="fw_3ZkNBrXgLw1EJ4y77kqSMBU5" \ +FIREWORKS_ACCOUNT_ID="pyroworks-dev" \ +FIREWORKS_BASE_URL="https://dev.api.fireworks.ai" \ +python -m pytest \ + "training/tests/e2e/test_dpo_resume_e2e.py::TestDPOResumeE2E::test_dpo_resume_from_checkpoint" \ + -v -s +``` + +## Cleanup + +Tests automatically delete trainer jobs on completion. If a test is interrupted, +clean up stale jobs manually: + +```bash +# List running jobs +curl -s "https://dev.api.fireworks.ai/v1/accounts/pyroworks-dev/rlorTrainerJobs" \ + -H "Authorization: Bearer fw_3ZkNBrXgLw1EJ4y77kqSMBU5" | python3 -m json.tool + +# Delete a specific job +curl -s -X DELETE \ + "https://dev.api.fireworks.ai/v1/accounts/pyroworks-dev/rlorTrainerJobs/?ignoreChecks=true" \ + -H "Authorization: Bearer fw_3ZkNBrXgLw1EJ4y77kqSMBU5" +``` From e02c02b3c1953fb87764343b79c379df1f4771fd Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Tue, 10 Mar 2026 00:15:31 -0700 Subject: [PATCH 21/28] Revert "docs: add e2e resume test README with setup instructions" This reverts commit eda0f23cfc5efb4d1270431d6c59f722275276af. --- training/tests/e2e/README.md | 81 ------------------------------------ 1 file changed, 81 deletions(-) delete mode 100644 training/tests/e2e/README.md diff --git a/training/tests/e2e/README.md b/training/tests/e2e/README.md deleted file mode 100644 index b92e74b6..00000000 --- a/training/tests/e2e/README.md +++ /dev/null @@ -1,81 +0,0 @@ -# E2E Resume Tests - -End-to-end tests for DCP checkpoint save/load across SFT, GRPO, and DPO recipes. - -## Prerequisites - -```bash -pip install -e "training[dev]" -``` - -## Environment Variables - -| Variable | Required | Description | -|----------|----------|-------------| -| `FIREWORKS_API_KEY` | Yes | Shared dev key: `fw_3ZkNBrXgLw1EJ4y77kqSMBU5` | -| `FIREWORKS_ACCOUNT_ID` | Yes | `pyroworks-dev` | -| `FIREWORKS_BASE_URL` | Yes | `https://dev.api.fireworks.ai` | - -## Training Shapes (pyroworks-dev) - -| Shape | Nodes | GPUs | Model | Trainer Tag | Region | -|-------|-------|------|-------|-------------|--------| -| `ts-qwen3-30b-a3b-policy` | 1 | 8xB200 | qwen3-30b-a3b | 0.33.0 | US_OHIO_1 | -| `ts-qwen3-30b-a3b-128k-2node` | 2 | 16xB200 | qwen3-30b-a3b | 0.0.0-dev-chengxili-dcp-ci-v3 | US_OHIO_1 | - -## Running Tests - -### SFT Resume - -```bash -FIREWORKS_API_KEY="fw_3ZkNBrXgLw1EJ4y77kqSMBU5" \ -FIREWORKS_ACCOUNT_ID="pyroworks-dev" \ -FIREWORKS_BASE_URL="https://dev.api.fireworks.ai" \ -python -m pytest \ - "training/tests/e2e/test_sft_resume_e2e.py::TestSFTResumeE2E::test_sft_resume_from_checkpoint" \ - -v -s -``` - -Phase 1 trains SFT on synthetic data (20 examples, 5 steps), saves DCP checkpoints. -Phase 2 uses `init_from_checkpoint` to load DCP from phase 1 and trains 5 more steps. -Verifies model weights are preserved (loss starts low, not from scratch). - -Expected runtime: ~15-25 minutes (2x job creation + model download + training). - -### GRPO Resume - -```bash -FIREWORKS_API_KEY="fw_3ZkNBrXgLw1EJ4y77kqSMBU5" \ -FIREWORKS_ACCOUNT_ID="pyroworks-dev" \ -FIREWORKS_BASE_URL="https://dev.api.fireworks.ai" \ -python -m pytest \ - "training/tests/e2e/test_grpo_resume_e2e.py::TestGRPOResumeE2E::test_grpo_resume_from_checkpoint" \ - -v -s -``` - -### DPO Resume - -```bash -FIREWORKS_API_KEY="fw_3ZkNBrXgLw1EJ4y77kqSMBU5" \ -FIREWORKS_ACCOUNT_ID="pyroworks-dev" \ -FIREWORKS_BASE_URL="https://dev.api.fireworks.ai" \ -python -m pytest \ - "training/tests/e2e/test_dpo_resume_e2e.py::TestDPOResumeE2E::test_dpo_resume_from_checkpoint" \ - -v -s -``` - -## Cleanup - -Tests automatically delete trainer jobs on completion. If a test is interrupted, -clean up stale jobs manually: - -```bash -# List running jobs -curl -s "https://dev.api.fireworks.ai/v1/accounts/pyroworks-dev/rlorTrainerJobs" \ - -H "Authorization: Bearer fw_3ZkNBrXgLw1EJ4y77kqSMBU5" | python3 -m json.tool - -# Delete a specific job -curl -s -X DELETE \ - "https://dev.api.fireworks.ai/v1/accounts/pyroworks-dev/rlorTrainerJobs/?ignoreChecks=true" \ - -H "Authorization: Bearer fw_3ZkNBrXgLw1EJ4y77kqSMBU5" -``` From d559f870cdb7368204558aa0a39f1e577cb15f08 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Tue, 10 Mar 2026 11:16:49 -0700 Subject: [PATCH 22/28] refactor: use tinker_cookbook checkpoint format + dataset utils Switch checkpoint persistence from custom local_checkpoint_state.jsonl to tinker_cookbook's checkpoints.jsonl format. Import get_last_checkpoint from tinker_cookbook for reading; implement sync save_checkpoint that stores cross-job refs resolvable by any future trainer job. - Rewrite checkpoint_utils.py: ResumeState -> ResumeInfo, save_loop_state -> save_checkpoint, load_dcp -> resolve_resume (loads directly) - SFT loop: use SupervisedDatasetFromHFDataset for batch-indexed iteration with per-epoch reshuffling via set_epoch() - RL loop: add RLPromptDataset for batch-indexed iteration, pass fw_api_key to ReconnectableClient (fixes shared dev auth) - Remove backward-compat fallbacks in supervised.py (datum_from_tokens_weights) - Pin tinker-cookbook to latest commit (d7fb423) - Update all recipe files, examples, and tests Made-with: Cursor --- .../examples/frozen_lake/train_frozen_lake.py | 7 +- training/pyproject.toml | 2 +- training/recipes/dpo_loop.py | 9 +- training/recipes/orpo_loop.py | 7 +- training/recipes/rl_loop.py | 66 ++--- training/recipes/sft_loop.py | 117 ++++----- training/tests/e2e/test_grpo_resume_e2e.py | 67 +++-- training/tests/e2e/test_sft_resume_e2e.py | 59 +++-- training/tests/unit/test_checkpoint_utils.py | 236 +++++++++++------- training/tests/unit/test_dpo_loop.py | 9 + training/tests/unit/test_orpo_loop.py | 13 + training/tests/unit/test_rl_loop.py | 37 ++- training/tests/unit/test_sft_loop.py | 150 ++--------- training/utils/__init__.py | 2 + training/utils/checkpoint_utils.py | 199 ++++++--------- training/utils/data.py | 21 ++ training/utils/supervised.py | 43 +--- 17 files changed, 510 insertions(+), 534 deletions(-) diff --git a/training/examples/frozen_lake/train_frozen_lake.py b/training/examples/frozen_lake/train_frozen_lake.py index b626ec33..c69d09bf 100644 --- a/training/examples/frozen_lake/train_frozen_lake.py +++ b/training/examples/frozen_lake/train_frozen_lake.py @@ -543,10 +543,9 @@ def _make_job(label: str, precreated_id: str | None, **extra_kw): infra_boot_time = time.time() - _infra_start wandb_log({"train/step": 0, "infra/total_boot_time": infra_boot_time}, step=0) - from training.utils.checkpoint_utils import resolve_resume, load_dcp - state = resolve_resume("./frozen_lake_logs") - load_dcp(policy, state) - step_offset = state.step + from training.utils.checkpoint_utils import resolve_resume + resume_info = resolve_resume(policy, "./frozen_lake_logs") + step_offset = resume_info.step if resume_info else 0 if hotload_cfg.hot_load_before_training and deploy_cfg.deployment_id: name = f"resume-{step_offset}-base" if step_offset > 0 else "step-0-base" weight_syncer.save_and_hotload(name, checkpoint_type="base") diff --git a/training/pyproject.toml b/training/pyproject.toml index 03095c37..383d02da 100644 --- a/training/pyproject.toml +++ b/training/pyproject.toml @@ -8,7 +8,7 @@ version = "0.1.0" requires-python = ">=3.11" dependencies = [ "fireworks-ai>=1.0.0a40,<2", - "tinker-cookbook @ git+https://github.com/thinking-machines-lab/tinker-cookbook.git@934f0d9b2f53c3edff02cbf23ec6da8682047fa5", + "tinker-cookbook @ git+https://github.com/thinking-machines-lab/tinker-cookbook.git@d7fb423674c0bdd27c20a9b948b6091457874039", "eval-protocol>=0.3.23", "tqdm", "torch", diff --git a/training/recipes/dpo_loop.py b/training/recipes/dpo_loop.py index ad103544..15f233ac 100644 --- a/training/recipes/dpo_loop.py +++ b/training/recipes/dpo_loop.py @@ -56,7 +56,7 @@ ) from fireworks.training.sdk.deployment import DEFAULT_DELTA_COMPRESSION from fireworks.training.sdk.weight_syncer import WeightSyncer -from training.utils.checkpoint_utils import resolve_resume, load_dcp, save_loop_state, dataset_fingerprint, validate_dataset +from training.utils.checkpoint_utils import resolve_resume from training.utils.timer import timer, flush_timing logger = logging.getLogger(__name__) @@ -485,10 +485,9 @@ def _signal_handler(signum, frame): compression_format=DEFAULT_DELTA_COMPRESSION, ) - state = resolve_resume(cfg.log_path, cfg.init_from_checkpoint) - dcp_load_time = load_dcp(policy, state) - wandb_log({"perf/dcp_load_time": dcp_load_time}, state.step) - step_offset = state.step + resume_info = resolve_resume(policy, cfg.log_path, cfg.init_from_checkpoint) + step_offset = resume_info.step if resume_info else 0 + wandb_log({"train/step": step_offset}, step_offset) adam_params = tinker.AdamParams(learning_rate=cfg.learning_rate, **DEFAULT_ADAM) # -- Cache reference logprobs concurrently ------------------------------ diff --git a/training/recipes/orpo_loop.py b/training/recipes/orpo_loop.py index 80fed530..193fe21d 100644 --- a/training/recipes/orpo_loop.py +++ b/training/recipes/orpo_loop.py @@ -61,7 +61,7 @@ render_preference_pair, resolve_renderer_name, ) -from training.utils.checkpoint_utils import resolve_resume, load_dcp +from training.utils.checkpoint_utils import resolve_resume logger = logging.getLogger(__name__) @@ -175,9 +175,8 @@ def _signal_handler(signum, frame): ) job_id = endpoint.job_id - state = resolve_resume(cfg.log_path, cfg.init_from_checkpoint) - dcp_load_time = load_dcp(client, state) - step_offset = state.step + resume_info = resolve_resume(client, cfg.log_path, cfg.init_from_checkpoint) + step_offset = resume_info.step if resume_info else 0 adam_params = tinker.AdamParams(learning_rate=cfg.learning_rate, **DEFAULT_ADAM) # -- Data ---------------------------------------------------------------- diff --git a/training/recipes/rl_loop.py b/training/recipes/rl_loop.py index 5a420787..beef7510 100644 --- a/training/recipes/rl_loop.py +++ b/training/recipes/rl_loop.py @@ -36,6 +36,7 @@ DeployConfig, HotloadConfig, ReconnectableClient, + RLPromptDataset, wandb_log, setup_wandb, wandb_finish, @@ -49,8 +50,7 @@ ) from training.utils.checkpoint_utils import ( resolve_resume, - load_dcp, - save_loop_state, + save_checkpoint, dataset_fingerprint, validate_dataset, ) @@ -351,11 +351,13 @@ def _signal_handler(signum, frame): policy = ReconnectableClient( rlor_mgr, policy_ep.job_id, cfg.base_model, cfg.lora_rank, + fw_api_key=api_key, endpoint=policy_ep if cfg.policy_base_url else None, ) reference = ( ReconnectableClient( rlor_mgr, reference_ep.job_id, cfg.base_model, cfg.lora_rank, + fw_api_key=api_key, endpoint=reference_ep if cfg.reference_base_url else None, ) if reference_ep else None @@ -388,10 +390,9 @@ def _signal_handler(signum, frame): # -- Resume --------------------------------------------------------------- - state = resolve_resume(cfg.log_path, cfg.init_from_checkpoint) - dcp_load_time = load_dcp(policy, state) - wandb_log({"perf/dcp_load_time": dcp_load_time}, state.step) - step_offset = state.step + resume_info = resolve_resume(policy, cfg.log_path, cfg.init_from_checkpoint) + step_offset = resume_info.step if resume_info else 0 + wandb_log({"train/step": step_offset}, step_offset) if cfg.hotload.hot_load_before_training and cfg.deployment.deployment_id: name = f"resume-{step_offset}-base" if step_offset > 0 else "step-0-base" @@ -399,9 +400,12 @@ def _signal_handler(signum, frame): # -- Prepare sampling and training -------------------------------------- - dataset = load_jsonl_dataset(cfg.dataset, cfg.max_rows) - fp = dataset_fingerprint(dataset) - validate_dataset(state.dataset_fingerprint, fp, state.data_consumed) + raw_dataset = load_jsonl_dataset(cfg.dataset, cfg.max_rows) + fp = dataset_fingerprint(raw_dataset) + if resume_info: + validate_dataset(resume_info.dataset_fingerprint, fp, resume_info.data_consumed) + all_rows = raw_dataset * cfg.epochs + rl_dataset = RLPromptDataset(all_rows, prompts_per_step=prompt_groups_per_step) adam_params = tinker.AdamParams(learning_rate=cfg.learning_rate, **DEFAULT_ADAM) loss_builder = build_loss_fn( policy_loss=cfg.policy_loss, kl_beta=cfg.kl_beta, @@ -531,20 +535,19 @@ def fwd_bwd_one(prompt_groups: list[PromptGroup]): """One minibatch forward/backward call after reference forward.""" if not prompt_groups: raise ValueError("fwd_bwd_one requires at least one prompt group") - import time as _t data, adv, ref_lp, prompt_lens, inf_lp = combine_prompt_groups(prompt_groups) - t0 = _t.time() + t0 = _time.time() prox_fwd = policy.forward(data, "cross_entropy") prox_lp = [prox_fwd.loss_fn_outputs[i]["logprobs"].data for i in range(len(data))] - logger.info("prox_forward: done (%.1fs)", _t.time() - t0) + logger.info("prox_forward: done (%.1fs)", _time.time() - t0) - t0 = _t.time() + t0 = _time.time() fwd_bwd_result = policy.forward_backward_custom( data, loss_builder(adv, ref_lp, prompt_lens, inf_lp, prox_lp), ) - logger.info("fwd_bwd: done (%.1fs)", _t.time() - t0) + logger.info("fwd_bwd: done (%.1fs)", _time.time() - t0) return fwd_bwd_result def finish_step( @@ -555,30 +558,27 @@ def finish_step( loop_stats: dict | None = None, ) -> tuple[int, dict]: """optim_step + hotload + metrics after all minibatches in a step.""" - import time as _t - - t0 = _t.time() + t0 = _time.time() optim_result = policy.optim_step(adam_params) step += 1 - logger.info("[step %d] optim_step: done (%.1fs)", step, _t.time() - t0) + logger.info("[step %d] optim_step: done (%.1fs)", step, _time.time() - t0) if cfg.hotload.hot_load_interval > 0 and step % cfg.hotload.hot_load_interval == 0: logger.info("[step %d] hotload: saving + loading...", step) - t0 = _t.time() + t0 = _time.time() with timer("weight_sync"): weight_syncer.save_and_hotload(f"step-{step}") - logger.info("[step %d] hotload: done (%.1fs)", step, _t.time() - t0) + logger.info("[step %d] hotload: done (%.1fs)", step, _time.time() - t0) if cfg.hotload.dcp_save_interval > 0 and step % cfg.hotload.dcp_save_interval == 0: logger.info("[step %d] dcp_save...", step) - t0 = _t.time() + t0 = _time.time() with timer("dcp_save"): weight_syncer.save_dcp(f"step-{step}") - logger.info("[step %d] dcp_save: done (%.1fs)", step, _t.time() - t0) - data_consumed = state.data_consumed + (step - state.step) * prompt_groups_per_step - save_loop_state(cfg.log_path, { + logger.info("[step %d] dcp_save: done (%.1fs)", step, _time.time() - t0) + _data_consumed = (resume_info.data_consumed if resume_info else 0) + (step - step_offset) * prompt_groups_per_step + save_checkpoint(policy, f"step-{step}", cfg.log_path, { "step": step, - "data_consumed": data_consumed, - "dcp_name": f"step-{step}", + "data_consumed": _data_consumed, "dataset_fingerprint": fp, "training_shape_id": getattr(cfg.infra, "training_shape_id", None), "source_job_id": policy_job_id, @@ -622,10 +622,12 @@ def _loop_metrics_callback(loop_metrics: dict) -> None: finish_step=finish_step, ) - all_rows = (dataset * cfg.epochs)[state.data_consumed:] + remaining_rows = [] + for i_step in range(step_offset, len(rl_dataset)): + remaining_rows.extend(rl_dataset.get_batch(i_step)) global_step = asyncio.run(run_rl_loop( - sample_fns=(sample_one_prompt(row) for row in all_rows), + sample_fns=(sample_one_prompt(row) for row in remaining_rows), minibatch_fns=train_fns, prompt_groups_per_step=prompt_groups_per_step, max_concurrent=cfg.max_concurrent, @@ -639,12 +641,10 @@ def _loop_metrics_callback(loop_metrics: dict) -> None: if global_step > step_offset: try: - policy.save_state(f"step-{global_step}", timeout=cfg.hotload.dcp_timeout) - data_consumed = state.data_consumed + (global_step - state.step) * prompt_groups_per_step - save_loop_state(cfg.log_path, { + _data_consumed = (resume_info.data_consumed if resume_info else 0) + (global_step - step_offset) * prompt_groups_per_step + save_checkpoint(policy, f"step-{global_step}", cfg.log_path, { "step": global_step, - "data_consumed": data_consumed, - "dcp_name": f"step-{global_step}", + "data_consumed": _data_consumed, "dataset_fingerprint": fp, "training_shape_id": getattr(cfg.infra, "training_shape_id", None), "source_job_id": policy_job_id, diff --git a/training/recipes/sft_loop.py b/training/recipes/sft_loop.py index 5d17eb90..75d32449 100644 --- a/training/recipes/sft_loop.py +++ b/training/recipes/sft_loop.py @@ -25,10 +25,12 @@ import tinker import json +import datasets as hf_datasets import transformers from dotenv import load_dotenv from fireworks.training.sdk import TrainerJobManager +from tinker_cookbook.supervised.data import SupervisedDatasetFromHFDataset from training.utils import ( DEFAULT_ADAM, InfraConfig, @@ -48,8 +50,7 @@ ) from training.utils.checkpoint_utils import ( resolve_resume, - load_dcp, - save_loop_state, + save_checkpoint, dataset_fingerprint, validate_dataset, ) @@ -184,61 +185,64 @@ def _signal_handler(signum, frame): break logger.info("Loaded %d examples from %s", len(raw_data), cfg.dataset) - training_data: List[tinker.Datum] = [] + max_seq_len = cfg.max_seq_len filtered_count = 0 - for row in raw_data: + + def _map_fn(row: dict) -> tinker.Datum | None: + nonlocal filtered_count messages = row.get("messages", []) if not messages: - continue - + filtered_count += 1 + return None rendered = render_messages_to_datum( - messages, - renderer=renderer, - train_on_what=train_on_what, + messages, renderer=renderer, train_on_what=train_on_what, ) - if len(rendered.token_ids) > cfg.max_seq_len or len(rendered.token_ids) < 2: + if len(rendered.token_ids) > max_seq_len or len(rendered.token_ids) < 2: filtered_count += 1 - continue - - training_data.append(rendered.datum) + return None + return rendered.datum + training_data = [d for row in raw_data if (d := _map_fn(row)) is not None] if filtered_count > 0: logger.info( "Seq-length filter: %d/%d examples filtered (len > %d or len < 2)", - filtered_count, - len(raw_data), - cfg.max_seq_len, + filtered_count, len(raw_data), max_seq_len, ) logger.info("Prepared %d training examples", len(training_data)) if not training_data: raise RuntimeError("No valid training examples after tokenization") + sft_dataset = SupervisedDatasetFromHFDataset( + hf_datasets.Dataset.from_dict({"datum_idx": list(range(len(training_data)))}), + batch_size=cfg.batch_size, + map_fn=lambda row: training_data[row["datum_idx"]], + ) + total_batches_per_epoch = len(sft_dataset) + logger.info("Dataset: %d examples, %d batches/epoch, %d epochs", + len(training_data), total_batches_per_epoch, cfg.epochs) + # -- Resume --------------------------------------------------------------- - state = resolve_resume(cfg.log_path, cfg.init_from_checkpoint) - dcp_load_time = load_dcp(client, state) - wandb_log({"perf/dcp_load_time": dcp_load_time}, state.step) + resume_info = resolve_resume(client, cfg.log_path, cfg.init_from_checkpoint) + step = resume_info.step if resume_info else 0 + data_consumed = resume_info.data_consumed if resume_info else 0 + wandb_log({"train/step": step}, step) fp = dataset_fingerprint(raw_data) - validate_dataset(state.dataset_fingerprint, fp, state.data_consumed) + if resume_info: + validate_dataset(resume_info.dataset_fingerprint, fp, data_consumed) adam_params = tinker.AdamParams(learning_rate=cfg.learning_rate, **DEFAULT_ADAM) - # -- Training loop (batched) ------------------------------------------- + # -- Training loop (batch-indexed) ------------------------------------- - batch_size = cfg.batch_size - all_examples = training_data * cfg.epochs - examples_to_process = all_examples[state.data_consumed:] - data_consumed = state.data_consumed - - step = state.step - total_steps = len(all_examples) // (cfg.grad_accum * batch_size) + start_batch = data_consumed // cfg.batch_size + total_steps_estimate = (total_batches_per_epoch * cfg.epochs) // cfg.grad_accum accum = 0 agg_loss_sum = 0.0 agg_resp_tokens = 0 def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[int, int]: - """Send a batch through forward_backward_custom and return (step, accum).""" nonlocal agg_loss_sum, agg_resp_tokens loss_fn = make_batch_weighted_sft_loss_fn() @@ -259,15 +263,13 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ if cfg.dcp_save_interval > 0 and step % cfg.dcp_save_interval == 0: with timer("dcp_save"): logger.info("Saving DCP checkpoint at step %d", step) - client.save_state(f"step-{step}") - save_loop_state(cfg.log_path, { - "step": step, - "data_consumed": data_consumed, - "dcp_name": f"step-{step}", - "dataset_fingerprint": fp, - "training_shape_id": getattr(cfg.infra, "training_shape_id", None), - "source_job_id": job_id, - }) + save_checkpoint(client, f"step-{step}", cfg.log_path, { + "step": step, + "data_consumed": data_consumed, + "dataset_fingerprint": fp, + "training_shape_id": getattr(cfg.infra, "training_shape_id", None), + "source_job_id": job_id, + }) step_metrics: Dict[str, Any] = flush_timing() @@ -280,10 +282,7 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ ppl = torch.exp(torch.tensor(avg_loss)).item() logger.info( "Step %d/%d | Loss: %.4f | PPL: %.2f", - step, - total_steps, - avg_loss, - ppl, + step, total_steps_estimate, avg_loss, ppl, ) log_metrics_json(step, ce_loss=avg_loss, ppl=ppl) step_metrics.update({ @@ -298,16 +297,13 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ return step, accum - batch_buffer: list[tinker.Datum] = [] - for ex in examples_to_process: - batch_buffer.append(ex) - data_consumed += 1 - if len(batch_buffer) >= batch_size: - step, accum = _flush_batch(batch_buffer, step, accum) - batch_buffer = [] - - if batch_buffer: - step, accum = _flush_batch(batch_buffer, step, accum) + for epoch in range(cfg.epochs): + sft_dataset.set_epoch(epoch) + epoch_start = start_batch if epoch == 0 else 0 + for i_batch in range(epoch_start, total_batches_per_epoch): + batch = sft_dataset.get_batch(i_batch) + data_consumed += len(batch) + step, accum = _flush_batch(batch, step, accum) if accum > 0: client.optim_step(adam_params) @@ -315,23 +311,16 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ # -- Final checkpoint -------------------------------------------------- - if step > state.step: - logger.info("Saving final DCP checkpoint (step %d)...", step) - client.save_state(f"step-{step}") - save_loop_state(cfg.log_path, { + start_step = resume_info.step if resume_info else 0 + if step > start_step: + logger.info("Saving final checkpoint (step %d)...", step) + save_checkpoint(client, f"step-{step}", cfg.log_path, { "step": step, "data_consumed": data_consumed, - "dcp_name": f"step-{step}", "dataset_fingerprint": fp, "training_shape_id": getattr(cfg.infra, "training_shape_id", None), "source_job_id": job_id, - }) - - logger.info("Saving final base checkpoint (step %d)...", step) - result = client.save_weights_for_sampler_ext( - f"final-step-{step}", checkpoint_type="base" - ) - logger.info("Final base checkpoint saved: %s", result.path) + }, kind="both") logger.info("Training complete: %d optimizer steps", step) return {"steps": step, "job_id": job_id} diff --git a/training/tests/e2e/test_grpo_resume_e2e.py b/training/tests/e2e/test_grpo_resume_e2e.py index 0b4b5cf7..4e393401 100644 --- a/training/tests/e2e/test_grpo_resume_e2e.py +++ b/training/tests/e2e/test_grpo_resume_e2e.py @@ -1,10 +1,11 @@ -"""E2E test: GRPO training -> DCP checkpoint -> resume. +"""E2E test: GRPO training -> DCP checkpoint -> resume from dataloader position. -Two-phase test on qwen3-30b-a3b (MoE) with Router Replay, TIS, and -hotloading: +Two-phase test on qwen3-30b-a3b (MoE) with Router Replay and TIS: - Phase 1: Train ~2 steps with hotloading and dcp_save_interval=2. - Phase 2: Create new RLOR jobs, reuse deployment, resume from checkpoint. + Phase 1: Train a few steps, save DCP checkpoints to checkpoints.jsonl. + Phase 2: New RLOR jobs, same log_dir -- resolve_resume picks up from + checkpoints.jsonl, loads cross-job checkpoint, resumes at the + saved step (continues dataloader, not from beginning). Requires: FIREWORKS_API_KEY -- API key with training/deployment access @@ -18,11 +19,12 @@ import os import re import logging +import tempfile import pytest +from tinker_cookbook.checkpoint_utils import get_last_checkpoint from training.utils import InfraConfig, DeployConfig, HotloadConfig -from training.utils.rl import ISConfig from training.tests.e2e.conftest import GSM8K_SAMPLE_URL from training.recipes.rl_loop import Config, main @@ -64,12 +66,14 @@ def test_grpo_resume_from_checkpoint( grpo_mod.reward_fn = _gsm8k_reward deployment_id = os.environ.get("GRPO_RESUME_DEPLOYMENT_ID") + log_dir = tempfile.mkdtemp(prefix="grpo_resume_") + training_shape_id = os.environ.get("FIREWORKS_E2E_TRAINING_SHAPE", "ts-qwen3-30b-a3b-policy") shared_infra = InfraConfig( region=e2e_region, - skip_validations=True, accelerator_type=e2e_training_accelerator, custom_image_tag=custom_image_tag, + training_shape_id=training_shape_id, ) # Phase 1: train ~2 steps, save DCP @@ -81,8 +85,8 @@ def test_grpo_resume_from_checkpoint( completions_per_prompt=4, max_rows=8, epochs=1, - router_replay=True, - is_correction=ISConfig(tis_cap=10.0), + kl_beta=0, + log_path=log_dir, infra=shared_infra, deployment=DeployConfig( deployment_id=deployment_id, @@ -106,21 +110,28 @@ def test_grpo_resume_from_checkpoint( phase1_steps = phase1_metrics["steps"] assert phase1_steps >= 2, f"Expected >= 2 steps in phase 1, got {phase1_steps}" - phase1_policy_job_id = phase1_metrics["policy_job_id"] - dcp_name = f"step-{phase1_steps}" - logger.info("Phase 1 done: %d steps, job=%s", phase1_steps, phase1_policy_job_id) + logger.info("Phase 1 done: %d steps, job=%s", + phase1_steps, phase1_metrics["policy_job_id"]) - # Phase 2: resume from checkpoint - logger.info("PHASE 2: resume from '%s' (source job: %s)", dcp_name, phase1_policy_job_id) + last_ckpt = get_last_checkpoint(log_dir) + assert last_ckpt is not None, "Expected at least one checkpoint in checkpoints.jsonl" + assert "state_path" in last_ckpt + saved_step = last_ckpt["step"] + saved_data_consumed = last_ckpt["data_consumed"] + logger.info("Phase 1 checkpoint: step=%d data_consumed=%d", + saved_step, saved_data_consumed) + + # Phase 2: new jobs, same log_dir -- resume from checkpoints.jsonl + logger.info("PHASE 2: resume from checkpoints.jsonl (step=%d)", saved_step) phase2_config = Config( base_model=e2e_model, dataset=GSM8K_SAMPLE_URL, completions_per_prompt=4, - max_rows=6, + max_rows=8, epochs=1, - router_replay=True, - is_correction=ISConfig(tis_cap=10.0), + kl_beta=0, + log_path=log_dir, infra=shared_infra, deployment=DeployConfig( deployment_id=deployment_id, @@ -132,7 +143,6 @@ def test_grpo_resume_from_checkpoint( hot_load_before_training=True, hot_load_timeout=600, ), - init_from_checkpoint=f"{phase1_policy_job_id}:{dcp_name}", ) phase2_metrics = main(phase2_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) @@ -140,5 +150,22 @@ def test_grpo_resume_from_checkpoint( assert isinstance(phase2_metrics, dict) assert "steps" in phase2_metrics phase2_steps = phase2_metrics["steps"] - assert phase2_steps > phase1_steps, f"Expected global_step > {phase1_steps} after resume, got {phase2_steps}" - logger.info("Resume verified: phase1=%d, phase2=%d", phase1_steps, phase2_steps) + + assert phase2_steps > saved_step, ( + f"Phase 2 should continue beyond phase 1's saved step {saved_step}, " + f"but got {phase2_steps}" + ) + + final_ckpt = get_last_checkpoint(log_dir) + assert final_ckpt is not None + assert final_ckpt["data_consumed"] > saved_data_consumed, ( + f"Phase 2 data_consumed ({final_ckpt['data_consumed']}) should exceed " + f"phase 1's ({saved_data_consumed}) -- dataloader should continue, not restart" + ) + + logger.info( + "Resume verified: phase1=%d steps (data_consumed=%d), " + "phase2=%d steps (data_consumed=%d)", + saved_step, saved_data_consumed, + phase2_steps, final_ckpt["data_consumed"], + ) diff --git a/training/tests/e2e/test_sft_resume_e2e.py b/training/tests/e2e/test_sft_resume_e2e.py index 5539312b..ccfb5825 100644 --- a/training/tests/e2e/test_sft_resume_e2e.py +++ b/training/tests/e2e/test_sft_resume_e2e.py @@ -1,6 +1,10 @@ -"""E2E test: SFT training -> DCP checkpoint -> resume -> verify continuation. +"""E2E test: SFT training -> DCP checkpoint -> resume from dataloader position. -Two-phase test on qwen3-30b-a3b. +Two-phase test on qwen3-30b-a3b: + Phase 1: Train on first portion of data, save DCP checkpoints to checkpoints.jsonl. + Phase 2: New trainer job, same log_dir -- resolve_resume picks up from + checkpoints.jsonl, loads state_path, resumes at the saved step + and data_consumed offset (continues dataloader, not from beginning). Requires: FIREWORKS_API_KEY -- API key with training access @@ -17,6 +21,7 @@ import pytest +from tinker_cookbook.checkpoint_utils import get_last_checkpoint from training.utils import InfraConfig from training.recipes.sft_loop import Config, main @@ -72,11 +77,10 @@ def test_sft_resume_from_checkpoint( custom_image_tag=custom_image_tag or "0.33.0", ) - # Phase 1: train, save DCP - logger.info("PHASE 1: initial SFT training") + log_dir = tempfile.mkdtemp(prefix="sft_resume_") - import tempfile as _tf - log_dir = _tf.mkdtemp(prefix="sft_resume_") + # Phase 1: train, save DCP checkpoints to checkpoints.jsonl + logger.info("PHASE 1: initial SFT training") phase1_config = Config( base_model=e2e_model, @@ -100,12 +104,23 @@ def test_sft_resume_from_checkpoint( phase1_steps = phase1_metrics["steps"] assert phase1_steps >= 2, f"Expected >= 2 steps in phase 1, got {phase1_steps}" - phase1_job_id = phase1_metrics["job_id"] - dcp_name = f"step-{phase1_steps}" - logger.info("Phase 1 done: %d steps, job=%s", phase1_steps, phase1_job_id) + logger.info("Phase 1 done: %d steps, job=%s", phase1_steps, phase1_metrics["job_id"]) + + last_ckpt = get_last_checkpoint(log_dir) + assert last_ckpt is not None, "Expected at least one checkpoint in checkpoints.jsonl" + assert "state_path" in last_ckpt, f"Checkpoint missing state_path: {last_ckpt}" + saved_step = last_ckpt["step"] + saved_data_consumed = last_ckpt["data_consumed"] + logger.info( + "Phase 1 checkpoint: step=%d data_consumed=%d state_path=%s", + saved_step, saved_data_consumed, last_ckpt["state_path"], + ) - # Phase 2: resume from checkpoint - logger.info("PHASE 2: resume from '%s' (source job: %s)", dcp_name, phase1_job_id) + # Phase 2: new job, same log_dir -- resume from checkpoints.jsonl + # resolve_resume reads state_path + step + data_consumed and + # continues the dataloader from where phase 1 left off. + logger.info("PHASE 2: resume from checkpoints.jsonl (step=%d, data_consumed=%d)", + saved_step, saved_data_consumed) phase2_config = Config( base_model=e2e_model, @@ -119,7 +134,6 @@ def test_sft_resume_from_checkpoint( max_examples=20, log_path=log_dir, infra=shared_infra, - init_from_checkpoint=f"{phase1_job_id}:{dcp_name}", ) phase2_metrics = main(phase2_config, rlor_mgr=rlor_mgr) @@ -127,14 +141,25 @@ def test_sft_resume_from_checkpoint( assert isinstance(phase2_metrics, dict) assert "steps" in phase2_metrics phase2_steps = phase2_metrics["steps"] - assert phase2_steps >= 2, ( - f"Expected >= 2 steps in phase 2 (init_from_checkpoint), got {phase2_steps}" + + assert phase2_steps > saved_step, ( + f"Phase 2 should continue beyond phase 1's saved step {saved_step}, " + f"but got {phase2_steps}" + ) + + final_ckpt = get_last_checkpoint(log_dir) + assert final_ckpt is not None + assert final_ckpt["step"] == phase2_steps + assert final_ckpt["data_consumed"] > saved_data_consumed, ( + f"Phase 2 data_consumed ({final_ckpt['data_consumed']}) should exceed " + f"phase 1's ({saved_data_consumed}) -- dataloader should continue, not restart" ) logger.info( - "Resume verified: phase1=%d steps, phase2=%d steps (from init_from_checkpoint)", - phase1_steps, - phase2_steps, + "Resume verified: phase1=%d steps (data_consumed=%d), " + "phase2=%d steps (data_consumed=%d) -- dataloader continued correctly", + saved_step, saved_data_consumed, + phase2_steps, final_ckpt["data_consumed"], ) finally: os.unlink(dataset_path) diff --git a/training/tests/unit/test_checkpoint_utils.py b/training/tests/unit/test_checkpoint_utils.py index cf652c04..ea40f6b0 100644 --- a/training/tests/unit/test_checkpoint_utils.py +++ b/training/tests/unit/test_checkpoint_utils.py @@ -1,20 +1,20 @@ -"""Unit tests for checkpoint_utils -- resume state, data position, fingerprints.""" +"""Unit tests for checkpoint_utils -- resume, save, fingerprints.""" import json import os import tempfile +from unittest.mock import MagicMock import pytest from training.utils.checkpoint_utils import ( - ResumeState, + ResumeInfo, resolve_resume, - save_loop_state, + save_checkpoint, dataset_fingerprint, validate_dataset, validate_training_shape, - _load_last_loop_state, - STATE_FILE, + CHECKPOINTS_BASE_NAME, ) @@ -24,92 +24,146 @@ def log_dir(): yield d +def _make_mock_client(job_id="test-job"): + """Create a mock ReconnectableClient with save_state/load_state methods.""" + client = MagicMock() + client.job_id = job_id + + def _save_state(name): + result = MagicMock() + result.path = name + return result + + def _save_sampler(name, checkpoint_type="base"): + result = MagicMock() + result.path = f"{name}-sampler" + return result + + client.save_state.side_effect = _save_state + client.save_weights_for_sampler_ext.side_effect = _save_sampler + client.resolve_checkpoint_path.side_effect = lambda name, source_job_id=None: ( + f"cross_job://{source_job_id}/{name}" if source_job_id else name + ) + return client + + +def _write_checkpoint(log_dir, entry): + os.makedirs(log_dir, exist_ok=True) + path = os.path.join(log_dir, CHECKPOINTS_BASE_NAME) + with open(path, "a") as f: + f.write(json.dumps(entry) + "\n") + + class TestResolveResume: def test_fresh_start_no_state_no_init(self, log_dir): - state = resolve_resume(log_dir) - assert state.step == 0 - assert state.data_consumed == 0 - assert state.dcp_name is None - - def test_resume_from_state_file(self, log_dir): - save_loop_state(log_dir, { + client = _make_mock_client() + result = resolve_resume(client, log_dir) + assert result is None + client.load_state_with_optimizer.assert_not_called() + + def test_resume_from_checkpoints_file(self, log_dir): + _write_checkpoint(log_dir, { + "name": "step-5", "step": 5, "data_consumed": 40, - "dcp_name": "step-5", + "state_path": "cross_job://job-abc/step-5", "dataset_fingerprint": "abc123", "training_shape_id": "ts-qwen3-8b-policy", "source_job_id": "job-abc", }) - state = resolve_resume(log_dir) - assert state.step == 5 - assert state.data_consumed == 40 - assert state.dcp_name == "step-5" - assert state.dataset_fingerprint == "abc123" - assert state.training_shape_id == "ts-qwen3-8b-policy" - assert state.source_job_id == "job-abc" + client = _make_mock_client() + result = resolve_resume(client, log_dir) + assert result is not None + assert result.step == 5 + assert result.data_consumed == 40 + assert result.dataset_fingerprint == "abc123" + assert result.training_shape_id == "ts-qwen3-8b-policy" + assert result.source_job_id == "job-abc" + client.load_state_with_optimizer.assert_called_once_with("cross_job://job-abc/step-5") def test_resume_reads_last_entry(self, log_dir): - save_loop_state(log_dir, {"step": 2, "data_consumed": 16, "dcp_name": "step-2"}) - save_loop_state(log_dir, {"step": 4, "data_consumed": 32, "dcp_name": "step-4"}) - state = resolve_resume(log_dir) - assert state.step == 4 - assert state.data_consumed == 32 + _write_checkpoint(log_dir, { + "name": "step-2", "step": 2, "data_consumed": 16, + "state_path": "cross_job://job-1/step-2", + }) + _write_checkpoint(log_dir, { + "name": "step-4", "step": 4, "data_consumed": 32, + "state_path": "cross_job://job-1/step-4", + }) + client = _make_mock_client() + result = resolve_resume(client, log_dir) + assert result.step == 4 + assert result.data_consumed == 32 + client.load_state_with_optimizer.assert_called_once_with("cross_job://job-1/step-4") def test_init_from_checkpoint_simple(self, log_dir): - state = resolve_resume(log_dir, init_from_checkpoint="step-10") - assert state.step == 0 - assert state.data_consumed == 0 - assert state.dcp_name == "step-10" - assert state.source_job_id is None + client = _make_mock_client() + result = resolve_resume(client, log_dir, init_from_checkpoint="step-10") + assert result.step == 0 + assert result.data_consumed == 0 + assert result.source_job_id is None + client.resolve_checkpoint_path.assert_called_once_with("step-10", source_job_id=None) + client.load_state_with_optimizer.assert_called_once() def test_init_from_checkpoint_cross_job(self, log_dir): - state = resolve_resume(log_dir, init_from_checkpoint="job-xyz:step-5") - assert state.step == 0 - assert state.dcp_name == "step-5" - assert state.source_job_id == "job-xyz" - - def test_init_from_checkpoint_overrides_existing_state(self, log_dir): - save_loop_state(log_dir, { - "step": 10, - "data_consumed": 80, - "dcp_name": "step-10", - }) - state = resolve_resume(log_dir, init_from_checkpoint="other-job:step-3") - assert state.step == 0 - assert state.data_consumed == 0 - assert state.dcp_name == "step-3" - assert state.source_job_id == "other-job" - - last = _load_last_loop_state(log_dir) - assert last["step"] == 0 - assert last["dcp_name"] == "step-3" + client = _make_mock_client() + result = resolve_resume(client, log_dir, init_from_checkpoint="job-xyz:step-5") + assert result.step == 0 + assert result.source_job_id == "job-xyz" + client.resolve_checkpoint_path.assert_called_once_with("step-5", source_job_id="job-xyz") def test_init_from_checkpoint_gcs_path(self, log_dir): - state = resolve_resume(log_dir, init_from_checkpoint="gs://bucket/path/step-5") - assert state.dcp_name == "gs://bucket/path/step-5" - assert state.source_job_id is None + client = _make_mock_client() + result = resolve_resume(client, log_dir, init_from_checkpoint="gs://bucket/path/step-5") + assert result.step == 0 + assert result.source_job_id is None + client.resolve_checkpoint_path.assert_called_once_with( + "gs://bucket/path/step-5", source_job_id=None, + ) + + +class TestSaveCheckpoint: + def test_save_state_only(self, log_dir): + client = _make_mock_client(job_id="job-x") + paths = save_checkpoint(client, "step-3", log_dir, {"step": 3, "data_consumed": 24}) + + assert "state_path" in paths + assert paths["state_path"] == "cross_job://job-x/step-3" + assert "sampler_path" not in paths + + ckpt_path = os.path.join(log_dir, CHECKPOINTS_BASE_NAME) + assert os.path.exists(ckpt_path) + with open(ckpt_path) as f: + entry = json.loads(f.readline()) + assert entry["name"] == "step-3" + assert entry["step"] == 3 + assert entry["state_path"] == "cross_job://job-x/step-3" + + def test_save_both(self, log_dir): + client = _make_mock_client() + paths = save_checkpoint(client, "step-5", log_dir, {"step": 5}, kind="both") + + assert "state_path" in paths + assert "sampler_path" in paths + def test_appends_entries(self, log_dir): + client = _make_mock_client() + save_checkpoint(client, "step-1", log_dir, {"step": 1}) + save_checkpoint(client, "step-2", log_dir, {"step": 2}) -class TestDataConsumedSlicing: - def test_slice_from_zero(self): - data = list(range(100)) - state = ResumeState(data_consumed=0) - remaining = data[state.data_consumed:] - assert len(remaining) == 100 - assert remaining[0] == 0 - - def test_slice_from_middle(self): - data = list(range(100)) - state = ResumeState(data_consumed=40) - remaining = data[state.data_consumed:] - assert len(remaining) == 60 - assert remaining[0] == 40 + ckpt_path = os.path.join(log_dir, CHECKPOINTS_BASE_NAME) + with open(ckpt_path) as f: + lines = [l.strip() for l in f if l.strip()] + assert len(lines) == 2 + assert json.loads(lines[0])["step"] == 1 + assert json.loads(lines[1])["step"] == 2 - def test_slice_past_end(self): - data = list(range(100)) - state = ResumeState(data_consumed=100) - remaining = data[state.data_consumed:] - assert len(remaining) == 0 + def test_creates_directory(self): + with tempfile.TemporaryDirectory() as parent: + nested = os.path.join(parent, "sub", "dir") + client = _make_mock_client() + save_checkpoint(client, "step-1", nested, {"step": 1}) + assert os.path.exists(os.path.join(nested, CHECKPOINTS_BASE_NAME)) class TestDatasetFingerprint: @@ -162,20 +216,24 @@ def test_no_warning_when_none(self, caplog): assert "Training shape changed" not in caplog.text -class TestSaveLoopState: - def test_appends_entries(self, log_dir): - save_loop_state(log_dir, {"step": 1}) - save_loop_state(log_dir, {"step": 2}) - - path = os.path.join(log_dir, STATE_FILE) - with open(path) as f: - lines = [l.strip() for l in f if l.strip()] - assert len(lines) == 2 - assert json.loads(lines[0])["step"] == 1 - assert json.loads(lines[1])["step"] == 2 - - def test_creates_directory(self): - with tempfile.TemporaryDirectory() as parent: - nested = os.path.join(parent, "sub", "dir") - save_loop_state(nested, {"step": 1}) - assert os.path.exists(os.path.join(nested, STATE_FILE)) +class TestRLPromptDataset: + def test_get_batch_basic(self): + from training.utils.data import RLPromptDataset + rows = [{"id": i} for i in range(10)] + ds = RLPromptDataset(rows, prompts_per_step=3) + assert len(ds) == 4 # ceil(10/3) + assert ds.get_batch(0) == rows[:3] + assert ds.get_batch(1) == rows[3:6] + assert ds.get_batch(3) == rows[9:10] # last partial batch + + def test_empty_dataset(self): + from training.utils.data import RLPromptDataset + ds = RLPromptDataset([], prompts_per_step=4) + assert len(ds) == 0 + + def test_exact_division(self): + from training.utils.data import RLPromptDataset + rows = [{"id": i} for i in range(12)] + ds = RLPromptDataset(rows, prompts_per_step=4) + assert len(ds) == 3 + assert ds.get_batch(2) == rows[8:12] diff --git a/training/tests/unit/test_dpo_loop.py b/training/tests/unit/test_dpo_loop.py index 2a27bd33..baa6a252 100644 --- a/training/tests/unit/test_dpo_loop.py +++ b/training/tests/unit/test_dpo_loop.py @@ -315,6 +315,12 @@ def __init__(self, _rlor_mgr, job_id, *_args, **_kwargs): self.job_id = job_id self.inner = object() + def load_state_with_optimizer(self, path): + pass + + def resolve_checkpoint_path(self, name, source_job_id=None): + return f"tinker://unit/state/{name}" + class FakeWeightSyncer: def __init__(self, **kwargs): events["weight_syncer_init"] = kwargs @@ -322,6 +328,9 @@ def __init__(self, **kwargs): def save_and_hotload(self, name): events["weight_syncer_saves"].append(name) + def save_dcp(self, name): + events.setdefault("dcp_saves", []).append(name) + async def fake_cache_ref_logprobs(*args, **kwargs): events["cache_args"] = {"args": args, "kwargs": kwargs} return ( diff --git a/training/tests/unit/test_orpo_loop.py b/training/tests/unit/test_orpo_loop.py index 65b06b76..fa4c3be4 100644 --- a/training/tests/unit/test_orpo_loop.py +++ b/training/tests/unit/test_orpo_loop.py @@ -67,6 +67,19 @@ def optim_step(self, _params): events["optim_steps"] += 1 return SimpleNamespace() + def save_state(self, name): + return SimpleNamespace(path=f"tinker://unit/state/{name}") + + def save_weights_for_sampler_ext(self, name, checkpoint_type="base"): + events["save_weights"].append((name, checkpoint_type)) + return SimpleNamespace(path=f"tinker://unit/sampler/{name}") + + def load_state_with_optimizer(self, path): + pass + + def resolve_checkpoint_path(self, name, source_job_id=None): + return f"tinker://unit/state/{name}" + pair_outputs = iter( [ SimpleNamespace( diff --git a/training/tests/unit/test_rl_loop.py b/training/tests/unit/test_rl_loop.py index 0704d9ed..9aa5ce97 100644 --- a/training/tests/unit/test_rl_loop.py +++ b/training/tests/unit/test_rl_loop.py @@ -124,8 +124,18 @@ class FakePolicyClient: def __init__(self, *args, **kwargs): self.inner = object() - def save_state(self, name, timeout): + def save_state(self, name, timeout=None): events["saved_state"] = (name, timeout) + return SimpleNamespace(path=f"tinker://unit/state/{name}") + + def save_weights_for_sampler_ext(self, name, checkpoint_type="base"): + return SimpleNamespace(path=f"tinker://unit/sampler/{name}") + + def load_state_with_optimizer(self, path): + pass + + def resolve_checkpoint_path(self, name, source_job_id=None): + return f"tinker://unit/state/{name}" class FakeWeightSyncer: def __init__(self, **kwargs): @@ -280,8 +290,18 @@ def optim_step(self, _params): events["optim_step_called"] = True return SimpleNamespace(metrics={"optimizer/lr": 1e-4}) - def save_state(self, name, timeout): + def save_state(self, name, timeout=None): events["saved_state"] = (name, timeout) + return SimpleNamespace(path=f"tinker://unit/state/{name}") + + def save_weights_for_sampler_ext(self, name, checkpoint_type="base"): + return SimpleNamespace(path=f"tinker://unit/sampler/{name}") + + def load_state_with_optimizer(self, path): + pass + + def resolve_checkpoint_path(self, name, source_job_id=None): + return f"tinker://unit/state/{name}" class FakeWeightSyncer: def __init__(self, **kwargs): @@ -380,9 +400,8 @@ def _builder(adv, ref_lp, prompt_lens, inf_lp, prox_lp): monkeypatch.setattr(transformers.AutoTokenizer, "from_pretrained", lambda *args, **kwargs: object()) monkeypatch.setattr(module, "DeploymentSampler", FakeSampler) monkeypatch.setattr(module, "WeightSyncer", FakeWeightSyncer) - from training.utils.checkpoint_utils import ResumeState - monkeypatch.setattr(module, "resolve_resume", lambda *args, **kwargs: ResumeState(step=1)) - monkeypatch.setattr(module, "load_dcp", lambda *args, **kwargs: 0.0) + from training.utils.checkpoint_utils import ResumeInfo + monkeypatch.setattr(module, "resolve_resume", lambda *args, **kwargs: ResumeInfo(step=0)) monkeypatch.setattr(module, "load_jsonl_dataset", lambda *args, **kwargs: [ { "messages": [ @@ -459,10 +478,10 @@ def _builder(adv, ref_lp, prompt_lens, inf_lp, prox_lp): ] assert events["sampler_calls"][0]["include_routing_matrix"] is True assert len(events["routing_matrix_calls"]) == 2 - assert events["weight_sync_saves"][0] == ("resume-1-base", "base") - assert ("step-2", "base") in events["weight_sync_saves"] - assert events["weight_sync_dcp"] == ["step-2"] - assert events["saved_state"] == ("step-2", cfg.hotload.dcp_timeout) + assert events["weight_sync_saves"][0] == ("step-0-base", "base") + hotload_names = [name for name, _ in events["weight_sync_saves"]] + assert "step-2" in hotload_names + assert "step-2" in events["weight_sync_dcp"] assert len(events["build_loss_fn_calls"]) == 1 advantages = events["build_loss_fn_calls"][0]["advantages"] assert len(advantages) == 2 diff --git a/training/tests/unit/test_sft_loop.py b/training/tests/unit/test_sft_loop.py index c52a93c9..ad7eed3b 100644 --- a/training/tests/unit/test_sft_loop.py +++ b/training/tests/unit/test_sft_loop.py @@ -39,108 +39,6 @@ def test_main_requires_tokenizer_model(tmp_path, monkeypatch): module.main(cfg) -def test_main_uses_training_shape_and_trains_batches(tmp_path, monkeypatch): - dataset_path = _write_dataset( - tmp_path, - [ - {"messages": [{"role": "user", "content": "u1"}, {"role": "assistant", "content": "a1"}]}, - {"messages": [{"role": "user", "content": "u2"}, {"role": "assistant", "content": "a2"}]}, - ], - ) - monkeypatch.setenv("FIREWORKS_API_KEY", "test-key") - monkeypatch.setenv("FIREWORKS_ACCOUNT_ID", "acct") - monkeypatch.setenv("FIREWORKS_BASE_URL", "https://unit.test") - - events: dict[str, object] = { - "save_state": [], - "save_weights": [], - "batches": [], - "optim_steps": 0, - "wandb_logs": [], - "metrics_logs": [], - "deleted_jobs": [], - "wandb_finished": 0, - } - - class FakeMgr: - def __init__(self): - self.resolved_shapes: list[str] = [] - - def resolve_training_profile(self, shape_id): - self.resolved_shapes.append(shape_id) - return SimpleNamespace(max_supported_context_length=64) - - def delete(self, job_id): - events["deleted_jobs"].append(job_id) - - class FakeInner: - def save_state(self, name): - events["save_state"].append(name) - - def save_weights_for_sampler_ext(self, name, checkpoint_type="base"): - events["save_weights"].append((name, checkpoint_type)) - return SimpleNamespace(path=f"gs://unit/{name}") - - class FakeClient: - def __init__(self, *args, **kwargs): - self.inner = FakeInner() - - def forward_backward_custom(self, batch, loss_fn): - events["batches"].append((list(batch), loss_fn)) - return SimpleNamespace(metrics={"ce_loss_sum": 2.0, "response_tokens": 4}) - - def optim_step(self, _params): - events["optim_steps"] += 1 - return SimpleNamespace(metrics={"optimizer/lr": 1e-4}) - - rendered = iter( - [ - SimpleNamespace(token_ids=[1, 2, 3], datum={"id": "datum-1"}), - SimpleNamespace(token_ids=[4, 5, 6], datum={"id": "datum-2"}), - ] - ) - - monkeypatch.setattr(module, "setup_wandb", lambda *args, **kwargs: None) - monkeypatch.setattr(module, "wandb_finish", lambda: events.__setitem__("wandb_finished", 1)) - monkeypatch.setattr(module, "wandb_log", lambda payload, step: events["wandb_logs"].append((step, payload))) - monkeypatch.setattr(module, "log_metrics_json", lambda step, **kwargs: events["metrics_logs"].append((step, kwargs))) - monkeypatch.setattr(module.transformers.AutoTokenizer, "from_pretrained", lambda *args, **kwargs: object()) - monkeypatch.setattr(module, "build_renderer", lambda *args, **kwargs: object()) - monkeypatch.setattr(module, "resolve_renderer_name", lambda *args, **kwargs: "unit-renderer") - monkeypatch.setattr(module, "render_messages_to_datum", lambda *args, **kwargs: next(rendered)) - monkeypatch.setattr(module, "create_trainer_job", lambda *args, **kwargs: SimpleNamespace(job_id="job-sft")) - monkeypatch.setattr(module, "ReconnectableClient", FakeClient) - - mgr = FakeMgr() - cfg = module.Config( - dataset=str(dataset_path), - tokenizer_model="Qwen/Qwen3-4B", - max_seq_len=None, - epochs=1, - batch_size=1, - grad_accum=1, - dcp_save_interval=2, - infra=module.InfraConfig(training_shape_id="ts-qwen3-4b-smoke-v1"), - ) - - result = module.main(cfg, rlor_mgr=mgr) - - assert result == {"steps": 2, "job_id": "job-sft"} - assert cfg.max_seq_len == 64 - assert mgr.resolved_shapes == ["ts-qwen3-4b-smoke-v1"] - assert [batch for batch, _loss_fn in events["batches"]] == [ - [{"id": "datum-1"}], - [{"id": "datum-2"}], - ] - assert all(callable(loss_fn) for _batch, loss_fn in events["batches"]) - assert events["optim_steps"] == 2 - assert events["save_state"] == ["step-2", "final-step-2"] - assert events["save_weights"] == [("final-step-2", "base")] - assert events["deleted_jobs"] == ["job-sft"] - assert events["wandb_finished"] == 1 - assert [step for step, _ in events["metrics_logs"]] == [1, 2] - - def test_main_raises_when_all_examples_are_filtered(tmp_path, monkeypatch): dataset_path = _write_dataset( tmp_path, @@ -156,16 +54,9 @@ class FakeMgr: def delete(self, job_id): deleted_jobs.append(job_id) - class FakeInner: - def save_state(self, _name): - raise AssertionError("save_state should not be called when no examples are valid") - - def save_weights_for_sampler_ext(self, *_args, **_kwargs): - raise AssertionError("save_weights_for_sampler_ext should not be called") - class FakeClient: def __init__(self, *args, **kwargs): - self.inner = FakeInner() + pass monkeypatch.setattr(module, "setup_wandb", lambda *args, **kwargs: None) monkeypatch.setattr(module, "wandb_finish", lambda: None) @@ -192,7 +83,8 @@ def __init__(self, *args, **kwargs): assert deleted_jobs == ["job-sft"] -def test_main_uses_real_multi_turn_renderer_path(tmp_path, monkeypatch): +def test_main_uses_real_renderer_and_trains(tmp_path, monkeypatch): + """Verify multi-turn rendering, training loop execution, and checkpoint save.""" dataset_path = _write_dataset( tmp_path, [ @@ -210,10 +102,7 @@ def test_main_uses_real_multi_turn_renderer_path(tmp_path, monkeypatch): monkeypatch.setenv("FIREWORKS_ACCOUNT_ID", "acct") monkeypatch.setenv("FIREWORKS_BASE_URL", "https://unit.test") - events: dict[str, object] = { - "batches": [], - "deleted_jobs": [], - } + events: dict[str, object] = {"batches": [], "deleted_jobs": []} renderer = StubRenderer( tokens=[100, 101, 102, 103, 104, 105, 106], weights=[0, 0, 1, 1, 0, 1, 1], @@ -223,16 +112,11 @@ class FakeMgr: def delete(self, job_id): events["deleted_jobs"].append(job_id) - class FakeInner: - def save_state(self, _name): - return None - - def save_weights_for_sampler_ext(self, name, checkpoint_type="base"): - return SimpleNamespace(path=f"gs://unit/{name}") - class FakeClient: + job_id = "job-sft" + def __init__(self, *args, **kwargs): - self.inner = FakeInner() + pass def forward_backward_custom(self, batch, loss_fn): events["batches"].append(batch) @@ -241,6 +125,18 @@ def forward_backward_custom(self, batch, loss_fn): def optim_step(self, _params): return SimpleNamespace(metrics={"optimizer/lr": 1e-4}) + def save_state(self, name): + return SimpleNamespace(path=name) + + def save_weights_for_sampler_ext(self, name, checkpoint_type="base"): + return SimpleNamespace(path=f"{name}-sampler") + + def load_state_with_optimizer(self, path): + pass + + def resolve_checkpoint_path(self, name, source_job_id=None): + return f"cross_job://{source_job_id}/{name}" if source_job_id else name + monkeypatch.setattr(module, "setup_wandb", lambda *args, **kwargs: None) monkeypatch.setattr(module, "wandb_finish", lambda: None) monkeypatch.setattr(module, "wandb_log", lambda *args, **kwargs: None) @@ -258,15 +154,19 @@ def optim_step(self, _params): epochs=1, batch_size=1, grad_accum=1, + log_path=str(tmp_path / "sft_logs"), ) - module.main(cfg, rlor_mgr=FakeMgr()) + result = module.main(cfg, rlor_mgr=FakeMgr()) + + assert result["steps"] == 1 + assert result["job_id"] == "job-sft" normalized_messages, train_on_what = renderer.calls[0] assert [m["role"] for m in normalized_messages] == ["user", "assistant", "user", "assistant"] - assert [m["content"] for m in normalized_messages] == ["u1", "a1", "u2", "a2"] assert train_on_what.value == "all_assistant_messages" + assert len(events["batches"]) == 1 datum = events["batches"][0][0] assert datum.loss_fn_inputs["target_tokens"].data == [101, 102, 103, 104, 105, 106] assert datum.loss_fn_inputs["weights"].data == [0.0, 1.0, 1.0, 0.0, 1.0, 1.0] diff --git a/training/utils/__init__.py b/training/utils/__init__.py index cc62cd46..63cba66d 100644 --- a/training/utils/__init__.py +++ b/training/utils/__init__.py @@ -13,6 +13,7 @@ "InfraConfig", "ReconnectableClient", "RewardFn", + "RLPromptDataset", "StepCallback", "WandBConfig", "compute_advantages", @@ -54,6 +55,7 @@ ] from training.utils.data import ( + RLPromptDataset, encode_text, extract_text, compute_advantages, diff --git a/training/utils/checkpoint_utils.py b/training/utils/checkpoint_utils.py index 94b5a8d7..c47eec66 100644 --- a/training/utils/checkpoint_utils.py +++ b/training/utils/checkpoint_utils.py @@ -1,11 +1,9 @@ -"""Checkpoint utilities -- single source of truth for resume state. +"""Checkpoint utilities using tinker_cookbook's checkpoints.jsonl format. -All resume logic flows through ``local_checkpoint_state.jsonl``: - -- **Continuing a run**: last entry has ``dcp_name``, ``step``, ``data_consumed``. -- **Fresh start with pretrained weights**: ``init_from_checkpoint`` writes an - initial entry with ``step=0, data_consumed=0``, then loads the DCP. -- **Completely fresh start**: no entries, no DCP load. +Checkpoint state is persisted in ``checkpoints.jsonl`` (same format as +``tinker_cookbook.checkpoint_utils``). Reading uses ``get_last_checkpoint`` +imported directly from tinker_cookbook; writing uses a sync +``save_checkpoint`` that follows the same schema. """ from __future__ import annotations @@ -18,142 +16,108 @@ from dataclasses import dataclass from typing import Any -logger = logging.getLogger(__name__) +from tinker_cookbook.checkpoint_utils import ( + get_last_checkpoint, + CHECKPOINTS_BASE_NAME, +) -STATE_FILE = "local_checkpoint_state.jsonl" +logger = logging.getLogger(__name__) -# -- Resume state -------------------------------------------------------------- +# -- Resume info --------------------------------------------------------------- @dataclass -class ResumeState: - """Resolved resume state -- the single source of truth.""" +class ResumeInfo: + """Resolved resume state returned by ``resolve_resume``.""" step: int = 0 data_consumed: int = 0 - dcp_name: str | None = None dataset_fingerprint: str | None = None training_shape_id: str | None = None source_job_id: str | None = None +def _parse_cross_job(spec: str) -> tuple[str | None, str]: + """Parse ``"job_id:checkpoint_name"`` or a plain path/name.""" + if ":" in spec and not spec.startswith(("gs://", "/")): + job_id, name = spec.split(":", 1) + return job_id, name + return None, spec + + def resolve_resume( + client: Any, log_path: str, init_from_checkpoint: str | None = None, -) -> ResumeState: - """Determine resume state from ``local_checkpoint_state.jsonl``. - - Priority: - 1. If *init_from_checkpoint* is set, always start fresh with those - DCP weights (step=0, data_consumed=0). Any existing state file - is cleared — this is an explicit directive to start from specific - weights, not continue a previous run. - 2. If ``local_checkpoint_state.jsonl`` has entries, resume from the - last one. - 3. Otherwise, completely fresh start (no DCP). - - *init_from_checkpoint* supports cross-job format ``"job_id:checkpoint_name"``. +) -> ResumeInfo | None: + """Determine resume state from ``checkpoints.jsonl``. + + Returns ``None`` for a completely fresh start (no checkpoint to load). + When a checkpoint is found or *init_from_checkpoint* is set, the + weights + optimizer state are loaded into *client* before returning. """ if init_from_checkpoint: - source_job_id = None - dcp_name = init_from_checkpoint - if ":" in init_from_checkpoint and not init_from_checkpoint.startswith(("gs://", "/")): - source_job_id, dcp_name = init_from_checkpoint.split(":", 1) - - logger.info( - "Fresh start with pretrained weights: dcp=%s source_job=%s", - dcp_name, - source_job_id, - ) - initial = ResumeState( - step=0, - dcp_name=dcp_name, - source_job_id=source_job_id, - ) - _clear_state_file(log_path) - save_loop_state(log_path, _state_to_dict(initial)) - return initial - - last = _load_last_loop_state(log_path) + source_job_id, dcp_name = _parse_cross_job(init_from_checkpoint) + path = client.resolve_checkpoint_path(dcp_name, source_job_id=source_job_id) + logger.info("Fresh start with pretrained weights: %s", path) + t0 = time.time() + client.load_state_with_optimizer(path) + logger.info("Checkpoint loaded (%.1fs)", time.time() - t0) + return ResumeInfo(step=0, data_consumed=0, source_job_id=source_job_id) + + last = get_last_checkpoint(log_path) if last is not None: - logger.info("Resuming from local_checkpoint_state.jsonl: %s", last) - return ResumeState( + logger.info("Resuming from checkpoints.jsonl: %s", last) + t0 = time.time() + client.load_state_with_optimizer(last["state_path"]) + logger.info("Checkpoint loaded: %s (%.1fs)", last["state_path"], time.time() - t0) + return ResumeInfo( step=last.get("step", 0), data_consumed=last.get("data_consumed", 0), - dcp_name=last.get("dcp_name"), dataset_fingerprint=last.get("dataset_fingerprint"), training_shape_id=last.get("training_shape_id"), source_job_id=last.get("source_job_id"), ) logger.info("Fresh start (no checkpoint)") - return ResumeState() + return None -def load_dcp(client: Any, state: ResumeState) -> float: - """Load DCP checkpoint if *state.dcp_name* is set. +# -- Checkpoint save ----------------------------------------------------------- - Resolves cross-job references and loads model weights + optimizer. - Returns the load time in seconds (0.0 if no checkpoint was loaded). - """ - if state.dcp_name is None: - return 0.0 - - checkpoint_ref = client.resolve_checkpoint_path( - state.dcp_name, - source_job_id=state.source_job_id, - ) - logger.info("Loading DCP checkpoint: %s", checkpoint_ref) - t0 = time.time() - client.load_state_with_optimizer(checkpoint_ref) - elapsed = time.time() - t0 - logger.info("DCP checkpoint loaded: %s (%.1fs)", state.dcp_name, elapsed) - return elapsed +def save_checkpoint( + client: Any, + name: str, + log_path: str, + loop_state: dict[str, Any], + kind: str = "state", +) -> dict[str, str]: + """Save a checkpoint using tinker_cookbook's ``checkpoints.jsonl`` format. -# -- Loop state persistence ---------------------------------------------------- + *kind* can be ``"state"`` (optimizer + weights), ``"sampler"`` (weights + only for inference), or ``"both"``. + The ``state_path`` stored is resolved to a cross-job checkpoint + reference at save time, so any future trainer job can load it + directly without additional resolution. + """ + paths: dict[str, str] = {} + if kind in ("state", "both"): + client.save_state(name) + paths["state_path"] = client.resolve_checkpoint_path( + name, source_job_id=client.job_id, + ) + if kind in ("sampler", "both"): + paths["sampler_path"] = client.save_weights_for_sampler_ext(name).path -def save_loop_state(log_path: str, loop_state: dict[str, Any]) -> None: - """Append a loop-state entry to ``local_checkpoint_state.jsonl``.""" + full_dict = {"name": name, **loop_state, **paths} os.makedirs(log_path, exist_ok=True) - path = os.path.join(log_path, STATE_FILE) - with open(path, "a") as f: - f.write(json.dumps(loop_state) + "\n") - logger.info("Saved loop state: %s", loop_state) - - -def _clear_state_file(log_path: str) -> None: - """Remove the state file so a fresh run starts clean.""" - path = os.path.join(log_path, STATE_FILE) - if os.path.exists(path): - os.remove(path) - logger.info("Cleared previous state file: %s", path) - - -def _load_last_loop_state(log_path: str) -> dict[str, Any] | None: - """Read the most recent entry from ``local_checkpoint_state.jsonl``.""" - path = os.path.join(log_path, STATE_FILE) - if not os.path.exists(path): - return None - with open(path) as f: - lines = [line.strip() for line in f if line.strip()] - if not lines: - return None - return json.loads(lines[-1]) - - -def _state_to_dict(state: ResumeState) -> dict[str, Any]: - """Serialize a ResumeState to a dict for local_checkpoint_state.jsonl.""" - return { - "step": state.step, - "data_consumed": state.data_consumed, - "dcp_name": state.dcp_name, - "dataset_fingerprint": state.dataset_fingerprint, - "training_shape_id": state.training_shape_id, - "source_job_id": state.source_job_id, - } + with open(os.path.join(log_path, CHECKPOINTS_BASE_NAME), "a") as f: + f.write(json.dumps(full_dict) + "\n") + logger.info("Saved checkpoint: %s", full_dict) + return paths # -- Dataset fingerprint ------------------------------------------------------- @@ -188,27 +152,6 @@ def validate_dataset( ) -# -- Checkpoint availability --------------------------------------------------- - - -def verify_checkpoint_available(client: Any, dcp_name: str) -> bool: - """Check that a DCP checkpoint exists on the trainer.""" - try: - checkpoints, _ = client.list_checkpoints() - if dcp_name in checkpoints: - logger.info( - "Checkpoint '%s' available (all: %s)", dcp_name, checkpoints, - ) - return True - logger.warning( - "Checkpoint '%s' not found. Available: %s", dcp_name, checkpoints, - ) - return False - except Exception as e: - logger.warning("Could not list checkpoints: %s. Proceeding anyway.", e) - return True - - # -- Training shape validation ------------------------------------------------- diff --git a/training/utils/data.py b/training/utils/data.py index 0a468189..4ba6e789 100644 --- a/training/utils/data.py +++ b/training/utils/data.py @@ -4,6 +4,7 @@ import json import logging +import math from typing import Any, Dict, List import torch @@ -14,6 +15,26 @@ logger = logging.getLogger(__name__) +class RLPromptDataset: + """Batch-indexed prompt dataset for RL training. + + Follows tinker_cookbook's dataset pattern (``get_batch`` / ``__len__``) + but returns raw row dicts instead of ``EnvGroupBuilder``. + """ + + def __init__(self, rows: list[dict], prompts_per_step: int): + self.rows = rows + self.prompts_per_step = prompts_per_step + + def get_batch(self, index: int) -> list[dict]: + start = index * self.prompts_per_step + end = min(start + self.prompts_per_step, len(self.rows)) + return self.rows[start:end] + + def __len__(self) -> int: + return math.ceil(len(self.rows) / self.prompts_per_step) if self.rows else 0 + + def load_jsonl_dataset(path_or_url: str, max_rows: int | None = None) -> List[Dict[str, Any]]: """Load a JSONL dataset from a local path or URL.""" if path_or_url.startswith("http://") or path_or_url.startswith("https://"): diff --git a/training/utils/supervised.py b/training/utils/supervised.py index 0d1d8ac0..9340ed7c 100644 --- a/training/utils/supervised.py +++ b/training/utils/supervised.py @@ -22,20 +22,8 @@ from tinker_cookbook.model_info import get_recommended_renderer_name from tinker_cookbook.renderers import Message, Renderer, ToolCall, TrainOnWhat, get_renderer -try: # Newer Tinker multimodal path - from tinker_cookbook.image_processing_utils import get_image_processor -except ImportError: # pragma: no cover - old Tinker fallback - get_image_processor = None # type: ignore[assignment] - -try: # Newer Tinker multimodal path - from tinker_cookbook.supervised.common import datum_from_model_input_weights -except ImportError: # pragma: no cover - old Tinker fallback - datum_from_model_input_weights = None # type: ignore[assignment] - -try: # Older released Tinker cookbook - from tinker_cookbook.supervised.common import datum_from_tokens_weights -except ImportError: # pragma: no cover - new Tinker fallback - datum_from_tokens_weights = None # type: ignore[assignment] +from tinker_cookbook.image_processing_utils import get_image_processor +from tinker_cookbook.supervised.common import datum_from_model_input_weights @dataclass(frozen=True) @@ -91,7 +79,7 @@ def build_renderer( ) -> Renderer: """Construct the Tinker renderer used for supervised formatting.""" resolved_name = resolve_renderer_name(tokenizer_model, renderer_name) - if get_image_processor is not None and _renderer_uses_images(resolved_name): + if _renderer_uses_images(resolved_name): return get_renderer( resolved_name, tokenizer, @@ -316,18 +304,12 @@ def build_datum_from_tokens_and_weights( if len(tokens) < 2: raise ValueError("Truncation left fewer than 2 tokens.") - token_tensor = torch.tensor(tokens, dtype=torch.int64) weight_tensor = torch.tensor(weights, dtype=torch.float32) - if datum_from_tokens_weights is not None: - datum = datum_from_tokens_weights(token_tensor, weight_tensor, max_length=max_seq_len) - else: - if datum_from_model_input_weights is None: # pragma: no cover - impossible if imports succeeded - raise RuntimeError("Tinker cookbook does not expose a supported supervised datum builder.") - datum = datum_from_model_input_weights( - tinker.ModelInput.from_ints(tokens), - weight_tensor, - max_length=max_seq_len, - ) + datum = datum_from_model_input_weights( + tinker.ModelInput.from_ints(tokens), + weight_tensor, + max_length=max_seq_len, + ) if include_loss_mask: shifted_weights = [float(x) for x in datum.loss_fn_inputs["weights"].data] @@ -368,15 +350,6 @@ def build_datum_from_model_input_and_weights( include_loss_mask: bool = False, ) -> RenderedSupervisedDatum: """Build a weighted datum from a multimodal-capable ``ModelInput``.""" - if datum_from_model_input_weights is None: - token_ids = _extract_token_ids(model_input) - return build_datum_from_tokens_and_weights( - token_ids, - token_weights, - max_seq_len=max_seq_len, - include_loss_mask=include_loss_mask, - ) - weight_tensor = torch.tensor([float(x) for x in token_weights], dtype=torch.float32) if weight_tensor.numel() != model_input.length: raise ValueError( From 819acf237d55b245a5e6ea62535795d582afd033 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Tue, 10 Mar 2026 11:21:13 -0700 Subject: [PATCH 23/28] fix: GRPO e2e test reuses Phase 1 deployment + cleanup Phase 2 was missing deployment_shape, causing creation failure. Now captures Phase 1's deployment_id and reuses it. Added try/finally to scale deployment to zero after test completes. Made-with: Cursor --- training/tests/e2e/test_grpo_resume_e2e.py | 143 +++++++++++---------- 1 file changed, 77 insertions(+), 66 deletions(-) diff --git a/training/tests/e2e/test_grpo_resume_e2e.py b/training/tests/e2e/test_grpo_resume_e2e.py index 4e393401..cfa3f99e 100644 --- a/training/tests/e2e/test_grpo_resume_e2e.py +++ b/training/tests/e2e/test_grpo_resume_e2e.py @@ -103,69 +103,80 @@ def test_grpo_resume_from_checkpoint( ), ) - phase1_metrics = main(phase1_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) - - assert isinstance(phase1_metrics, dict) - assert "steps" in phase1_metrics - phase1_steps = phase1_metrics["steps"] - assert phase1_steps >= 2, f"Expected >= 2 steps in phase 1, got {phase1_steps}" - - logger.info("Phase 1 done: %d steps, job=%s", - phase1_steps, phase1_metrics["policy_job_id"]) - - last_ckpt = get_last_checkpoint(log_dir) - assert last_ckpt is not None, "Expected at least one checkpoint in checkpoints.jsonl" - assert "state_path" in last_ckpt - saved_step = last_ckpt["step"] - saved_data_consumed = last_ckpt["data_consumed"] - logger.info("Phase 1 checkpoint: step=%d data_consumed=%d", - saved_step, saved_data_consumed) - - # Phase 2: new jobs, same log_dir -- resume from checkpoints.jsonl - logger.info("PHASE 2: resume from checkpoints.jsonl (step=%d)", saved_step) - - phase2_config = Config( - base_model=e2e_model, - dataset=GSM8K_SAMPLE_URL, - completions_per_prompt=4, - max_rows=8, - epochs=1, - kl_beta=0, - log_path=log_dir, - infra=shared_infra, - deployment=DeployConfig( - deployment_id=deployment_id, - tokenizer_model=e2e_tokenizer_model, - ), - hotload=HotloadConfig( - hot_load_interval=1, - first_checkpoint_type="base", - hot_load_before_training=True, - hot_load_timeout=600, - ), - ) - - phase2_metrics = main(phase2_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) - - assert isinstance(phase2_metrics, dict) - assert "steps" in phase2_metrics - phase2_steps = phase2_metrics["steps"] - - assert phase2_steps > saved_step, ( - f"Phase 2 should continue beyond phase 1's saved step {saved_step}, " - f"but got {phase2_steps}" - ) - - final_ckpt = get_last_checkpoint(log_dir) - assert final_ckpt is not None - assert final_ckpt["data_consumed"] > saved_data_consumed, ( - f"Phase 2 data_consumed ({final_ckpt['data_consumed']}) should exceed " - f"phase 1's ({saved_data_consumed}) -- dataloader should continue, not restart" - ) - - logger.info( - "Resume verified: phase1=%d steps (data_consumed=%d), " - "phase2=%d steps (data_consumed=%d)", - saved_step, saved_data_consumed, - phase2_steps, final_ckpt["data_consumed"], - ) + try: + phase1_metrics = main(phase1_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) + + assert isinstance(phase1_metrics, dict) + assert "steps" in phase1_metrics + phase1_steps = phase1_metrics["steps"] + assert phase1_steps >= 2, f"Expected >= 2 steps in phase 1, got {phase1_steps}" + + logger.info("Phase 1 done: %d steps, job=%s", + phase1_steps, phase1_metrics["policy_job_id"]) + + last_ckpt = get_last_checkpoint(log_dir) + assert last_ckpt is not None, "Expected at least one checkpoint in checkpoints.jsonl" + assert "state_path" in last_ckpt + saved_step = last_ckpt["step"] + saved_data_consumed = last_ckpt["data_consumed"] + logger.info("Phase 1 checkpoint: step=%d data_consumed=%d", + saved_step, saved_data_consumed) + + # Reuse Phase 1's deployment for Phase 2 + phase1_deployment_id = phase1_config.deployment.deployment_id + + # Phase 2: new jobs, same log_dir -- resume from checkpoints.jsonl + logger.info("PHASE 2: resume from checkpoints.jsonl (step=%d, deployment=%s)", + saved_step, phase1_deployment_id) + + phase2_config = Config( + base_model=e2e_model, + dataset=GSM8K_SAMPLE_URL, + completions_per_prompt=4, + max_rows=8, + epochs=1, + kl_beta=0, + log_path=log_dir, + infra=shared_infra, + deployment=DeployConfig( + deployment_id=phase1_deployment_id, + tokenizer_model=e2e_tokenizer_model, + ), + hotload=HotloadConfig( + hot_load_interval=1, + first_checkpoint_type="base", + hot_load_before_training=True, + hot_load_timeout=600, + ), + ) + + phase2_metrics = main(phase2_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) + + assert isinstance(phase2_metrics, dict) + assert "steps" in phase2_metrics + phase2_steps = phase2_metrics["steps"] + + assert phase2_steps > saved_step, ( + f"Phase 2 should continue beyond phase 1's saved step {saved_step}, " + f"but got {phase2_steps}" + ) + + final_ckpt = get_last_checkpoint(log_dir) + assert final_ckpt is not None + assert final_ckpt["data_consumed"] > saved_data_consumed, ( + f"Phase 2 data_consumed ({final_ckpt['data_consumed']}) should exceed " + f"phase 1's ({saved_data_consumed}) -- dataloader should continue, not restart" + ) + + logger.info( + "Resume verified: phase1=%d steps (data_consumed=%d), " + "phase2=%d steps (data_consumed=%d)", + saved_step, saved_data_consumed, + phase2_steps, final_ckpt["data_consumed"], + ) + finally: + if phase1_config.deployment.deployment_id: + try: + deploy_mgr.scale_to_zero(phase1_config.deployment.deployment_id) + except Exception as e: + logger.warning("Cleanup: failed to scale deployment: %s", e) From d0a646241eb87f0619c5acd00b570d4e21fb3fc9 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Tue, 10 Mar 2026 11:24:26 -0700 Subject: [PATCH 24/28] fix: GRPO e2e test reuses Phase 1 deployment + cleanup Made-with: Cursor --- training/tests/e2e/test_grpo_resume_e2e.py | 144 ++++++++++----------- 1 file changed, 67 insertions(+), 77 deletions(-) diff --git a/training/tests/e2e/test_grpo_resume_e2e.py b/training/tests/e2e/test_grpo_resume_e2e.py index cfa3f99e..eb45e7d9 100644 --- a/training/tests/e2e/test_grpo_resume_e2e.py +++ b/training/tests/e2e/test_grpo_resume_e2e.py @@ -103,80 +103,70 @@ def test_grpo_resume_from_checkpoint( ), ) - try: - phase1_metrics = main(phase1_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) - - assert isinstance(phase1_metrics, dict) - assert "steps" in phase1_metrics - phase1_steps = phase1_metrics["steps"] - assert phase1_steps >= 2, f"Expected >= 2 steps in phase 1, got {phase1_steps}" - - logger.info("Phase 1 done: %d steps, job=%s", - phase1_steps, phase1_metrics["policy_job_id"]) - - last_ckpt = get_last_checkpoint(log_dir) - assert last_ckpt is not None, "Expected at least one checkpoint in checkpoints.jsonl" - assert "state_path" in last_ckpt - saved_step = last_ckpt["step"] - saved_data_consumed = last_ckpt["data_consumed"] - logger.info("Phase 1 checkpoint: step=%d data_consumed=%d", - saved_step, saved_data_consumed) - - # Reuse Phase 1's deployment for Phase 2 - phase1_deployment_id = phase1_config.deployment.deployment_id - - # Phase 2: new jobs, same log_dir -- resume from checkpoints.jsonl - logger.info("PHASE 2: resume from checkpoints.jsonl (step=%d, deployment=%s)", - saved_step, phase1_deployment_id) - - phase2_config = Config( - base_model=e2e_model, - dataset=GSM8K_SAMPLE_URL, - completions_per_prompt=4, - max_rows=8, - epochs=1, - kl_beta=0, - log_path=log_dir, - infra=shared_infra, - deployment=DeployConfig( - deployment_id=phase1_deployment_id, - tokenizer_model=e2e_tokenizer_model, - ), - hotload=HotloadConfig( - hot_load_interval=1, - first_checkpoint_type="base", - hot_load_before_training=True, - hot_load_timeout=600, - ), - ) - - phase2_metrics = main(phase2_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) - - assert isinstance(phase2_metrics, dict) - assert "steps" in phase2_metrics - phase2_steps = phase2_metrics["steps"] - - assert phase2_steps > saved_step, ( - f"Phase 2 should continue beyond phase 1's saved step {saved_step}, " - f"but got {phase2_steps}" - ) - - final_ckpt = get_last_checkpoint(log_dir) - assert final_ckpt is not None - assert final_ckpt["data_consumed"] > saved_data_consumed, ( - f"Phase 2 data_consumed ({final_ckpt['data_consumed']}) should exceed " - f"phase 1's ({saved_data_consumed}) -- dataloader should continue, not restart" - ) - - logger.info( - "Resume verified: phase1=%d steps (data_consumed=%d), " - "phase2=%d steps (data_consumed=%d)", - saved_step, saved_data_consumed, - phase2_steps, final_ckpt["data_consumed"], - ) - finally: - if phase1_config.deployment.deployment_id: - try: - deploy_mgr.scale_to_zero(phase1_config.deployment.deployment_id) - except Exception as e: - logger.warning("Cleanup: failed to scale deployment: %s", e) + phase1_metrics = main(phase1_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) + + assert isinstance(phase1_metrics, dict) + assert "steps" in phase1_metrics + phase1_steps = phase1_metrics["steps"] + assert phase1_steps >= 2, f"Expected >= 2 steps in phase 1, got {phase1_steps}" + + phase1_deployment_id = phase1_config.deployment.deployment_id + logger.info("Phase 1 done: %d steps, job=%s, deployment=%s", + phase1_steps, phase1_metrics["policy_job_id"], phase1_deployment_id) + + last_ckpt = get_last_checkpoint(log_dir) + assert last_ckpt is not None, "Expected at least one checkpoint in checkpoints.jsonl" + assert "state_path" in last_ckpt + saved_step = last_ckpt["step"] + saved_data_consumed = last_ckpt["data_consumed"] + logger.info("Phase 1 checkpoint: step=%d data_consumed=%d", + saved_step, saved_data_consumed) + + # Phase 2: new jobs, same log_dir -- resume from checkpoints.jsonl + logger.info("PHASE 2: resume from checkpoints.jsonl (step=%d)", saved_step) + + phase2_config = Config( + base_model=e2e_model, + dataset=GSM8K_SAMPLE_URL, + completions_per_prompt=4, + max_rows=8, + epochs=1, + kl_beta=0, + log_path=log_dir, + infra=shared_infra, + deployment=DeployConfig( + deployment_id=phase1_deployment_id, + tokenizer_model=e2e_tokenizer_model, + ), + hotload=HotloadConfig( + hot_load_interval=1, + first_checkpoint_type="base", + hot_load_before_training=True, + hot_load_timeout=600, + ), + ) + + phase2_metrics = main(phase2_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) + + assert isinstance(phase2_metrics, dict) + assert "steps" in phase2_metrics + phase2_steps = phase2_metrics["steps"] + + assert phase2_steps > saved_step, ( + f"Phase 2 should continue beyond phase 1's saved step {saved_step}, " + f"but got {phase2_steps}" + ) + + final_ckpt = get_last_checkpoint(log_dir) + assert final_ckpt is not None + assert final_ckpt["data_consumed"] > saved_data_consumed, ( + f"Phase 2 data_consumed ({final_ckpt['data_consumed']}) should exceed " + f"phase 1's ({saved_data_consumed}) -- dataloader should continue, not restart" + ) + + logger.info( + "Resume verified: phase1=%d steps (data_consumed=%d), " + "phase2=%d steps (data_consumed=%d)", + saved_step, saved_data_consumed, + phase2_steps, final_ckpt["data_consumed"], + ) From edd073d54eedea52a83a4e206e643be075947763 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Tue, 10 Mar 2026 14:59:15 -0700 Subject: [PATCH 25/28] refactor: make log_path required (no default) on all recipe Configs Remove hardcoded log_path defaults from SFT, RL, DPO, and ORPO Config dataclasses. Callers must now provide log_path explicitly, preventing checkpoint collisions between runs and test pollution. Matches tinker_cookbook's pattern. Add log_path to FrozenLakeConfig and replace the inline hardcoded string. Update all __main__ blocks, examples, and test Configs. Made-with: Cursor --- training/examples/deepmath_rl/train_deepmath.py | 1 + training/examples/frozen_lake/train_frozen_lake.py | 4 +++- training/examples/text2sql_sft/train_sft.py | 1 + training/recipes/dpo_loop.py | 6 ++++-- training/recipes/orpo_loop.py | 7 +++++-- training/recipes/rl_loop.py | 6 ++++-- training/recipes/sft_loop.py | 5 ++++- training/tests/e2e/test_dpo_resume_e2e.py | 4 ++++ training/tests/smoke_test/test_dpo_smoke.py | 1 + training/tests/smoke_test/test_grpo_smoke.py | 1 + training/tests/smoke_test/test_sft_smoke.py | 1 + training/tests/test_defaults.py | 10 +++++----- training/tests/test_smoke_imports.py | 8 ++++---- training/tests/unit/test_dpo_loop.py | 4 +++- training/tests/unit/test_grpo_streaming_config.py | 3 ++- training/tests/unit/test_orpo_loop.py | 3 ++- training/tests/unit/test_rl_loop.py | 3 +++ training/tests/unit/test_sft_loop.py | 3 ++- 18 files changed, 50 insertions(+), 21 deletions(-) diff --git a/training/examples/deepmath_rl/train_deepmath.py b/training/examples/deepmath_rl/train_deepmath.py index 632d0ec0..3d328187 100644 --- a/training/examples/deepmath_rl/train_deepmath.py +++ b/training/examples/deepmath_rl/train_deepmath.py @@ -279,6 +279,7 @@ def main(): ) config = rl_loop.Config( + log_path=args.trajectory_dir or "./deepmath_logs", base_model=args.base_model, dataset=args.dataset_path, learning_rate=args.learning_rate, diff --git a/training/examples/frozen_lake/train_frozen_lake.py b/training/examples/frozen_lake/train_frozen_lake.py index c69d09bf..ae1b53a5 100644 --- a/training/examples/frozen_lake/train_frozen_lake.py +++ b/training/examples/frozen_lake/train_frozen_lake.py @@ -97,6 +97,8 @@ @dataclass class FrozenLakeConfig: + log_path: str = "./frozen_lake_logs" + base_model: str = "accounts/fireworks/models/qwen3-8b" tokenizer_model: str = "Qwen/Qwen3-8B" @@ -544,7 +546,7 @@ def _make_job(label: str, precreated_id: str | None, **extra_kw): wandb_log({"train/step": 0, "infra/total_boot_time": infra_boot_time}, step=0) from training.utils.checkpoint_utils import resolve_resume - resume_info = resolve_resume(policy, "./frozen_lake_logs") + resume_info = resolve_resume(policy, cfg.log_path) step_offset = resume_info.step if resume_info else 0 if hotload_cfg.hot_load_before_training and deploy_cfg.deployment_id: name = f"resume-{step_offset}-base" if step_offset > 0 else "step-0-base" diff --git a/training/examples/text2sql_sft/train_sft.py b/training/examples/text2sql_sft/train_sft.py index f3693ede..f8326b2d 100644 --- a/training/examples/text2sql_sft/train_sft.py +++ b/training/examples/text2sql_sft/train_sft.py @@ -72,6 +72,7 @@ def main(): ) config = sft_loop.Config( + log_path="./text2sql_logs", base_model=args.base_model, dataset=args.dataset_path, tokenizer_model=args.tokenizer_model, diff --git a/training/recipes/dpo_loop.py b/training/recipes/dpo_loop.py index 15f233ac..967053d7 100644 --- a/training/recipes/dpo_loop.py +++ b/training/recipes/dpo_loop.py @@ -95,6 +95,9 @@ async def gather_with_progress( @dataclass class Config: + log_path: str + """Directory for checkpoints and logs. Required, no default.""" + base_model: str = "accounts/fireworks/models/qwen3-8b" dataset: str = "" tokenizer_model: str = "" # HuggingFace model name for client-side tokenization @@ -119,7 +122,6 @@ class Config: deployment: DeployConfig = field(default_factory=DeployConfig) hotload: HotloadConfig = field(default_factory=lambda: HotloadConfig(hot_load_interval=0)) wandb: WandBConfig = field(default_factory=lambda: WandBConfig(project="dpo-tinker")) - log_path: str = "./dpo_logs" init_from_checkpoint: str | None = None @@ -560,4 +562,4 @@ def _signal_handler(signum, frame): if __name__ == "__main__": logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") - main(Config()) + main(Config(log_path="./dpo_logs")) diff --git a/training/recipes/orpo_loop.py b/training/recipes/orpo_loop.py index 193fe21d..340caa6e 100644 --- a/training/recipes/orpo_loop.py +++ b/training/recipes/orpo_loop.py @@ -72,6 +72,9 @@ @dataclass class Config: + log_path: str + """Directory for checkpoints and logs. Required, no default.""" + base_model: str = "accounts/fireworks/models/qwen3-235b-a22b-instruct-2507" dataset: str = "" tokenizer_model: str = "" @@ -93,7 +96,6 @@ class Config: project="dsv3-training", ) ) - log_path: str = "./orpo_logs" init_from_checkpoint: str | None = None @@ -357,7 +359,8 @@ def _signal_handler(signum, frame): ) cfg = Config( - dataset=os.environ.get("ORPO_DATASET_PATH"), + log_path="./orpo_logs", + dataset=os.environ.get("ORPO_DATASET_PATH"), tokenizer_model=os.environ.get("ORPO_TOKENIZER", "Qwen/Qwen3-235B-A22B-Instruct-2507"), ) main(cfg) diff --git a/training/recipes/rl_loop.py b/training/recipes/rl_loop.py index beef7510..b9e4ac3c 100644 --- a/training/recipes/rl_loop.py +++ b/training/recipes/rl_loop.py @@ -77,6 +77,9 @@ @dataclass class Config: + log_path: str + """Directory for checkpoints and logs. Required, no default.""" + base_model: str = "accounts/fireworks/models/qwen3-8b" dataset: str = "https://raw.githubusercontent.com/eval-protocol/python-sdk/main/development/gsm8k_sample.jsonl" @@ -131,7 +134,6 @@ class Config: reference_base_url: str | None = None """Base URL for the reference trainer (bypass direct route).""" - log_path: str = "./rl_logs" init_from_checkpoint: str | None = None """Load pretrained DCP weights on a fresh dataset. Supports cross-job format ``"job_id:checkpoint_name"``.""" @@ -683,4 +685,4 @@ def _loop_metrics_callback(loop_metrics: dict) -> None: if __name__ == "__main__": logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") - main(Config()) + main(Config(log_path="./rl_logs")) diff --git a/training/recipes/sft_loop.py b/training/recipes/sft_loop.py index 75d32449..6193246f 100644 --- a/training/recipes/sft_loop.py +++ b/training/recipes/sft_loop.py @@ -67,6 +67,9 @@ @dataclass class Config: + log_path: str + """Directory for checkpoints and logs. Required, no default.""" + base_model: str = "accounts/fireworks/models/qwen3-8b" dataset: str = "" tokenizer_model: str = "" # HuggingFace model name for chat template, e.g. "Qwen/Qwen3-1.7B" @@ -83,7 +86,6 @@ class Config: dcp_save_interval: int = 0 # save DCP checkpoint every N steps (0 = off) - log_path: str = "./sft_logs" init_from_checkpoint: str | None = None """Load pretrained DCP weights on a fresh dataset. Supports cross-job format ``"job_id:checkpoint_name"``.""" @@ -336,6 +338,7 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ if __name__ == "__main__": logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") cfg = Config( + log_path="./sft_logs", dataset="kimi2_deid_sample_100_formatted.jsonl", tokenizer_model="Qwen/Qwen3-8B", max_seq_len=4096, diff --git a/training/tests/e2e/test_dpo_resume_e2e.py b/training/tests/e2e/test_dpo_resume_e2e.py index 66309ae8..a12b2f3c 100644 --- a/training/tests/e2e/test_dpo_resume_e2e.py +++ b/training/tests/e2e/test_dpo_resume_e2e.py @@ -61,6 +61,8 @@ def test_dpo_resume_from_checkpoint( with tempfile.NamedTemporaryFile(mode="w", suffix=".jsonl", delete=False) as f: dataset_path = f.name + log_dir = tempfile.mkdtemp(prefix="dpo_resume_") + try: _make_preference_dataset(dataset_path, num_pairs=8) @@ -75,6 +77,7 @@ def test_dpo_resume_from_checkpoint( logger.info("PHASE 1: initial DPO training") phase1_config = Config( + log_path=log_dir, base_model=e2e_model, dataset=dataset_path, beta=0.1, @@ -102,6 +105,7 @@ def test_dpo_resume_from_checkpoint( logger.info("PHASE 2: resume from '%s' (source job: %s)", dcp_name, phase1_job_id) phase2_config = Config( + log_path=log_dir, base_model=e2e_model, dataset=dataset_path, beta=0.1, diff --git a/training/tests/smoke_test/test_dpo_smoke.py b/training/tests/smoke_test/test_dpo_smoke.py index 83278551..b5df00f9 100644 --- a/training/tests/smoke_test/test_dpo_smoke.py +++ b/training/tests/smoke_test/test_dpo_smoke.py @@ -49,6 +49,7 @@ def test_dpo_smoke( try: _make_preference_dataset(dataset_path, num_pairs=4) config = Config( + log_path=tempfile.mkdtemp(prefix="dpo_smoke_"), base_model=smoke_base_model, dataset=dataset_path, tokenizer_model=smoke_tokenizer_model, diff --git a/training/tests/smoke_test/test_grpo_smoke.py b/training/tests/smoke_test/test_grpo_smoke.py index 0bdfba40..d0521754 100644 --- a/training/tests/smoke_test/test_grpo_smoke.py +++ b/training/tests/smoke_test/test_grpo_smoke.py @@ -71,6 +71,7 @@ def test_grpo_smoke( _make_prompt_dataset(dataset_path, n=4) config = Config( + log_path=tempfile.mkdtemp(prefix="grpo_smoke_"), base_model=smoke_base_model, dataset=dataset_path, learning_rate=1e-4, diff --git a/training/tests/smoke_test/test_sft_smoke.py b/training/tests/smoke_test/test_sft_smoke.py index 509f7d70..7db68ac2 100644 --- a/training/tests/smoke_test/test_sft_smoke.py +++ b/training/tests/smoke_test/test_sft_smoke.py @@ -41,6 +41,7 @@ def test_sft_smoke( try: _make_chat_dataset(dataset_path, num_examples=4) config = Config( + log_path=tempfile.mkdtemp(prefix="sft_smoke_"), base_model=smoke_base_model, dataset=dataset_path, tokenizer_model=smoke_tokenizer_model, diff --git a/training/tests/test_defaults.py b/training/tests/test_defaults.py index 80270305..7d2157c2 100644 --- a/training/tests/test_defaults.py +++ b/training/tests/test_defaults.py @@ -6,13 +6,13 @@ def test_grpo_temperature(): from training.recipes.rl_loop import Config - assert Config().temperature == 1.0 + assert Config(log_path="/tmp/test").temperature == 1.0 def test_grpo_max_completion_tokens(): from training.recipes.rl_loop import Config - assert Config().max_completion_tokens == 1024 + assert Config(log_path="/tmp/test").max_completion_tokens == 1024 def test_is_config_defaults(): @@ -26,7 +26,7 @@ def test_is_config_defaults(): def test_cispo_config_defaults(): from training.recipes.rl_loop import Config - cfg = Config() + cfg = Config(log_path="/tmp/test") assert cfg.cispo.eps_low == 0.2 assert cfg.cispo.eps_high == 0.28 @@ -34,11 +34,11 @@ def test_cispo_config_defaults(): def test_cispo_is_valid_policy_loss(): from training.recipes.rl_loop import Config - cfg = Config(policy_loss="cispo") + cfg = Config(log_path="/tmp/test", policy_loss="cispo") assert cfg.policy_loss == "cispo" def test_dpo_has_tokenizer_model(): from training.recipes.dpo_loop import Config - assert hasattr(Config(), "tokenizer_model") + assert hasattr(Config(log_path="/tmp/test"), "tokenizer_model") diff --git a/training/tests/test_smoke_imports.py b/training/tests/test_smoke_imports.py index 679b9e00..43de2f54 100644 --- a/training/tests/test_smoke_imports.py +++ b/training/tests/test_smoke_imports.py @@ -156,28 +156,28 @@ def test_sdk_symbols(module: str, attr: str): def test_sft_config_defaults(): from training.recipes.sft_loop import Config - cfg = Config() + cfg = Config(log_path="/tmp/test") assert cfg.base_model def test_rl_config_defaults(): from training.recipes.rl_loop import Config - cfg = Config() + cfg = Config(log_path="/tmp/test") assert cfg.base_model def test_dpo_config_defaults(): from training.recipes.dpo_loop import Config - cfg = Config() + cfg = Config(log_path="/tmp/test") assert cfg.base_model def test_orpo_config_defaults(): from training.recipes.orpo_loop import Config - cfg = Config() + cfg = Config(log_path="/tmp/test") assert cfg.base_model diff --git a/training/tests/unit/test_dpo_loop.py b/training/tests/unit/test_dpo_loop.py index baa6a252..69440308 100644 --- a/training/tests/unit/test_dpo_loop.py +++ b/training/tests/unit/test_dpo_loop.py @@ -261,7 +261,7 @@ def forward_backward_custom(self, datums, loss_fn): def test_main_requires_tokenizer_model(monkeypatch): monkeypatch.setattr(module, "setup_wandb", lambda *args, **kwargs: None) - cfg = module.Config(dataset="/tmp/pairs.jsonl", tokenizer_model="") + cfg = module.Config(log_path="/tmp/dpo_test_logs", dataset="/tmp/pairs.jsonl", tokenizer_model="") with pytest.raises(ValueError, match="tokenizer_model"): module.main(cfg) @@ -377,6 +377,7 @@ def fake_create_trainer_job(*args, **kwargs): monkeypatch.setattr(transformers.AutoTokenizer, "from_pretrained", lambda *args, **kwargs: object()) cfg = module.Config( + log_path="/tmp/dpo_test_logs", base_model="accounts/test/models/qwen3-4b", dataset="/tmp/pairs.jsonl", tokenizer_model="Qwen/Qwen3-4B", @@ -458,6 +459,7 @@ def save_dcp(self, name): 1: {"chosen_datum": {"id": "c1"}, "rejected_datum": {"id": "r1"}, "ref_chosen": [-0.3], "ref_rejected": [-0.4], "response_start": 4}, } cfg = module.Config( + log_path="/tmp/dpo_test_logs", beta=0.2, epochs=1, batch_size=1, diff --git a/training/tests/unit/test_grpo_streaming_config.py b/training/tests/unit/test_grpo_streaming_config.py index 93afe746..da78f374 100644 --- a/training/tests/unit/test_grpo_streaming_config.py +++ b/training/tests/unit/test_grpo_streaming_config.py @@ -7,12 +7,13 @@ class TestConfigDefaults: def test_defaults(self): - cfg = Config() + cfg = Config(log_path="/tmp/test") assert cfg.completions_per_prompt == 4 assert cfg.prompt_groups_per_step == 1 def test_custom_values(self): cfg = Config( + log_path="/tmp/test", completions_per_prompt=8, prompt_groups_per_step=16, ) diff --git a/training/tests/unit/test_orpo_loop.py b/training/tests/unit/test_orpo_loop.py index fa4c3be4..4f01594c 100644 --- a/training/tests/unit/test_orpo_loop.py +++ b/training/tests/unit/test_orpo_loop.py @@ -10,7 +10,7 @@ def test_main_rejects_invalid_base_model(monkeypatch): monkeypatch.setattr(module, "setup_wandb", lambda *args, **kwargs: None) - cfg = module.Config(base_model="qwen3-4b", dataset="/tmp/pairs.jsonl", tokenizer_model="Qwen/Qwen3-4B") + cfg = module.Config(log_path="/tmp/orpo_test_logs", base_model="qwen3-4b", dataset="/tmp/pairs.jsonl", tokenizer_model="Qwen/Qwen3-4B") with pytest.raises(ValueError, match="Invalid base_model"): module.main(cfg) @@ -122,6 +122,7 @@ def resolve_checkpoint_path(self, name, source_job_id=None): mgr = FakeMgr() cfg = module.Config( + log_path="/tmp/orpo_test_logs", base_model="accounts/test/models/qwen3-4b", dataset="/tmp/pairs.jsonl", tokenizer_model="Qwen/Qwen3-4B", diff --git a/training/tests/unit/test_rl_loop.py b/training/tests/unit/test_rl_loop.py index 9aa5ce97..5fcce5ee 100644 --- a/training/tests/unit/test_rl_loop.py +++ b/training/tests/unit/test_rl_loop.py @@ -81,6 +81,7 @@ def test_dump_trajectory_writes_one_record_per_completion(tmp_path): def test_main_requires_deployment_tokenizer_model(monkeypatch): monkeypatch.setattr(module, "setup_wandb", lambda *args, **kwargs: None) cfg = module.Config( + log_path="/tmp/rl_test_logs", dataset="/tmp/prompts.jsonl", deployment=module.DeployConfig(tokenizer_model=""), ) @@ -167,6 +168,7 @@ async def fake_run_rl_loop(**kwargs): monkeypatch.setattr(module, "run_rl_loop", fake_run_rl_loop) cfg = module.Config( + log_path="/tmp/rl_test_logs", base_model="accounts/test/models/qwen3-4b", dataset="/tmp/prompts.jsonl", kl_beta=0.0, @@ -439,6 +441,7 @@ def _builder(adv, ref_lp, prompt_lens, inf_lp, prox_lp): ) cfg = module.Config( + log_path=str(tmp_path / "rl_logs"), base_model="accounts/test/models/qwen3-4b", dataset="/tmp/prompts.jsonl", kl_beta=0.1, diff --git a/training/tests/unit/test_sft_loop.py b/training/tests/unit/test_sft_loop.py index ad7eed3b..fe561538 100644 --- a/training/tests/unit/test_sft_loop.py +++ b/training/tests/unit/test_sft_loop.py @@ -33,7 +33,7 @@ def test_main_requires_tokenizer_model(tmp_path, monkeypatch): ) monkeypatch.setattr(module, "setup_wandb", lambda *args, **kwargs: None) - cfg = module.Config(dataset=str(dataset_path), tokenizer_model="", max_seq_len=32) + cfg = module.Config(log_path=str(tmp_path / "logs"), dataset=str(dataset_path), tokenizer_model="", max_seq_len=32) with pytest.raises(ValueError, match="tokenizer_model"): module.main(cfg) @@ -72,6 +72,7 @@ def __init__(self, *args, **kwargs): monkeypatch.setattr(module, "ReconnectableClient", FakeClient) cfg = module.Config( + log_path=str(tmp_path / "logs"), dataset=str(dataset_path), tokenizer_model="Qwen/Qwen3-4B", max_seq_len=32, From d11964de5c9b4247054f8f8fb3fc9dab7e4804b1 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Tue, 10 Mar 2026 15:00:17 -0700 Subject: [PATCH 26/28] test: add log_path required tests + save/resume roundtrip Verify all 4 recipe Configs reject construction without log_path, accept it when provided, and test the full save -> resume roundtrip with directory creation. Made-with: Cursor --- training/tests/unit/test_checkpoint_utils.py | 73 ++++++++++++++++++++ 1 file changed, 73 insertions(+) diff --git a/training/tests/unit/test_checkpoint_utils.py b/training/tests/unit/test_checkpoint_utils.py index ea40f6b0..0c7eca55 100644 --- a/training/tests/unit/test_checkpoint_utils.py +++ b/training/tests/unit/test_checkpoint_utils.py @@ -237,3 +237,76 @@ def test_exact_division(self): ds = RLPromptDataset(rows, prompts_per_step=4) assert len(ds) == 3 assert ds.get_batch(2) == rows[8:12] + + +class TestLogPathRequired: + """Verify log_path is a required field on all recipe Configs (no default).""" + + def test_sft_config_requires_log_path(self): + from training.recipes.sft_loop import Config + with pytest.raises(TypeError, match="log_path"): + Config() + + def test_rl_config_requires_log_path(self): + from training.recipes.rl_loop import Config + with pytest.raises(TypeError, match="log_path"): + Config() + + def test_dpo_config_requires_log_path(self): + from training.recipes.dpo_loop import Config + with pytest.raises(TypeError, match="log_path"): + Config() + + def test_orpo_config_requires_log_path(self): + from training.recipes.orpo_loop import Config + with pytest.raises(TypeError, match="log_path"): + Config() + + def test_sft_config_accepts_log_path(self): + from training.recipes.sft_loop import Config + cfg = Config(log_path="/tmp/test_sft") + assert cfg.log_path == "/tmp/test_sft" + + def test_rl_config_accepts_log_path(self): + from training.recipes.rl_loop import Config + cfg = Config(log_path="/tmp/test_rl") + assert cfg.log_path == "/tmp/test_rl" + + def test_save_checkpoint_creates_log_dir(self, log_dir): + """save_checkpoint creates log_path if it doesn't exist.""" + nested = os.path.join(log_dir, "deep", "nested") + assert not os.path.exists(nested) + client = _make_mock_client() + save_checkpoint(client, "step-1", nested, {"step": 1}) + assert os.path.exists(os.path.join(nested, CHECKPOINTS_BASE_NAME)) + + def test_resolve_resume_with_empty_log_dir(self, log_dir): + """resolve_resume returns None for a fresh log_path with no checkpoints.""" + client = _make_mock_client() + result = resolve_resume(client, log_dir) + assert result is None + client.load_state_with_optimizer.assert_not_called() + + def test_save_then_resume_roundtrip(self, log_dir): + """Full roundtrip: save a checkpoint, then resume from it.""" + client = _make_mock_client(job_id="job-roundtrip") + save_checkpoint(client, "step-3", log_dir, { + "step": 3, + "data_consumed": 24, + "source_job_id": "job-roundtrip", + }) + + ckpt_path = os.path.join(log_dir, CHECKPOINTS_BASE_NAME) + assert os.path.exists(ckpt_path) + with open(ckpt_path) as f: + entry = json.loads(f.readline()) + assert entry["name"] == "step-3" + assert entry["step"] == 3 + assert "state_path" in entry + + client2 = _make_mock_client(job_id="job-new") + result = resolve_resume(client2, log_dir) + assert result is not None + assert result.step == 3 + assert result.data_consumed == 24 + client2.load_state_with_optimizer.assert_called_once() From 134db619ca3454bb95e1e2681adb2941f3dd973d Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Tue, 10 Mar 2026 15:10:47 -0700 Subject: [PATCH 27/28] docs: explain why save/resume are local, not from tinker_cookbook Made-with: Cursor --- training/utils/checkpoint_utils.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/training/utils/checkpoint_utils.py b/training/utils/checkpoint_utils.py index c47eec66..d0f8eba5 100644 --- a/training/utils/checkpoint_utils.py +++ b/training/utils/checkpoint_utils.py @@ -1,9 +1,8 @@ """Checkpoint utilities using tinker_cookbook's checkpoints.jsonl format. -Checkpoint state is persisted in ``checkpoints.jsonl`` (same format as -``tinker_cookbook.checkpoint_utils``). Reading uses ``get_last_checkpoint`` -imported directly from tinker_cookbook; writing uses a sync -``save_checkpoint`` that follows the same schema. +Reading uses ``get_last_checkpoint`` from tinker_cookbook directly. +Writing (``save_checkpoint``) and resume (``resolve_resume``) are +implemented locally for Fireworks RLOR compatibility. """ from __future__ import annotations From 36e72d1b791ef09244f68bbb0612ca5e8c6cdcc7 Mon Sep 17 00:00:00 2001 From: Chengxi Li Date: Tue, 10 Mar 2026 15:27:42 -0700 Subject: [PATCH 28/28] refactor: remove redundant validation helpers Delete dataset_fingerprint, validate_dataset, validate_training_shape from checkpoint_utils -- these were Fireworks-specific validation checks that added complexity without real value. Remove dataset_fingerprint and training_shape_id from ResumeInfo and all loop_state dicts in save_checkpoint calls. Made-with: Cursor --- training/recipes/rl_loop.py | 9 --- training/recipes/sft_loop.py | 10 ---- training/tests/unit/test_checkpoint_utils.py | 59 +------------------- training/utils/checkpoint_utils.py | 54 ------------------ 4 files changed, 1 insertion(+), 131 deletions(-) diff --git a/training/recipes/rl_loop.py b/training/recipes/rl_loop.py index b9e4ac3c..36767ce9 100644 --- a/training/recipes/rl_loop.py +++ b/training/recipes/rl_loop.py @@ -51,8 +51,6 @@ from training.utils.checkpoint_utils import ( resolve_resume, save_checkpoint, - dataset_fingerprint, - validate_dataset, ) from fireworks.training.sdk.deployment import DeploymentSampler from training.utils.rl import PromptGroup @@ -403,9 +401,6 @@ def _signal_handler(signum, frame): # -- Prepare sampling and training -------------------------------------- raw_dataset = load_jsonl_dataset(cfg.dataset, cfg.max_rows) - fp = dataset_fingerprint(raw_dataset) - if resume_info: - validate_dataset(resume_info.dataset_fingerprint, fp, resume_info.data_consumed) all_rows = raw_dataset * cfg.epochs rl_dataset = RLPromptDataset(all_rows, prompts_per_step=prompt_groups_per_step) adam_params = tinker.AdamParams(learning_rate=cfg.learning_rate, **DEFAULT_ADAM) @@ -581,8 +576,6 @@ def finish_step( save_checkpoint(policy, f"step-{step}", cfg.log_path, { "step": step, "data_consumed": _data_consumed, - "dataset_fingerprint": fp, - "training_shape_id": getattr(cfg.infra, "training_shape_id", None), "source_job_id": policy_job_id, }) @@ -647,8 +640,6 @@ def _loop_metrics_callback(loop_metrics: dict) -> None: save_checkpoint(policy, f"step-{global_step}", cfg.log_path, { "step": global_step, "data_consumed": _data_consumed, - "dataset_fingerprint": fp, - "training_shape_id": getattr(cfg.infra, "training_shape_id", None), "source_job_id": policy_job_id, }) except Exception as e: diff --git a/training/recipes/sft_loop.py b/training/recipes/sft_loop.py index 6193246f..37653d24 100644 --- a/training/recipes/sft_loop.py +++ b/training/recipes/sft_loop.py @@ -51,8 +51,6 @@ from training.utils.checkpoint_utils import ( resolve_resume, save_checkpoint, - dataset_fingerprint, - validate_dataset, ) from training.utils.timer import timer, flush_timing @@ -230,10 +228,6 @@ def _map_fn(row: dict) -> tinker.Datum | None: data_consumed = resume_info.data_consumed if resume_info else 0 wandb_log({"train/step": step}, step) - fp = dataset_fingerprint(raw_data) - if resume_info: - validate_dataset(resume_info.dataset_fingerprint, fp, data_consumed) - adam_params = tinker.AdamParams(learning_rate=cfg.learning_rate, **DEFAULT_ADAM) # -- Training loop (batch-indexed) ------------------------------------- @@ -268,8 +262,6 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ save_checkpoint(client, f"step-{step}", cfg.log_path, { "step": step, "data_consumed": data_consumed, - "dataset_fingerprint": fp, - "training_shape_id": getattr(cfg.infra, "training_shape_id", None), "source_job_id": job_id, }) @@ -319,8 +311,6 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ save_checkpoint(client, f"step-{step}", cfg.log_path, { "step": step, "data_consumed": data_consumed, - "dataset_fingerprint": fp, - "training_shape_id": getattr(cfg.infra, "training_shape_id", None), "source_job_id": job_id, }, kind="both") diff --git a/training/tests/unit/test_checkpoint_utils.py b/training/tests/unit/test_checkpoint_utils.py index 0c7eca55..21b063c3 100644 --- a/training/tests/unit/test_checkpoint_utils.py +++ b/training/tests/unit/test_checkpoint_utils.py @@ -1,4 +1,4 @@ -"""Unit tests for checkpoint_utils -- resume, save, fingerprints.""" +"""Unit tests for checkpoint_utils -- resume, save.""" import json import os @@ -11,9 +11,6 @@ ResumeInfo, resolve_resume, save_checkpoint, - dataset_fingerprint, - validate_dataset, - validate_training_shape, CHECKPOINTS_BASE_NAME, ) @@ -67,8 +64,6 @@ def test_resume_from_checkpoints_file(self, log_dir): "step": 5, "data_consumed": 40, "state_path": "cross_job://job-abc/step-5", - "dataset_fingerprint": "abc123", - "training_shape_id": "ts-qwen3-8b-policy", "source_job_id": "job-abc", }) client = _make_mock_client() @@ -76,8 +71,6 @@ def test_resume_from_checkpoints_file(self, log_dir): assert result is not None assert result.step == 5 assert result.data_consumed == 40 - assert result.dataset_fingerprint == "abc123" - assert result.training_shape_id == "ts-qwen3-8b-policy" assert result.source_job_id == "job-abc" client.load_state_with_optimizer.assert_called_once_with("cross_job://job-abc/step-5") @@ -166,56 +159,6 @@ def test_creates_directory(self): assert os.path.exists(os.path.join(nested, CHECKPOINTS_BASE_NAME)) -class TestDatasetFingerprint: - def test_consistent(self): - rows = [{"a": 1}, {"b": 2}, {"c": 3}] - fp1 = dataset_fingerprint(rows) - fp2 = dataset_fingerprint(rows) - assert fp1 == fp2 - assert len(fp1) == 12 - - def test_different_data(self): - fp1 = dataset_fingerprint([{"a": 1}]) - fp2 = dataset_fingerprint([{"a": 2}]) - assert fp1 != fp2 - - def test_different_length(self): - fp1 = dataset_fingerprint([{"a": 1}]) - fp2 = dataset_fingerprint([{"a": 1}, {"b": 2}]) - assert fp1 != fp2 - - def test_empty(self): - assert dataset_fingerprint([]) == "empty" - - -class TestValidateDataset: - def test_no_warning_on_match(self, log_dir, caplog): - validate_dataset("abc123", "abc123", 40) - assert "Dataset changed" not in caplog.text - - def test_warning_on_mismatch(self, log_dir, caplog): - validate_dataset("abc123", "xyz789", 40) - assert "Dataset changed" in caplog.text - - def test_no_warning_when_saved_is_none(self, log_dir, caplog): - validate_dataset(None, "abc123", 0) - assert "Dataset changed" not in caplog.text - - -class TestValidateTrainingShape: - def test_no_warning_on_match(self, caplog): - validate_training_shape("ts-qwen3-8b", "ts-qwen3-8b") - assert "Training shape changed" not in caplog.text - - def test_warning_on_mismatch(self, caplog): - validate_training_shape("ts-qwen3-8b", "ts-qwen3-30b") - assert "Training shape changed" in caplog.text - - def test_no_warning_when_none(self, caplog): - validate_training_shape(None, "ts-qwen3-8b") - assert "Training shape changed" not in caplog.text - - class TestRLPromptDataset: def test_get_batch_basic(self): from training.utils.data import RLPromptDataset diff --git a/training/utils/checkpoint_utils.py b/training/utils/checkpoint_utils.py index d0f8eba5..54736508 100644 --- a/training/utils/checkpoint_utils.py +++ b/training/utils/checkpoint_utils.py @@ -7,7 +7,6 @@ from __future__ import annotations -import hashlib import json import logging import os @@ -32,8 +31,6 @@ class ResumeInfo: step: int = 0 data_consumed: int = 0 - dataset_fingerprint: str | None = None - training_shape_id: str | None = None source_job_id: str | None = None @@ -74,8 +71,6 @@ def resolve_resume( return ResumeInfo( step=last.get("step", 0), data_consumed=last.get("data_consumed", 0), - dataset_fingerprint=last.get("dataset_fingerprint"), - training_shape_id=last.get("training_shape_id"), source_job_id=last.get("source_job_id"), ) @@ -117,52 +112,3 @@ def save_checkpoint( f.write(json.dumps(full_dict) + "\n") logger.info("Saved checkpoint: %s", full_dict) return paths - - -# -- Dataset fingerprint ------------------------------------------------------- - - -def dataset_fingerprint(rows: list[dict]) -> str: - """Short hash of row count + first/last row content.""" - if not rows: - return "empty" - content = ( - f"{len(rows)}:" - f"{json.dumps(rows[0], sort_keys=True)}:" - f"{json.dumps(rows[-1], sort_keys=True)}" - ) - return hashlib.sha256(content.encode()).hexdigest()[:12] - - -def validate_dataset( - saved_fingerprint: str | None, - current_fingerprint: str, - data_consumed: int, -) -> None: - """Warn if the dataset changed between checkpoint save and resume.""" - if saved_fingerprint and saved_fingerprint != current_fingerprint: - logger.warning( - "Dataset changed since checkpoint! " - "fingerprint: saved=%s current=%s. " - "data_consumed=%d may point to different data.", - saved_fingerprint, - current_fingerprint, - data_consumed, - ) - - -# -- Training shape validation ------------------------------------------------- - - -def validate_training_shape( - saved_shape_id: str | None, - current_shape_id: str | None, -) -> None: - """Warn if the training shape changed between checkpoint save and resume.""" - if saved_shape_id and current_shape_id and saved_shape_id != current_shape_id: - logger.warning( - "Training shape changed! saved=%s current=%s. " - "Model parallelism may be incompatible.", - saved_shape_id, - current_shape_id, - )