Fix review findings: engine unload before training, caption precedence, dataset caption counts, local diffusers tagging, family LoRA targets

- Unload the ACTIVE image engine (sd_cpp or diffusers) before diffusion training starts, not just the diffusers singleton
- Count metadata.jsonl captions in dataset summaries so metadata-captioned datasets are not reported as uncaptioned
- Sidecar captions now override metadata rows everywhere (grid edits win); trainer and dataset API agree
- Tag local diffusers image checkpoints with text-to-image so they appear in the Images picker
- Family LoRA targets (_FLUX_TARGETS etc) apply when the config carries the generic defaults; explicit overrides still win
This commit is contained in:
Daniel Han 2026-07-04 03:23:05 +00:00
commit bfbb902610
9 changed files with 259 additions and 59 deletions

View file

@ -32,6 +32,7 @@ from typing import Any, Callable, Optional
from core.training.diffusion_train_common import (
DEFAULT_LORA_FILENAME,
DEFAULT_LORA_TARGETS,
DiffusionLoraConfig,
EventCb,
StopCb,
@ -58,6 +59,20 @@ _QWEN_TARGETS = _FLUX_TARGETS
_ZIMAGE_TARGETS = ("to_q", "to_k", "to_v", "to_out.0")
def _select_lora_targets(
cfg_targets: tuple[str, ...], spec_targets: tuple[str, ...]
) -> tuple[str, ...]:
"""Pick the LoRA target modules for a DiT run.
``normalized()`` always fills ``lora_target_modules`` with the generic
``DEFAULT_LORA_TARGETS`` when a caller does not set it, so that value means "unset"
here: prefer the family's ``spec.lora_targets`` (which add the DiT-specific
projections). Any OTHER explicit tuple is a deliberate override and still wins."""
if tuple(cfg_targets) == DEFAULT_LORA_TARGETS:
return tuple(spec_targets)
return tuple(cfg_targets)
@dataclass
class _FamilySpec:
"""Everything the shared loop needs that differs by family."""
@ -513,7 +528,7 @@ def run_dit_lora_training(
# The flow-matching + 4-bit path is bf16 throughout (fp32 on a CPU-only box, which is
# unsupported for real runs but keeps import/unit tests architecture-agnostic).
weight_dtype = torch.bfloat16 if device == "cuda" else torch.float32
use_lora_targets = tuple(cfg.lora_target_modules) or spec.lora_targets
use_lora_targets = _select_lora_targets(cfg.lora_target_modules, spec.lora_targets)
_assert_trusted_base_model(cfg.base_model)
_assert_gated_access(cfg.base_model, cfg.hf_token)

View file

@ -292,11 +292,15 @@ def discover_image_caption_pairs(
"""Resolve ``(image_path, caption)`` pairs from a dataset directory.
Caption sources, in priority order per image:
1. a ``metadata.jsonl`` / ``captions.jsonl`` row keyed by ``file_name`` (or ``image``)
1. a per-image sidecar ``<stem>.txt`` / ``<stem>.caption``,
2. a ``metadata.jsonl`` / ``captions.jsonl`` row keyed by ``file_name`` (or ``image``)
carrying the caption in ``caption_column`` (default ``text``),
2. a per-image sidecar ``<stem>.txt`` / ``<stem>.caption``,
3. ``instance_prompt`` (dreambooth) for any remaining image.
A sidecar wins over the metadata row because it is the user's explicit per-image edit
(the labeling grid writes a .txt sidecar), which must override the bulk metadata file.
Must agree with ``routes.training._image_record``, which resolves captions the same way.
Images with no caption from any source are skipped. Pure filesystem + JSON, so it is
unit-testable without torch. Raises FileNotFoundError for a missing dir and ValueError
when nothing is captionable.
@ -328,15 +332,15 @@ def discover_image_caption_pairs(
pairs: list[tuple[str, str]] = []
for img in images:
caption: Optional[str] = None
# 1. metadata row keyed by file name (basename or the name as written).
caption = meta_caption.get(img.name) or meta_caption.get(str(img.relative_to(root)))
# 2. per-image sidecar caption file.
# 1. per-image sidecar caption file (the user's explicit edit; wins over metadata).
for ext in _CAPTION_EXTS:
sidecar = img.with_suffix(ext)
if sidecar.is_file():
caption = sidecar.read_text(encoding = "utf-8").strip()
break
# 2. metadata row keyed by file name (basename or the name as written).
if caption is None:
for ext in _CAPTION_EXTS:
sidecar = img.with_suffix(ext)
if sidecar.is_file():
caption = sidecar.read_text(encoding = "utf-8").strip()
break
caption = meta_caption.get(img.name) or meta_caption.get(str(img.relative_to(root)))
# 3. dreambooth instance prompt.
if caption is None and instance_prompt:
caption = instance_prompt

View file

@ -909,9 +909,10 @@ async def list_local_models(
try:
models = collect_local_models(models_root)
# Tag each GGUF with its task so the Images picker can filter to diffusion.
# Tag each model with its task so the Images picker can filter to diffusion
# (GGUF by architecture; local diffusers checkpoints by pipeline / family).
models = [
m.model_copy(update = {"task": _local_model_task(m.path, m.model_format)}) for m in models
m.model_copy(update = {"task": _local_model_task(m)}) for m in models
]
return LocalModelListResponse(
@ -3229,24 +3230,55 @@ def _repo_gguf_task(repo_info) -> Optional[str]:
return None
def _local_model_task(path: str, model_format: Optional[str]) -> Optional[str]:
"""Same classification for a local model: read its GGUF architecture. The
path may be the .gguf file itself or a folder containing one."""
if model_format != "gguf":
def _local_model_task(model: "LocalModelInfo") -> Optional[str]:
"""Classify a local model into an HF pipeline task so the Images picker can filter.
For a GGUF, read its architecture (the path may be the .gguf file itself or a folder
containing one). For a local non-GGUF image checkpoint (a diffusers pipeline dir or a
single-file safetensors), fall through to the diffusers detection so on-device image
models get the 'text-to-image' tag instead of being dropped as task=null; the load
path accepts these as a local pipeline."""
path = model.path
if model.model_format == "gguf":
try:
p = Path(path)
if p.suffix.lower() == ".gguf" and p.is_file():
return _arch_to_task(_gguf_architecture(str(p)))
for f in _iter_gguf_paths(p):
if _is_mmproj_filename(f.name):
continue
task = _arch_to_task(_gguf_architecture(str(f)))
if task is not None:
return task
except Exception:
pass
return None
if _local_is_diffusers(model):
return "text-to-image"
return None
def _local_is_diffusers(model: "LocalModelInfo") -> bool:
"""True for a local diffusers image checkpoint, mirroring the cached-repo
``_repo_is_diffusers`` heuristics: a full pipeline carries a top-level
``model_index.json``, while single-file / safetensors image checkpoints ship none, so
fall back to the model id resolving to a known diffusion family (the same resolver the
Images backend loads from). Family detection uses the clean model id / name, not the
on-disk path, so a parent directory keyword can't spuriously match."""
try:
p = Path(path)
if p.suffix.lower() == ".gguf" and p.is_file():
return _arch_to_task(_gguf_architecture(str(p)))
for f in _iter_gguf_paths(p):
if _is_mmproj_filename(f.name):
continue
task = _arch_to_task(_gguf_architecture(str(f)))
if task is not None:
return task
p = Path(model.path)
if p.is_dir() and (p / "model_index.json").is_file():
return True
except Exception:
pass
return None
try:
from core.inference.diffusion_families import detect_family
for needle in (model.model_id, model.display_name, model.id):
if needle and detect_family(needle) is not None:
return True
except Exception:
pass
return False
@router.get("/cached-gguf")

View file

@ -1111,9 +1111,13 @@ def _free_gpu_for_diffusion_training() -> None:
try:
from core.inference import gpu_arbiter
from core.inference.diffusion import get_diffusion_backend
from core.inference.diffusion_engine_router import get_active_diffusion_engine
diffusion = get_diffusion_backend()
# The ACTIVE engine, not the diffusers singleton: on a native (sd_cpp)
# selection the diffusers backend reports unloaded while the resident
# sd-server still holds the GPU, so unloading only the singleton is a no-op.
# Mirrors the LLM training start path.
diffusion = get_active_diffusion_engine()
if diffusion.is_loaded:
logger.info("Unloading resident Images pipeline to free GPU memory for training")
diffusion.unload() # no-op when nothing is loaded; also preempts an in-flight load
@ -1274,14 +1278,19 @@ _DIFFUSION_DATASET_TEXT_EXTS = {".txt", ".caption", ".jsonl"}
def _diffusion_dataset_summary(folder: Path) -> DiffusionDatasetSummary:
# Count an image as captioned when a metadata/captions.jsonl row or a per-image
# sidecar (.txt / .caption) resolves a caption for it -- the same sources the
# trainer reads. Counting metadata-only captions here keeps a metadata-captioned
# dataset from reporting caption_count=0 and being treated as uncaptioned.
meta_captions = _load_metadata_captions(folder)
images = captions = 0
for f in folder.iterdir():
if not f.is_file():
if not f.is_file() or f.suffix.lower() not in _DIFFUSION_DATASET_IMAGE_EXTS:
continue
ext = f.suffix.lower()
if ext in _DIFFUSION_DATASET_IMAGE_EXTS:
images += 1
elif ext in (".txt", ".caption"):
images += 1
if f.name in meta_captions or any(
f.with_suffix(ext).is_file() for ext in (".txt", ".caption")
):
captions += 1
return DiffusionDatasetSummary(
name = folder.name, path = str(folder), image_count = images, caption_count = captions
@ -1478,20 +1487,26 @@ def _load_metadata_captions(folder: Path) -> dict[str, str]:
def _image_record(
folder: Path, image_path: Path, meta_captions: dict[str, str]
) -> DiffusionDatasetImageRecord:
"""Build one image record, resolving its caption with metadata > sidecar precedence
(the same order the trainer uses)."""
caption: Optional[str] = meta_captions.get(image_path.name)
source = "metadata" if caption is not None else "none"
"""Build one image record, resolving its caption with sidecar > metadata precedence
(the same order the trainer uses). A per-image .txt / .caption sidecar wins because
it is the user's explicit edit from the labeling grid, which must override a
metadata.jsonl / captions.jsonl row for the image."""
caption: Optional[str] = None
source = "none"
for ext in (".txt", ".caption"):
sidecar = image_path.with_suffix(ext)
if sidecar.is_file():
try:
caption = sidecar.read_text(encoding = "utf-8").strip()
source = "sidecar"
except OSError:
caption = None
break
if caption is None:
for ext in (".txt", ".caption"):
sidecar = image_path.with_suffix(ext)
if sidecar.is_file():
try:
caption = sidecar.read_text(encoding = "utf-8").strip()
source = "sidecar"
except OSError:
caption = None
break
meta = meta_captions.get(image_path.name)
if meta is not None:
caption = meta
source = "metadata"
try:
size_bytes = image_path.stat().st_size
except OSError:

View file

@ -61,12 +61,14 @@ def test_list_images_caption_precedence(client, ds_root):
_write_png(folder / "a.png")
_write_png(folder / "b.png")
_write_png(folder / "c.png")
# a.png -> metadata (beats a stray sidecar), b.png -> sidecar, c.png -> none.
# a.png -> sidecar (an explicit edit beats the metadata row), b.png -> metadata-only,
# c.png -> none.
(folder / "metadata.jsonl").write_text(
json.dumps({"file_name": "a.png", "text": "from metadata"}) + "\n", encoding = "utf-8"
json.dumps({"file_name": "a.png", "text": "from metadata"}) + "\n"
+ json.dumps({"file_name": "b.png", "text": "from metadata"}) + "\n",
encoding = "utf-8",
)
(folder / "a.txt").write_text("stray sidecar", encoding = "utf-8")
(folder / "b.txt").write_text("from sidecar", encoding = "utf-8")
(folder / "a.txt").write_text("edited sidecar", encoding = "utf-8")
r = client.get("/api/train/diffusion/dataset/styleset/images")
assert r.status_code == 200, r.text
@ -74,10 +76,11 @@ def test_list_images_caption_precedence(client, ds_root):
assert body["name"] == "styleset"
recs = {rec["filename"]: rec for rec in body["images"]}
assert set(recs) == {"a.png", "b.png", "c.png"}
assert recs["a.png"]["caption"] == "from metadata"
assert recs["a.png"]["caption_source"] == "metadata"
assert recs["b.png"]["caption"] == "from sidecar"
assert recs["b.png"]["caption_source"] == "sidecar"
# A sidecar edit overrides the metadata row for the same image.
assert recs["a.png"]["caption"] == "edited sidecar"
assert recs["a.png"]["caption_source"] == "sidecar"
assert recs["b.png"]["caption"] == "from metadata"
assert recs["b.png"]["caption_source"] == "metadata"
assert recs["c.png"]["caption"] is None
assert recs["c.png"]["caption_source"] == "none"
assert recs["a.png"]["width"] == 8 and recs["a.png"]["height"] == 8
@ -133,6 +136,25 @@ def test_put_caption_roundtrip_and_clear(client, ds_root):
assert not (folder / "x.txt").exists()
def test_put_caption_overrides_metadata_row(client, ds_root):
# Editing a caption for an image that already has a metadata.jsonl row must take
# effect: the sidecar edit wins over the metadata caption in the response (and in
# the data the trainer reads), not the other way round.
folder = ds_root / "cap"
folder.mkdir()
_write_png(folder / "x.png")
(folder / "metadata.jsonl").write_text(
json.dumps({"file_name": "x.png", "text": "from metadata"}) + "\n", encoding = "utf-8"
)
r = client.put(
"/api/train/diffusion/dataset/cap/caption/x.png", json = {"caption": "edited caption"}
)
assert r.status_code == 200, r.text
assert r.json()["caption"] == "edited caption"
assert r.json()["caption_source"] == "sidecar"
def test_put_caption_missing_image_404(client, ds_root):
(ds_root / "cap").mkdir()
r = client.put("/api/train/diffusion/dataset/cap/caption/ghost.png", json = {"caption": "hi"})

View file

@ -12,13 +12,21 @@ from __future__ import annotations
import pytest
from core.training.diffusion_dit_trainer import (
_FLUX_TARGETS,
_GATED_TRAIN_REPOS,
_QWEN_TARGETS,
_SPECS,
_ZIMAGE_TARGETS,
_assert_gated_access,
_repo_is_prequantized,
_select_lora_targets,
run_dit_lora_training,
)
from core.training.diffusion_train_common import DiffusionLoraConfig, family_train_infos
from core.training.diffusion_train_common import (
DEFAULT_LORA_TARGETS,
DiffusionLoraConfig,
family_train_infos,
)
def test_specs_cover_the_three_dit_families():
@ -32,6 +40,27 @@ def test_specs_cover_the_three_dit_families():
assert _SPECS["qwen-image"].force_bf16 is True
def test_select_lora_targets_uses_family_default_for_generic_config():
# normalized() fills lora_target_modules with the generic DEFAULT_LORA_TARGETS when a
# caller doesn't set it, so that value must resolve to the family's targets (which add
# the DiT-specific projections), not stay stuck on the generic SDXL list.
assert _select_lora_targets(DEFAULT_LORA_TARGETS, _FLUX_TARGETS) == _FLUX_TARGETS
assert _select_lora_targets(DEFAULT_LORA_TARGETS, _QWEN_TARGETS) == _QWEN_TARGETS
assert _select_lora_targets(DEFAULT_LORA_TARGETS, _ZIMAGE_TARGETS) == _ZIMAGE_TARGETS
def test_select_lora_targets_explicit_override_wins():
# Any OTHER explicit tuple is a deliberate override and must win over the family spec.
override = ("to_q", "to_k")
assert _select_lora_targets(override, _FLUX_TARGETS) == override
# The default request path (config carrying the generic default) reaches the spec.
cfg = DiffusionLoraConfig(
base_model = "black-forest-labs/FLUX.1-dev", data_dir = "d", output_dir = "o"
).normalized()
assert cfg.lora_target_modules == DEFAULT_LORA_TARGETS
assert _select_lora_targets(cfg.lora_target_modules, _SPECS["flux.1"].lora_targets) == _FLUX_TARGETS
@pytest.mark.parametrize(
"repo, expected",
[

View file

@ -27,15 +27,15 @@ def _touch(p):
p.write_bytes(b"")
def test_discover_prefers_metadata_then_sidecar_then_instance(tmp_path):
def test_discover_prefers_sidecar_then_metadata_then_instance(tmp_path):
_touch(tmp_path / "a.png")
_touch(tmp_path / "b.jpg")
_touch(tmp_path / "c.webp")
# a.png captioned via metadata.jsonl
# a.png captioned via metadata.jsonl only
(tmp_path / "metadata.jsonl").write_text(
json.dumps({"file_name": "a.png", "text": "from metadata"}) + "\n", encoding = "utf-8"
)
# b.jpg captioned via sidecar
# b.jpg captioned via sidecar only
(tmp_path / "b.txt").write_text("from sidecar", encoding = "utf-8")
# c.webp falls back to the instance prompt
pairs = dict(discover_image_caption_pairs(tmp_path, instance_prompt = "from instance"))
@ -44,6 +44,18 @@ def test_discover_prefers_metadata_then_sidecar_then_instance(tmp_path):
assert pairs[str(tmp_path / "c.webp")] == "from instance"
def test_discover_sidecar_overrides_metadata_row(tmp_path):
# A per-image sidecar is the user's explicit edit and must win over a metadata row
# for the same image (the labeling grid writes sidecars).
_touch(tmp_path / "a.png")
(tmp_path / "metadata.jsonl").write_text(
json.dumps({"file_name": "a.png", "text": "from metadata"}) + "\n", encoding = "utf-8"
)
(tmp_path / "a.txt").write_text("edited sidecar", encoding = "utf-8")
pairs = dict(discover_image_caption_pairs(tmp_path))
assert pairs[str(tmp_path / "a.png")] == "edited sidecar"
def test_discover_skips_uncaptioned_without_instance_prompt(tmp_path):
_touch(tmp_path / "cap.png")
_touch(tmp_path / "nocap.png")

View file

@ -11,6 +11,7 @@ service, so wiring / validation / error mapping are covered without a GPU.
from __future__ import annotations
import json
import queue as _queue
import threading
import time
@ -417,6 +418,31 @@ def test_diffusion_info_lists_image_dataset_folders(client, dataset_roots):
assert body["datasets"][0]["caption_count"] == 1
def test_diffusion_info_counts_metadata_captions(client, dataset_roots):
# A dataset captioned via metadata.jsonl must not report caption_count=0 (which the
# Train UI treats as uncaptioned); metadata rows count like sidecars, without
# double-counting an image that has both.
ds_root, _ = dataset_roots
folder = ds_root / "meta-captioned"
folder.mkdir()
(folder / "a.png").write_bytes(b"x")
(folder / "b.png").write_bytes(b"x")
(folder / "c.png").write_bytes(b"x")
# a.png + b.png via metadata; a.png also has a sidecar (must count once); c.png none.
(folder / "metadata.jsonl").write_text(
json.dumps({"file_name": "a.png", "text": "cap a"}) + "\n"
+ json.dumps({"file_name": "b.png", "text": "cap b"}) + "\n",
encoding = "utf-8",
)
(folder / "a.txt").write_text("edited a", encoding = "utf-8")
r = client.get("/api/train/diffusion/info")
assert r.status_code == 200, r.text
summary = next(d for d in r.json()["datasets"] if d["name"] == "meta-captioned")
assert summary["image_count"] == 3
assert summary["caption_count"] == 2
def test_diffusion_dataset_upload_accumulates(client, dataset_roots):
ds_root, _ = dataset_roots
files = [

View file

@ -115,3 +115,48 @@ def test_scan_models_dir_classifies_root_gguf_with_config(tmp_path):
assert row.path == str(root)
assert row.model_format == "gguf"
# ── Images picker task tag for local (non-GGUF) diffusers models ──────────────
from models.models import LocalModelInfo # noqa: E402
def _local(path, *, model_format = None, model_id = None, display_name = "m", id = "m"):
return LocalModelInfo(
id = id,
display_name = display_name,
path = str(path),
source = "models_dir",
model_id = model_id,
model_format = model_format,
)
def test_local_task_tags_diffusers_pipeline_dir(tmp_path):
# A local diffusers pipeline (top-level model_index.json) is an image model even
# though its model_format is not "gguf": tag it so the Images picker keeps it.
d = tmp_path / "my-local-pipeline"
_touch(d / "model_index.json")
_touch(d / "unet" / "diffusion_pytorch_model.safetensors")
assert models_route._local_model_task(_local(d)) == "text-to-image"
def test_local_task_tags_diffusers_by_family_id(tmp_path):
# A single-file / safetensors image checkpoint ships no model_index.json; fall back
# to the model id resolving to a known diffusion family.
d = tmp_path / "flux-checkpoint"
_touch(d / "flux1-dev.safetensors")
assert (
models_route._local_model_task(
_local(d, model_id = "black-forest-labs/FLUX.1-dev")
)
== "text-to-image"
)
def test_local_task_none_for_plain_llm(tmp_path):
# A plain non-GGUF LLM checkpoint (no pipeline, no image family) stays untagged.
d = tmp_path / "llama"
_touch(d / "config.json")
_touch(d / "model.safetensors")
assert models_route._local_model_task(_local(d, model_id = "meta-llama/Llama-3.1-8B")) is None