| """ |
| 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) |
|
|
| |
| 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), |
| "--speaker_name", args.speaker_name, |
| ] |
|
|
| result = subprocess.run(cmd, check=True) |
|
|
| |
| ckpt_dir = output_dir / f"checkpoint-epoch-{epoch}" |
| if ckpt_dir.exists(): |
| push_checkpoint(ckpt_dir, args.hf_repo, epoch) |
| |
| |
| 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() |