"""Launches the training stages on Hugging Face Jobs. python3 main.py baseline render vanilla SD 1.5 on the eval prompts python3 main.py latents VAE-encode the dataset once python3 main.py train LoRA fine-tune on the cached latents python3 main.py sample checkpoint-2000 Everything runs on HF infrastructure; the M2 has no usable GPU. No -v volume is passed - each job pulls what it needs from the Hub itself. Order matters: `baseline` first (so there is something to compare against), then `latents`, then `train`, then `sample` per checkpoint. Per-run overrides go through the environment rather than flags here, since each job script already reads its own constants that way: EPOCHS=20 python3 main.py train """ import os import subprocess import sys from pathlib import Path HERE = Path(__file__).resolve().parent # Job scripts are self-contained PEP-723 files. They ship to HF standalone # and cannot import from this repo, so each carries its own constants and # writes them into a manifest - no shared config module to drift out of sync. STAGES = { "latents": (HERE / "cache_latents_job.py", "a10g-small", "2h"), "train": (HERE / "train_lora_job.py", "a10g-small", "6h"), "sample": (HERE / "sample_job.py", "a10g-small", "1h"), } # Variables worth forwarding to the job if they are set locally. FORWARDED = [ "RESOLUTION", "EPOCHS", "BATCH_SIZE", "GRAD_ACCUM", "LEARNING_RATE", "LORA_RANK", "LORA_ALPHA", "CAPTION_DROPOUT", "CHECKPOINT_EVERY", "RESUME_FROM", "SEED", "CHECKPOINT", "STEPS", "GUIDANCE", "LATENTS_REVISION", "SOURCE_REVISION", "TARGET_REVISION", "MODEL_REPO", ] DATASET_REPO = "whosouravsharma/text-to-image-diffusiondb-2M" def require_latents() -> None: """Refuse to launch training before stage 1 has produced latents. Checked here rather than only inside the job: a GPU job that dies on a missing branch still costs a scheduling round-trip, and the failure shows up as a 404 traceback rather than as the one line that explains it. """ from huggingface_hub import HfApi revision = os.environ.get("LATENTS_REVISION", "latents-512") branches = [ b.name for b in HfApi().list_repo_refs(DATASET_REPO, repo_type="dataset").branches ] if revision not in branches: raise SystemExit( f"No '{revision}' branch on {DATASET_REPO}.\n" f"Existing branches: {', '.join(branches)}\n\n" f"Cache the latents first:\n" f" python3 main.py latents" ) def run_stage(stage: str, checkpoint: str | None = None) -> int: if stage == "baseline": stage, checkpoint = "sample", "base" if stage not in STAGES: raise SystemExit( f"Unknown stage {stage!r}. " f"Choose from: baseline, {', '.join(STAGES)}" ) if stage == "train": require_latents() script, flavor, timeout = STAGES[stage] command = [ "hf", "jobs", "uv", "run", "--flavor", flavor, "--timeout", timeout, "--secrets", "HF_TOKEN", ] if checkpoint: command += ["--env", f"CHECKPOINT={checkpoint}"] for key in FORWARDED: if key in os.environ and not (checkpoint and key == "CHECKPOINT"): command += ["--env", f"{key}={os.environ[key]}"] command.append(str(script)) print("=" * 60) print(f"TRAINING STAGE: {stage}" + (f" ({checkpoint})" if checkpoint else "")) print("=" * 60) print("\nRunning:", " ".join(command), "\n") return subprocess.call(command) if __name__ == "__main__": if len(sys.argv) < 2: raise SystemExit(__doc__) sys.exit(run_stage(sys.argv[1], sys.argv[2] if len(sys.argv) > 2 else None))