Text-to-Image
Diffusers
stable-diffusion
stable-diffusion-diffusers
lora
Eval Results (legacy)
whosouravsharma's picture
Add the training scripts that produced these checkpoints
1ab3690 verified
Raw History Blame Contribute Delete
3.76 kB
"""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))