Perf core for the diffusion trainers, defaults preserving the training math:
- Phased model loading: the pipeline now loads without its transformer
(conditioning only), captions are encoded and the text encoders freed,
the VAE latent cache is built and the VAE freed, and only then does the
transformer load. The multi-GB denoiser never shares VRAM with the
encoders, cutting measured peak VRAM on B200: FLUX 17.1 -> 10.4 GB,
Qwen-Image 19.1 -> 12.8 GB, Z-Image 7.3 -> 4.7 GB.
- Latent cache (cache_latents, default on): per-image crop/flip variants
(cache_variants, default 4 vs the single frozen variant of the diffusers
--cache_latents) store the VAE posterior's affine parameters, so every
step still draws a fresh VAE sample; a cached center-crop Z-Image run
matches the uncached one at the bf16 nondeterminism floor.
- True batching: train_batch_size now actually batches the transformer
forward (it was silently 1). nf4 dequant dominates the step cost, so
batch 4 lands near batch-1 step time: 4.0x samples/s on Qwen-Image,
3.1x on FLUX, 2.1x on Z-Image, with multi-seed loss envelopes
overlapping batch-1.
- LR scheduler support in the DiT loop (lr_scheduler / lr_warmup_steps
were accepted but ignored); progress events now report the real
per-step LR.
- TF32 + high fp32 matmul precision under enable_tf32 (default on),
snapshot/restored around the run. cudnn.benchmark is scoped to a
caller opt-in only: autotuning the fp32 VAE convs doubled peak VRAM
on the DiT families for zero steady-state gain.
- Vectorized sigma gathering (drops a per-step Python search loop),
cached FLUX img_ids/guidance, fused torch AdamW fallback, steady-state
samples_per_second (excludes the first-step warmup).
- Regional torch.compile plumbing (compile_transformer off/on/auto with
eager fallback): auto stays off over a bitsandbytes base where compile
is a net loss (27 s warmup, slightly slower steady on Z-Image); it
arms automatically for the dense/quantized speed modes that follow.
- Stop parity with the LLM trainer: /api/train/diffusion/stop accepts an
optional {save} body and the service forwards save=False as a
no-save cancel; a new preparing event surfaces cache-build progress.
- SDXL trainer gets the same latent cache, perf flags, and fused
fallback; its batching, LR schedule, and min-SNR stay as they were.
Verified: 83 backend tests green; per-family 30-40 step runs with
adapter round-trip generation through the normal LoRA path (FLUX,
Qwen-Image, Z-Image all pass).
627 lines
23 KiB
Python
627 lines
23 KiB
Python
# SPDX-License-Identifier: AGPL-3.0-only
|
|
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
|
|
|
"""Tests for the diffusion LoRA training service + routes.
|
|
|
|
The service's subprocess context and target are injected with in-thread fakes, so the
|
|
full start -> event-pump -> status -> complete path is exercised without real
|
|
multiprocessing or torch. The routes are hit with the FastAPI TestClient and a mocked
|
|
service, so wiring / validation / error mapping are covered without a GPU.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import queue as _queue
|
|
import threading
|
|
import time
|
|
|
|
import pytest
|
|
from fastapi import FastAPI
|
|
from fastapi.testclient import TestClient
|
|
|
|
from auth.authentication import get_current_subject
|
|
from core.training.diffusion_training_service import DiffusionTrainingService
|
|
from routes.training import router as training_router
|
|
|
|
|
|
# ── fake spawn context (runs the "process" target on a thread) ────────────────
|
|
class _FakeQueue:
|
|
def __init__(self) -> None:
|
|
self._q: _queue.Queue = _queue.Queue()
|
|
|
|
def put(self, x):
|
|
self._q.put(x)
|
|
|
|
def get(self, timeout = None):
|
|
return self._q.get(timeout = timeout) # raises queue.Empty on timeout
|
|
|
|
def get_nowait(self):
|
|
return self._q.get_nowait()
|
|
|
|
def empty(self):
|
|
return self._q.empty()
|
|
|
|
|
|
class _FakeProc:
|
|
def __init__(self, target, kwargs, daemon):
|
|
self._target = target
|
|
self._kwargs = kwargs
|
|
self._thread: threading.Thread | None = None
|
|
self.pid = 4321
|
|
|
|
def start(self):
|
|
self._thread = threading.Thread(target = self._target, kwargs = self._kwargs, daemon = True)
|
|
self._thread.start()
|
|
|
|
def is_alive(self):
|
|
return self._thread is not None and self._thread.is_alive()
|
|
|
|
|
|
class _FakeCtx:
|
|
def Queue(self):
|
|
return _FakeQueue()
|
|
|
|
def Process(self, target, kwargs, daemon):
|
|
return _FakeProc(target, kwargs, daemon)
|
|
|
|
|
|
def _happy_target(*, event_queue, stop_queue, config):
|
|
event_queue.put({"type": "model_load_started", "num_images": 3})
|
|
event_queue.put({"type": "model_load_completed"})
|
|
event_queue.put(
|
|
{
|
|
"type": "progress",
|
|
"step": 1,
|
|
"total_steps": 2,
|
|
"loss": 0.5,
|
|
"avg_loss": 0.5,
|
|
"learning_rate": 1e-4,
|
|
}
|
|
)
|
|
event_queue.put(
|
|
{
|
|
"type": "progress",
|
|
"step": 2,
|
|
"total_steps": 2,
|
|
"loss": 0.4,
|
|
"avg_loss": 0.45,
|
|
"learning_rate": 1e-4,
|
|
}
|
|
)
|
|
event_queue.put(
|
|
{
|
|
"type": "complete",
|
|
"output_dir": config["output_dir"],
|
|
"lora_path": config["output_dir"] + "/pytorch_lora_weights.safetensors",
|
|
"stopped": False,
|
|
}
|
|
)
|
|
|
|
|
|
def _stoppable_target(*, event_queue, stop_queue, config):
|
|
event_queue.put({"type": "model_load_completed"})
|
|
stop_queue.get(timeout = 5.0) # block until stop() signals
|
|
event_queue.put(
|
|
{"type": "complete", "output_dir": config["output_dir"], "lora_path": "x", "stopped": True}
|
|
)
|
|
|
|
|
|
def _crashing_target(*, event_queue, stop_queue, config):
|
|
event_queue.put({"type": "model_load_started"})
|
|
# Exits without a terminal event -> the pump must mark it as an error.
|
|
|
|
|
|
_CFG = {"base_model": "b", "data_dir": "d", "output_dir": "/tmp/out", "train_steps": 2}
|
|
|
|
|
|
def _wait_status(
|
|
svc,
|
|
*terminal,
|
|
timeout = 3.0,
|
|
):
|
|
end = time.time() + timeout
|
|
while time.time() < end:
|
|
st = svc.status()
|
|
if st["status"] in terminal:
|
|
return st
|
|
time.sleep(0.02)
|
|
return svc.status()
|
|
|
|
|
|
def test_service_happy_path():
|
|
svc = DiffusionTrainingService(ctx = _FakeCtx(), target = _happy_target)
|
|
job_id = svc.start(dict(_CFG))
|
|
assert job_id
|
|
st = _wait_status(svc, "completed")
|
|
assert st["status"] == "completed"
|
|
assert st["step"] == 2 and st["total_steps"] == 2
|
|
assert st["num_images"] == 3
|
|
assert st["loss"] == 0.4 and st["avg_loss"] == 0.45
|
|
assert st["lora_path"].endswith("pytorch_lora_weights.safetensors")
|
|
assert st["active"] is False
|
|
|
|
|
|
def test_service_rejects_bad_config_before_spawn():
|
|
svc = DiffusionTrainingService(ctx = _FakeCtx(), target = _happy_target)
|
|
with pytest.raises(ValueError):
|
|
svc.start({**_CFG, "train_steps": 0})
|
|
# Nothing was spawned; still idle.
|
|
assert svc.status()["status"] == "idle"
|
|
|
|
|
|
def test_service_rejects_second_concurrent_job():
|
|
svc = DiffusionTrainingService(ctx = _FakeCtx(), target = _stoppable_target)
|
|
svc.start(dict(_CFG))
|
|
_wait_status(svc, "running")
|
|
with pytest.raises(RuntimeError):
|
|
svc.start(dict(_CFG))
|
|
assert svc.stop() is True
|
|
_wait_status(svc, "stopped")
|
|
|
|
|
|
def test_service_stop_marks_stopped():
|
|
svc = DiffusionTrainingService(ctx = _FakeCtx(), target = _stoppable_target)
|
|
svc.start(dict(_CFG))
|
|
_wait_status(svc, "running")
|
|
assert svc.stop() is True
|
|
st = _wait_status(svc, "stopped")
|
|
assert st["status"] == "stopped"
|
|
assert st["active"] is False
|
|
# Stopping again when idle is a no-op.
|
|
assert svc.stop() is False
|
|
|
|
|
|
def test_service_crash_without_terminal_event_is_error():
|
|
svc = DiffusionTrainingService(ctx = _FakeCtx(), target = _crashing_target)
|
|
svc.start(dict(_CFG))
|
|
st = _wait_status(svc, "error")
|
|
assert st["status"] == "error"
|
|
assert "unexpectedly" in st["message"]
|
|
|
|
|
|
def test_apply_event_transitions():
|
|
svc = DiffusionTrainingService(ctx = _FakeCtx(), target = _happy_target)
|
|
svc._apply_event({"type": "model_load_started", "num_images": 5})
|
|
assert svc.status()["in_model_load"] is True and svc.status()["num_images"] == 5
|
|
svc._apply_event({"type": "model_load_completed"})
|
|
assert svc.status()["in_model_load"] is False
|
|
svc._apply_event({"type": "error", "message": "boom"})
|
|
assert svc.status()["status"] == "error" and svc.status()["message"] == "boom"
|
|
|
|
|
|
def test_terminal_events_clear_model_load_flag():
|
|
# A stop or error during model load emits complete/error WITHOUT a preceding
|
|
# model_load_completed, so the terminal update must reset in_model_load or the
|
|
# client shows a stale loading indicator after the job ended.
|
|
svc = DiffusionTrainingService(ctx = _FakeCtx(), target = _happy_target)
|
|
svc._apply_event({"type": "model_load_started"})
|
|
assert svc.status()["in_model_load"] is True
|
|
svc._apply_event({"type": "complete", "stopped": True})
|
|
assert svc.status()["in_model_load"] is False and svc.status()["status"] == "stopped"
|
|
|
|
svc2 = DiffusionTrainingService(ctx = _FakeCtx(), target = _happy_target)
|
|
svc2._apply_event({"type": "model_load_started"})
|
|
svc2._apply_event({"type": "error", "message": "load failed"})
|
|
assert svc2.status()["in_model_load"] is False and svc2.status()["status"] == "error"
|
|
|
|
|
|
# ── route wiring (mocked service) ─────────────────────────────────────────────
|
|
class _FakeService:
|
|
def __init__(self):
|
|
self._running = False
|
|
self.started_with = None
|
|
self.stopped_with_save = None
|
|
# Extra keys merged into status() so a test can inject metric history / perf fields.
|
|
self.status_extra: dict = {}
|
|
|
|
def start(self, config):
|
|
self.started_with = config
|
|
self._running = True
|
|
return "job-123"
|
|
|
|
def stop(self, save = True):
|
|
self.stopped_with_save = save
|
|
was = self._running
|
|
self._running = False
|
|
return was
|
|
|
|
def status(self):
|
|
return {
|
|
"active": self._running,
|
|
"job_id": "job-123" if self._running else None,
|
|
"status": "running" if self._running else "idle",
|
|
"message": "",
|
|
"step": 1,
|
|
"total_steps": 2,
|
|
"loss": 0.5,
|
|
"avg_loss": 0.5,
|
|
"learning_rate": 1e-4,
|
|
"num_images": 3,
|
|
"in_model_load": False,
|
|
"output_dir": None,
|
|
"lora_path": None,
|
|
"started_at": None,
|
|
"updated_at": None,
|
|
**self.status_extra,
|
|
}
|
|
|
|
|
|
class _FakeLLMBackend:
|
|
def __init__(self, active = False):
|
|
self._active = active
|
|
|
|
def is_training_active(self):
|
|
return self._active
|
|
|
|
|
|
@pytest.fixture
|
|
def client(monkeypatch):
|
|
fake = _FakeService()
|
|
monkeypatch.setattr(
|
|
"core.training.diffusion_training_service.get_diffusion_training_service", lambda: fake
|
|
)
|
|
# Neutralize the LLM interlock + GPU-free for the wiring tests (their own tests below
|
|
# exercise those behaviors). The route imports get_training_backend at module scope.
|
|
import routes.training as tr
|
|
|
|
monkeypatch.setattr(tr, "get_training_backend", lambda: _FakeLLMBackend(active = False))
|
|
monkeypatch.setattr(tr, "_free_gpu_for_diffusion_training", lambda: None)
|
|
app = FastAPI()
|
|
app.include_router(training_router, prefix = "/api/train")
|
|
app.dependency_overrides[get_current_subject] = lambda: "test-user"
|
|
c = TestClient(app)
|
|
c._fake = fake # type: ignore[attr-defined]
|
|
return c
|
|
|
|
|
|
# Studio-relative paths: the route resolves/contains them before spawn.
|
|
_BODY = {
|
|
"base_model": "stabilityai/sdxl-turbo",
|
|
"data_dir": "uploads/my-images",
|
|
"output_dir": "my-lora-run",
|
|
"train_steps": 10,
|
|
}
|
|
|
|
|
|
def test_route_start_ok(client):
|
|
r = client.post("/api/train/diffusion/start", json = _BODY)
|
|
assert r.status_code == 200, r.text
|
|
assert r.json() == {"job_id": "job-123", "status": "running"}
|
|
assert client._fake.started_with["base_model"] == "stabilityai/sdxl-turbo"
|
|
# Paths were resolved to absolute Studio-contained locations before spawn.
|
|
from pathlib import Path
|
|
|
|
assert Path(client._fake.started_with["data_dir"]).is_absolute()
|
|
assert Path(client._fake.started_with["output_dir"]).is_absolute()
|
|
|
|
|
|
def test_route_start_forwards_extra_training_knobs(client):
|
|
# max_grad_norm and lora_target_modules must reach the service, not be silently dropped.
|
|
body = {**_BODY, "max_grad_norm": 0.5, "lora_target_modules": ["to_q", "to_v"]}
|
|
r = client.post("/api/train/diffusion/start", json = body)
|
|
assert r.status_code == 200, r.text
|
|
assert client._fake.started_with["max_grad_norm"] == 0.5
|
|
assert client._fake.started_with["lora_target_modules"] == ["to_q", "to_v"]
|
|
|
|
|
|
def test_route_start_rejects_uncontained_paths(client):
|
|
# An absolute path outside the Studio dataset roots is a 400, not silently accepted.
|
|
r = client.post("/api/train/diffusion/start", json = {**_BODY, "data_dir": "/etc"})
|
|
assert r.status_code == 400
|
|
|
|
|
|
def test_route_start_blocked_by_active_llm_training(client, monkeypatch):
|
|
import routes.training as tr
|
|
|
|
monkeypatch.setattr(tr, "get_training_backend", lambda: _FakeLLMBackend(active = True))
|
|
r = client.post("/api/train/diffusion/start", json = _BODY)
|
|
assert r.status_code == 409
|
|
assert "LLM training" in r.json()["detail"]
|
|
|
|
|
|
def test_route_start_missing_required_is_422(client):
|
|
r = client.post(
|
|
"/api/train/diffusion/start", json = {"base_model": "x"}
|
|
) # no data_dir/output_dir
|
|
assert r.status_code == 422
|
|
|
|
|
|
def test_route_start_bad_config_maps_to_400(client, monkeypatch):
|
|
def _raise(_cfg):
|
|
raise ValueError("resolution must be a multiple of 8")
|
|
|
|
client._fake.start = _raise # type: ignore[assignment]
|
|
r = client.post("/api/train/diffusion/start", json = _BODY)
|
|
assert r.status_code == 400
|
|
assert "multiple of 8" in r.json()["detail"]
|
|
|
|
|
|
def test_route_start_conflict_maps_to_409(client):
|
|
def _raise(_cfg):
|
|
raise RuntimeError("A diffusion training job is already running.")
|
|
|
|
client._fake.start = _raise # type: ignore[assignment]
|
|
r = client.post("/api/train/diffusion/start", json = _BODY)
|
|
assert r.status_code == 409
|
|
|
|
|
|
def test_route_status_and_stop(client):
|
|
client.post("/api/train/diffusion/start", json = _BODY)
|
|
s = client.get("/api/train/diffusion/status")
|
|
assert s.status_code == 200 and s.json()["status"] == "running"
|
|
st = client.post("/api/train/diffusion/stop")
|
|
assert st.status_code == 200 and st.json()["status"] == "stopping"
|
|
# After stopping, a stop with nothing running reports idle.
|
|
st2 = client.post("/api/train/diffusion/stop")
|
|
assert st2.json()["status"] == "idle"
|
|
|
|
|
|
def test_service_restart_after_completion():
|
|
# A finished job's pump is joined OUTSIDE the lock (it needs the lock for its
|
|
# final state writes), so a second start neither stalls nor deadlocks.
|
|
svc = DiffusionTrainingService(ctx = _FakeCtx(), target = _happy_target)
|
|
svc.start(dict(_CFG))
|
|
_wait_status(svc, "completed")
|
|
t0 = time.time()
|
|
job2 = svc.start(dict(_CFG))
|
|
assert job2
|
|
assert time.time() - t0 < 4.0 # no 5s join-under-lock stall
|
|
st = _wait_status(svc, "completed")
|
|
assert st["status"] == "completed"
|
|
|
|
|
|
def test_stale_pump_events_cannot_corrupt_new_job():
|
|
# An event carrying a superseded job's proc identity must be dropped, so a
|
|
# straggler pump can never overwrite the state of a newly started job.
|
|
svc = DiffusionTrainingService(ctx = _FakeCtx(), target = _happy_target)
|
|
svc.start(dict(_CFG))
|
|
_wait_status(svc, "completed")
|
|
current = svc._proc
|
|
svc._apply_event({"type": "error", "message": "stale boom"}, proc = object())
|
|
assert svc.status()["message"] != "stale boom"
|
|
# The current job's events still apply.
|
|
svc._apply_event({"type": "progress", "step": 9}, proc = current)
|
|
assert svc.status()["step"] == 9
|
|
|
|
|
|
# ── /diffusion/info + /diffusion/dataset (dataset discovery + upload) ─────────
|
|
@pytest.fixture
|
|
def dataset_roots(client, monkeypatch, tmp_path):
|
|
# The endpoints import these lazily per-request, so patching the package attr works.
|
|
import utils.paths as up
|
|
|
|
ds_root = tmp_path / "assets" / "datasets"
|
|
out_root = tmp_path / "outputs"
|
|
ds_root.mkdir(parents = True)
|
|
out_root.mkdir(parents = True)
|
|
monkeypatch.setattr(up, "datasets_root", lambda: ds_root)
|
|
monkeypatch.setattr(up, "outputs_root", lambda: out_root)
|
|
return ds_root, out_root
|
|
|
|
|
|
def test_diffusion_info_lists_image_dataset_folders(client, dataset_roots):
|
|
ds_root, out_root = dataset_roots
|
|
good = ds_root / "cat-photos"
|
|
good.mkdir()
|
|
(good / "a.png").write_bytes(b"x")
|
|
(good / "b.jpg").write_bytes(b"x")
|
|
(good / "a.txt").write_text("a cat")
|
|
(ds_root / "empty-dir").mkdir() # no images -> not a dataset
|
|
(ds_root / "stray.txt").write_text("not a folder")
|
|
|
|
r = client.get("/api/train/diffusion/info")
|
|
assert r.status_code == 200, r.text
|
|
body = r.json()
|
|
assert body["datasets_root"] == str(ds_root)
|
|
assert body["outputs_root"] == str(out_root)
|
|
assert [d["name"] for d in body["datasets"]] == ["cat-photos"]
|
|
assert body["datasets"][0]["image_count"] == 2
|
|
assert body["datasets"][0]["caption_count"] == 1
|
|
|
|
|
|
def test_diffusion_dataset_upload_accumulates(client, dataset_roots):
|
|
ds_root, _ = dataset_roots
|
|
files = [
|
|
("files", ("a.png", b"png-bytes", "image/png")),
|
|
("files", ("b.JPG", b"jpg-bytes", "image/jpeg")),
|
|
("files", ("a.txt", b"a caption", "text/plain")),
|
|
]
|
|
r = client.post("/api/train/diffusion/dataset", data = {"name": "my style"}, files = files)
|
|
assert r.status_code == 200, r.text
|
|
body = r.json()
|
|
assert body["name"] == "my style"
|
|
assert body["uploaded"] == 3
|
|
assert body["image_count"] == 2
|
|
assert body["caption_count"] == 1
|
|
assert (ds_root / "my style" / "a.png").read_bytes() == b"png-bytes"
|
|
|
|
# A second batch into the same name accumulates (large sets arrive in chunks).
|
|
r = client.post(
|
|
"/api/train/diffusion/dataset",
|
|
data = {"name": "my style"},
|
|
files = [("files", ("c.webp", b"w", "image/webp"))],
|
|
)
|
|
assert r.status_code == 200, r.text
|
|
assert r.json()["uploaded"] == 1
|
|
assert r.json()["image_count"] == 3
|
|
|
|
|
|
def test_diffusion_dataset_upload_rejects_traversal_names(client, dataset_roots):
|
|
for bad in ("../evil", "a/b", ".hidden", " "):
|
|
r = client.post(
|
|
"/api/train/diffusion/dataset",
|
|
data = {"name": bad},
|
|
files = [("files", ("a.png", b"x", "image/png"))],
|
|
)
|
|
assert r.status_code == 400, f"{bad!r}: {r.status_code}"
|
|
|
|
|
|
def test_diffusion_dataset_upload_rejects_unsupported_files(client, dataset_roots):
|
|
r = client.post(
|
|
"/api/train/diffusion/dataset",
|
|
data = {"name": "ok-name"},
|
|
files = [("files", ("weights.exe", b"mz", "application/octet-stream"))],
|
|
)
|
|
assert r.status_code == 400
|
|
assert "Unsupported file" in r.json()["detail"]
|
|
|
|
|
|
def test_route_start_refuses_non_sdxl_base_without_freeing_gpu(client, monkeypatch):
|
|
# A doomed start (non-SDXL base) must 400 BEFORE resident GPU workloads are freed,
|
|
# so a bad pick never unloads the user's working chat/Images model.
|
|
import routes.training as tr
|
|
|
|
freed = []
|
|
monkeypatch.setattr(tr, "_free_gpu_for_diffusion_training", lambda: freed.append(1))
|
|
r = client.post(
|
|
"/api/train/diffusion/start", json = {**_BODY, "base_model": "unsloth/FLUX.1-dev-GGUF"}
|
|
)
|
|
assert r.status_code == 400
|
|
assert "SDXL" in r.json()["detail"]
|
|
assert freed == []
|
|
assert client._fake.started_with is None
|
|
|
|
|
|
# ── metric history + perf/family fields (PR A platform) ──────────────────────
|
|
def test_apply_event_records_metric_history_and_perf():
|
|
svc = DiffusionTrainingService(ctx = _FakeCtx(), target = _happy_target)
|
|
svc._apply_event(
|
|
{
|
|
"type": "progress",
|
|
"step": 1,
|
|
"total_steps": 10,
|
|
"loss": 0.5,
|
|
"learning_rate": 1e-4,
|
|
"samples_per_second": 3.2,
|
|
"peak_memory_gb": 7.1,
|
|
}
|
|
)
|
|
svc._apply_event(
|
|
{"type": "progress", "step": 2, "total_steps": 10, "loss": 0.4, "learning_rate": 9e-5}
|
|
)
|
|
st = svc.status()
|
|
assert st["metric_steps"] == [1, 2]
|
|
assert st["metric_loss"] == [0.5, 0.4]
|
|
assert st["metric_lr"] == [1e-4, 9e-5]
|
|
assert st["samples_per_second"] == 3.2
|
|
assert st["peak_memory_gb"] == 7.1
|
|
|
|
|
|
def test_apply_event_metric_history_skips_bad_points():
|
|
svc = DiffusionTrainingService(ctx = _FakeCtx(), target = _happy_target)
|
|
# step 0 (warmup / no real step), a None loss, and a NaN loss must all be skipped.
|
|
svc._apply_event({"type": "progress", "step": 0, "loss": 0.9, "learning_rate": 1e-4})
|
|
svc._apply_event({"type": "progress", "step": 1, "loss": None, "learning_rate": 1e-4})
|
|
svc._apply_event({"type": "progress", "step": 2, "loss": float("nan"), "learning_rate": 1e-4})
|
|
svc._apply_event({"type": "progress", "step": 3, "loss": 0.3, "learning_rate": None})
|
|
st = svc.status()
|
|
assert st["metric_steps"] == [3]
|
|
assert st["metric_loss"] == [0.3]
|
|
assert st["metric_lr"] == [None] # lr None is retained so the series stays index-aligned
|
|
|
|
|
|
def test_metric_history_decimates_at_cap():
|
|
from core.training import diffusion_training_service as svc_mod
|
|
|
|
svc = DiffusionTrainingService(ctx = _FakeCtx(), target = _happy_target)
|
|
svc._apply_event({"type": "model_load_started"}) # initialise running state
|
|
n = svc_mod._METRIC_CAP + 50
|
|
for i in range(1, n + 1):
|
|
svc._apply_event({"type": "progress", "step": i, "loss": 1.0 / i, "learning_rate": 1e-4})
|
|
st = svc.status()
|
|
# Never exceeds the cap, and stays a valid paired history with matching lengths.
|
|
assert len(st["metric_steps"]) <= svc_mod._METRIC_CAP
|
|
assert len(st["metric_steps"]) == len(st["metric_loss"]) == len(st["metric_lr"])
|
|
# Decimation keeps the curve monotonic in step (still increasing, just sparser).
|
|
assert st["metric_steps"] == sorted(st["metric_steps"])
|
|
assert st["metric_steps"][-1] == n # the latest point is always retained
|
|
|
|
|
|
def test_complete_event_records_family_and_catalog():
|
|
svc = DiffusionTrainingService(ctx = _FakeCtx(), target = _happy_target)
|
|
svc._apply_event(
|
|
{
|
|
"type": "complete",
|
|
"output_dir": "/o",
|
|
"lora_path": "/o/w.safetensors",
|
|
"catalog_path": "/loras/w.safetensors",
|
|
"family": "sdxl",
|
|
"base_model": "b",
|
|
}
|
|
)
|
|
st = svc.status()
|
|
assert st["status"] == "completed"
|
|
assert st["family"] == "sdxl"
|
|
assert st["base_model"] == "b"
|
|
assert st["catalog_path"] == "/loras/w.safetensors"
|
|
|
|
|
|
def test_status_route_nests_metric_history(client):
|
|
# The status route folds the service's flat arrays into a nested metric_history object.
|
|
client._fake.status_extra = {
|
|
"metric_steps": [1, 2],
|
|
"metric_loss": [0.5, 0.4],
|
|
"metric_lr": [1e-4, 9e-5],
|
|
"family": "sdxl",
|
|
"samples_per_second": 2.0,
|
|
"peak_memory_gb": 6.0,
|
|
}
|
|
r = client.get("/api/train/diffusion/status")
|
|
assert r.status_code == 200, r.text
|
|
body = r.json()
|
|
assert body["metric_history"]["steps"] == [1, 2]
|
|
assert body["metric_history"]["loss"] == [0.5, 0.4]
|
|
assert body["family"] == "sdxl"
|
|
assert body["samples_per_second"] == 2.0
|
|
|
|
|
|
# ── /diffusion/info families + gated-repo preflight (PR B) ──────────────────────
|
|
def test_info_lists_trainable_families(client):
|
|
r = client.get("/api/train/diffusion/info")
|
|
assert r.status_code == 200, r.text
|
|
families = {f["name"]: f for f in r.json()["families"]}
|
|
for fam in ("sdxl", "flux.1", "qwen-image", "z-image"):
|
|
assert fam in families
|
|
assert families["flux.1"]["default_base"] == "black-forest-labs/FLUX.1-dev"
|
|
assert families["z-image"]["defaults"]["resolution"] == 768
|
|
|
|
|
|
def test_start_gated_base_without_access_is_400_and_keeps_gpu(client, monkeypatch):
|
|
# A gated FLUX base with no valid token must 400 from the HEAD preflight BEFORE the GPU
|
|
# residents are freed, so a doomed start never evicts the user's loaded model.
|
|
import urllib.error
|
|
import urllib.request
|
|
|
|
import routes.training as tr
|
|
|
|
freed: list[int] = []
|
|
monkeypatch.setattr(tr, "_free_gpu_for_diffusion_training", lambda: freed.append(1))
|
|
|
|
def _fake_urlopen(req, timeout = None):
|
|
raise urllib.error.HTTPError(req.full_url, 403, "Forbidden", {}, None)
|
|
|
|
monkeypatch.setattr(urllib.request, "urlopen", _fake_urlopen)
|
|
r = client.post(
|
|
"/api/train/diffusion/start",
|
|
json = {**_BODY, "base_model": "black-forest-labs/FLUX.1-dev"},
|
|
)
|
|
assert r.status_code == 400
|
|
assert "gated" in r.json()["detail"].lower()
|
|
assert freed == []
|
|
assert client._fake.started_with is None
|
|
|
|
|
|
def test_start_ungated_base_preflight_is_noop(client, monkeypatch):
|
|
# A reachable base (HEAD 200) proceeds to start normally.
|
|
import urllib.request
|
|
|
|
import routes.training as tr
|
|
|
|
monkeypatch.setattr(tr, "_free_gpu_for_diffusion_training", lambda: None)
|
|
monkeypatch.setattr(urllib.request, "urlopen", lambda req, timeout = None: object())
|
|
r = client.post(
|
|
"/api/train/diffusion/start",
|
|
json = {**_BODY, "base_model": "black-forest-labs/FLUX.1-dev"},
|
|
)
|
|
assert r.status_code == 200, r.text
|
|
assert client._fake.started_with["base_model"] == "black-forest-labs/FLUX.1-dev"
|