""" Wrapper around sft_12hz.py that pushes each checkpoint to HF Hub after every epoch. Reads sft_12hz.py's checkpoint output dirs and uploads them. """ import argparse import subprocess import sys import os import shutil from pathlib import Path from huggingface_hub import HfApi, create_repo def parse_args(): p = argparse.ArgumentParser() p.add_argument("--init_model_path", required=True) p.add_argument("--output_model_path", required=True) p.add_argument("--train_jsonl", required=True) p.add_argument("--batch_size", type=int, default=4) p.add_argument("--lr", type=float, default=2e-6) p.add_argument("--num_epochs", type=int, default=5) p.add_argument("--speaker_name", default="mostafa_speaker") p.add_argument("--hf_repo", required=True) return p.parse_args() def push_checkpoint(local_dir: Path, hf_repo: str, epoch: int): api = HfApi() try: create_repo(hf_repo, exist_ok=True, repo_type="model") except Exception: pass print(f"\n→ Pushing checkpoint-epoch-{epoch} to {hf_repo} ...") api.upload_folder( folder_path=str(local_dir), repo_id=hf_repo, repo_type="model", commit_message=f"epoch {epoch} checkpoint", ) print(f"āœ“ Pushed checkpoint-epoch-{epoch}") def main(): args = parse_args() output_dir = Path(args.output_model_path) # Run one epoch at a time so we can push after each for epoch in range(args.num_epochs): print(f"\n{'='*50}") print(f"Training epoch {epoch+1}/{args.num_epochs}") print(f"{'='*50}") cmd = [ sys.executable, "sft_12hz.py", "--init_model_path", args.init_model_path, "--output_model_path", args.output_model_path, "--train_jsonl", args.train_jsonl, "--batch_size", str(args.batch_size), "--lr", str(args.lr), "--num_epochs", str(epoch + 1), # train up to this epoch "--speaker_name", args.speaker_name, ] result = subprocess.run(cmd, check=True) # Push the latest checkpoint ckpt_dir = output_dir / f"checkpoint-epoch-{epoch}" if ckpt_dir.exists(): push_checkpoint(ckpt_dir, args.hf_repo, epoch) # --- NEW CLEANUP BLOCK --- try: shutil.rmtree(ckpt_dir) print(f"āœ“ Deleted local folder {ckpt_dir} to free up SSD space.") except Exception as e: print(f"Warning: Could not delete {ckpt_dir}: {e}") # ------------------------- else: print(f"Warning: {ckpt_dir} not found, skipping push") print("\nAll epochs done and pushed.") if __name__ == "__main__": main()