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/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 51cacdc4..ae1b53a5 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, @@ -99,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" @@ -545,7 +545,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) - step_offset, _ = setup_resume(policy, ResumeConfig()) + from training.utils.checkpoint_utils import resolve_resume + 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" weight_syncer.save_and_hotload(name, checkpoint_type="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/pyproject.toml b/training/pyproject.toml index 7adc0a3e..383d02da 100644 --- a/training/pyproject.toml +++ b/training/pyproject.toml @@ -7,8 +7,8 @@ name = "fireworks-training-cookbook" version = "0.1.0" requires-python = ">=3.11" dependencies = [ - "fireworks-ai>=1.0.0a39,<2", - "tinker-cookbook @ git+https://github.com/thinking-machines-lab/tinker-cookbook.git@934f0d9b2f53c3edff02cbf23ec6da8682047fa5", + "fireworks-ai>=1.0.0a40,<2", + "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 5d92ee84..967053d7 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 from training.utils.timer import timer, flush_timing logger = logging.getLogger(__name__) @@ -96,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 @@ -120,7 +122,7 @@ 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) + init_from_checkpoint: str | None = None # --------------------------------------------------------------------------- @@ -389,7 +391,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 +487,9 @@ def _signal_handler(signum, frame): compression_format=DEFAULT_DELTA_COMPRESSION, ) - step_offset, _ = setup_resume(policy, cfg.resume) + 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 ------------------------------ @@ -533,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} @@ -556,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 780f6555..340caa6e 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 logger = logging.getLogger(__name__) @@ -73,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 = "" @@ -94,7 +96,7 @@ class Config: project="dsv3-training", ) ) - resume: ResumeConfig = field(default_factory=ResumeConfig) + init_from_checkpoint: str | None = None # --------------------------------------------------------------------------- @@ -175,7 +177,8 @@ def _signal_handler(signum, frame): ) job_id = endpoint.job_id - step_offset, _ = setup_resume(client, cfg.resume) + 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 ---------------------------------------------------------------- @@ -325,8 +328,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) @@ -353,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 59a43f07..36767ce9 100644 --- a/training/recipes/rl_loop.py +++ b/training/recipes/rl_loop.py @@ -34,12 +34,11 @@ InfraConfig, WandBConfig, DeployConfig, - ResumeConfig, HotloadConfig, ReconnectableClient, + RLPromptDataset, wandb_log, setup_wandb, - setup_resume, wandb_finish, validate_config, log_metrics_json, @@ -49,6 +48,10 @@ load_jsonl_dataset, prepare_sampling_messages, ) +from training.utils.checkpoint_utils import ( + resolve_resume, + save_checkpoint, +) from fireworks.training.sdk.deployment import DeploymentSampler from training.utils.rl import PromptGroup from training.utils.rl.importance_sampling import ISConfig @@ -72,6 +75,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" @@ -126,11 +132,14 @@ class Config: reference_base_url: str | None = None """Base URL for the reference trainer (bypass direct route).""" + 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: @@ -342,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 @@ -377,14 +388,21 @@ 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 --------------------------------------------------------------- + + 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" weight_syncer.save_and_hotload(name, checkpoint_type="base") # -- Prepare sampling and training -------------------------------------- - dataset = load_jsonl_dataset(cfg.dataset, cfg.max_rows) + raw_dataset = load_jsonl_dataset(cfg.dataset, cfg.max_rows) + 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, @@ -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, @@ -514,20 +532,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( @@ -538,25 +555,29 @@ 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) + 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, + "source_job_id": policy_job_id, + }) metrics = compute_step_metrics( prompt_groups=prompt_groups, @@ -596,10 +617,12 @@ def _loop_metrics_callback(loop_metrics: dict) -> None: finish_step=finish_step, ) - all_rows = dataset * cfg.epochs + 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, @@ -613,7 +636,12 @@ 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 = (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, + "source_job_id": policy_job_id, + }) except Exception as e: logger.warning("Failed to save final checkpoint: %s", e) @@ -648,4 +676,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 b04e7bc8..37653d24 100644 --- a/training/recipes/sft_loop.py +++ b/training/recipes/sft_loop.py @@ -25,19 +25,19 @@ 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, WandBConfig, - ResumeConfig, ReconnectableClient, wandb_log, setup_wandb, - setup_resume, wandb_finish, validate_config, log_metrics_json, @@ -48,6 +48,10 @@ render_messages_to_datum, resolve_renderer_name, ) +from training.utils.checkpoint_utils import ( + resolve_resume, + save_checkpoint, +) from training.utils.timer import timer, flush_timing @@ -61,6 +65,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" @@ -77,9 +84,12 @@ class Config: dcp_save_interval: int = 0 # save DCP checkpoint every N steps (0 = off) + 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 +111,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, { @@ -151,7 +161,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: @@ -175,49 +185,60 @@ 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") - step_offset, _ = setup_resume(client, cfg.resume) + 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 --------------------------------------------------------------- + + 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) + adam_params = tinker.AdamParams(learning_rate=cfg.learning_rate, **DEFAULT_ADAM) - # -- Training loop (batched) ------------------------------------------- + # -- Training loop (batch-indexed) ------------------------------------- - batch_size = cfg.batch_size - step = step_offset - total_steps = len(training_data) * cfg.epochs // (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() @@ -238,7 +259,11 @@ 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}") + save_checkpoint(client, f"step-{step}", cfg.log_path, { + "step": step, + "data_consumed": data_consumed, + "source_job_id": job_id, + }) step_metrics: Dict[str, Any] = flush_timing() @@ -251,10 +276,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({ @@ -269,16 +291,13 @@ 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: - 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) @@ -286,15 +305,14 @@ def _flush_batch(batch_buf: list[tinker.Datum], step: int, accum: int) -> tuple[ # -- Final checkpoint -------------------------------------------------- - if step > step_offset: - logger.info("Saving final DCP checkpoint (step %d)...", step) - client.inner.save_state(f"final-step-{step}") - - logger.info("Saving final base checkpoint (step %d)...", step) - result = client.inner.save_weights_for_sampler_ext( - f"final-step-{step}", checkpoint_type="base" - ) - logger.info("Final base checkpoint saved: %s", result.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, + "source_job_id": job_id, + }, kind="both") logger.info("Training complete: %d optimizer steps", step) return {"steps": step, "job_id": job_id} @@ -310,6 +328,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 948fd211..a12b2f3c 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__) @@ -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, @@ -112,7 +116,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..eb45e7d9 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 training.utils import InfraConfig, DeployConfig, ResumeConfig, HotloadConfig -from training.utils.rl import ISConfig +from tinker_cookbook.checkpoint_utils import get_last_checkpoint +from training.utils import InfraConfig, DeployConfig, HotloadConfig 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,24 +110,32 @@ 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) + 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) - # 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, + deployment_id=phase1_deployment_id, tokenizer_model=e2e_tokenizer_model, ), hotload=HotloadConfig( @@ -132,10 +144,6 @@ 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, - ), ) phase2_metrics = main(phase2_config, rlor_mgr=rlor_mgr, deploy_mgr=deploy_mgr) @@ -143,5 +151,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 080080ad..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,7 +21,8 @@ import pytest -from training.utils import InfraConfig, DeployConfig, ResumeConfig, HotloadConfig +from tinker_cookbook.checkpoint_utils import get_last_checkpoint +from training.utils import InfraConfig from training.recipes.sft_loop import Config, main logger = logging.getLogger(__name__) @@ -65,16 +70,16 @@ 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=20) 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.33.0", ) - # Phase 1: train, save DCP + log_dir = tempfile.mkdtemp(prefix="sft_resume_") + + # Phase 1: train, save DCP checkpoints to checkpoints.jsonl logger.info("PHASE 1: initial SFT training") phase1_config = Config( @@ -83,26 +88,39 @@ def test_sft_resume_from_checkpoint( tokenizer_model=tokenizer_model, learning_rate=1e-4, epochs=2, + batch_size=4, grad_accum=2, - max_examples=10, + max_seq_len=4096, + max_examples=20, + dcp_save_interval=2, + log_path=log_dir, 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 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, @@ -110,23 +128,38 @@ def test_sft_resume_from_checkpoint( tokenizer_model=tokenizer_model, learning_rate=1e-4, epochs=2, + batch_size=4, grad_accum=2, - max_examples=10, + max_seq_len=4096, + max_examples=20, + log_path=log_dir, infra=shared_infra, - deployment=DeployConfig(), - hotload=HotloadConfig(hot_load_interval=0), - resume=ResumeConfig(resume_from=dcp_name, resume_job_id=phase1_job_id), ) - 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 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["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 (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/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 aab1d064..43de2f54 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", @@ -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_checkpoint_utils.py b/training/tests/unit/test_checkpoint_utils.py new file mode 100644 index 00000000..21b063c3 --- /dev/null +++ b/training/tests/unit/test_checkpoint_utils.py @@ -0,0 +1,255 @@ +"""Unit tests for checkpoint_utils -- resume, save.""" + +import json +import os +import tempfile +from unittest.mock import MagicMock + +import pytest + +from training.utils.checkpoint_utils import ( + ResumeInfo, + resolve_resume, + save_checkpoint, + CHECKPOINTS_BASE_NAME, +) + + +@pytest.fixture +def log_dir(): + with tempfile.TemporaryDirectory() as d: + 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): + 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, + "state_path": "cross_job://job-abc/step-5", + "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.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): + _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): + 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): + 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): + 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}) + + 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_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 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] + + +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() diff --git a/training/tests/unit/test_dpo_loop.py b/training/tests/unit/test_dpo_loop.py index 2a27bd33..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) @@ -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 ( @@ -368,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", @@ -449,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 65b06b76..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) @@ -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( @@ -109,6 +122,7 @@ def optim_step(self, _params): 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 408ae1a8..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=""), ) @@ -124,8 +125,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): @@ -157,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, @@ -280,8 +292,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,7 +402,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) - monkeypatch.setattr(module, "setup_resume", lambda *args, **kwargs: (1, None)) + 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": [ @@ -418,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, @@ -457,10 +481,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..fe561538 100644 --- a/training/tests/unit/test_sft_loop.py +++ b/training/tests/unit/test_sft_loop.py @@ -33,114 +33,12 @@ 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) -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) @@ -181,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, @@ -192,7 +84,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 +103,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 +113,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 +126,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 +155,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/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): diff --git a/training/utils/__init__.py b/training/utils/__init__.py index a59595d4..63cba66d 100644 --- a/training/utils/__init__.py +++ b/training/utils/__init__.py @@ -12,8 +12,8 @@ "HotloadConfig", "InfraConfig", "ReconnectableClient", - "ResumeConfig", "RewardFn", + "RLPromptDataset", "StepCallback", "WandBConfig", "compute_advantages", @@ -43,7 +43,6 @@ "resolve_renderer_name", "prepare_sampling_messages", "setup_deployment", - "setup_resume", "setup_training_client", "setup_wandb", "flush_timing", @@ -56,6 +55,7 @@ ] from training.utils.data import ( + RLPromptDataset, encode_text, extract_text, compute_advantages, @@ -78,7 +78,6 @@ InfraConfig, WandBConfig, DeployConfig, - ResumeConfig, StepCallback, HotloadConfig, ) @@ -102,7 +101,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..54736508 --- /dev/null +++ b/training/utils/checkpoint_utils.py @@ -0,0 +1,114 @@ +"""Checkpoint utilities using tinker_cookbook's checkpoints.jsonl format. + +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 + +import json +import logging +import os +import time +from dataclasses import dataclass +from typing import Any + +from tinker_cookbook.checkpoint_utils import ( + get_last_checkpoint, + CHECKPOINTS_BASE_NAME, +) + +logger = logging.getLogger(__name__) + + +# -- Resume info --------------------------------------------------------------- + + +@dataclass +class ResumeInfo: + """Resolved resume state returned by ``resolve_resume``.""" + + step: int = 0 + data_consumed: int = 0 + 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, +) -> 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, 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 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), + source_job_id=last.get("source_job_id"), + ) + + logger.info("Fresh start (no checkpoint)") + return None + + +# -- Checkpoint save ----------------------------------------------------------- + + +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. + + *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 + + full_dict = {"name": name, **loop_state, **paths} + os.makedirs(log_path, exist_ok=True) + 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 diff --git a/training/utils/client.py b/training/utils/client.py index 8109e057..c51fa486 100644 --- a/training/utils/client.py +++ b/training/utils/client.py @@ -112,9 +112,15 @@ 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) + 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/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/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, 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/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( 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(