Merge branch 'video-inference' into video-tab
This commit is contained in:
commit
4ecd7b6ef2
10 changed files with 448 additions and 33 deletions
|
|
@ -119,6 +119,22 @@ def _is_trusted_video_repo(repo_id: str) -> bool:
|
|||
return rid.startswith("unsloth/") or rid in _TRUSTED_NON_GGUF_VIDEO_REPOS
|
||||
|
||||
|
||||
def _detect_load_family(
|
||||
repo_id: str,
|
||||
gguf_filename: Optional[str],
|
||||
family_override: Optional[str],
|
||||
) -> Optional[VideoFamily]:
|
||||
"""Family detection shared by validate_load_request and the load worker: the
|
||||
repo id first, then the picked filename -- a local directory or generically
|
||||
named repo often carries the family token only in the checkpoint filename,
|
||||
and the worker must resolve the same family the validator accepted."""
|
||||
return detect_video_family(repo_id, family_override) or (
|
||||
detect_video_family(f"{repo_id}/{gguf_filename}")
|
||||
if gguf_filename and not family_override
|
||||
else None
|
||||
)
|
||||
|
||||
|
||||
def _ensure_mp4_encoder_available() -> None:
|
||||
"""Fail a load fast when PyAV is missing: the export otherwise dies AFTER a
|
||||
multi-minute denoise, which is the worst possible time to learn about it."""
|
||||
|
|
@ -193,11 +209,16 @@ class VideoBackend:
|
|||
) -> VideoFamily:
|
||||
"""Cheap, network-free validation shared by the route and the load path."""
|
||||
kind = resolve_video_model_kind(gguf_filename, model_kind)
|
||||
fam = detect_video_family(repo_id, family_override) or (
|
||||
detect_video_family(f"{repo_id}/{gguf_filename}")
|
||||
if gguf_filename and not family_override
|
||||
else None
|
||||
)
|
||||
# A -GGUF repo picked without a quant filename resolves to the pipeline
|
||||
# kind and would only fail minutes later in from_pretrained (no
|
||||
# model_index.json), AFTER the route evicted the current GPU owner.
|
||||
# Reject it here, where failing is still free.
|
||||
if kind == "pipeline" and repo_id.strip().lower().rstrip("/").endswith("-gguf"):
|
||||
raise ValueError(
|
||||
f"'{repo_id}' is a GGUF repo: pick one of its .gguf files "
|
||||
"(gguf_filename) instead of loading it as a diffusers pipeline."
|
||||
)
|
||||
fam = _detect_load_family(repo_id, gguf_filename, family_override)
|
||||
if fam is None:
|
||||
raise ValueError(
|
||||
f"'{repo_id}' is not a supported text-to-video model. Supported families: "
|
||||
|
|
@ -291,7 +312,9 @@ class VideoBackend:
|
|||
def _run_load(self, **kwargs: Any) -> None:
|
||||
token = kwargs.get("_load_token")
|
||||
try:
|
||||
fam = detect_video_family(kwargs["repo_id"], kwargs.get("family_override"))
|
||||
fam = _detect_load_family(
|
||||
kwargs["repo_id"], kwargs.get("gguf_filename"), kwargs.get("family_override")
|
||||
)
|
||||
kind = resolve_video_model_kind(kwargs.get("gguf_filename"), kwargs.get("model_kind"))
|
||||
base = (
|
||||
kwargs["repo_id"]
|
||||
|
|
@ -369,6 +392,11 @@ class VideoBackend:
|
|||
if self._load_token == token:
|
||||
self._loading = None
|
||||
except Exception as exc: # noqa: BLE001 -- surfaced via load_progress
|
||||
# A failed or cancelled load never commits _VideoLoadState, so the
|
||||
# teardown path has no snapshot to restore: roll back the process-wide
|
||||
# speed globals here (token-scoped, so a superseded load cannot clobber
|
||||
# the globals a newer in-flight load now owns).
|
||||
self._rollback_precommit_globals(token)
|
||||
if self._load_token != token:
|
||||
return
|
||||
logger.error("video.load_failed: %s", exc)
|
||||
|
|
@ -378,6 +406,25 @@ class VideoBackend:
|
|||
if self._load_token == token and self._loading is not None:
|
||||
self._loading.error = redact_native_paths(str(exc))
|
||||
|
||||
def _rollback_precommit_globals(self, token: Optional[int]) -> None:
|
||||
"""Restore process-wide speed globals (cudnn.benchmark / TF32 / the compiled
|
||||
GGUF dequantizer) for a load that died BEFORE committing _VideoLoadState.
|
||||
_teardown_state only restores from the committed state's snapshot, so an
|
||||
uncommitted load would otherwise leak its profile into the next speed=off
|
||||
load. Token-scoped: when a newer load has already taken the snapshot slot,
|
||||
the stale worker must leave the globals alone."""
|
||||
stored = getattr(self, "_precommit_globals", None)
|
||||
if stored is None:
|
||||
return
|
||||
stored_token, flags = stored
|
||||
if token is not None and stored_token is not None and stored_token != token:
|
||||
return
|
||||
self._precommit_globals = None
|
||||
restore_backend_flags(flags)
|
||||
from . import diffusion_gguf_compile
|
||||
|
||||
diffusion_gguf_compile.uninstall_all()
|
||||
|
||||
# Base-repo subfolders an LTX-2.3 assembly reads: the checkpoint (plus the GGUF
|
||||
# repo's extras files) supplies the DiT, connectors, both VAEs and the vocoder,
|
||||
# so only the 2.0 base's scheduler / text encoder / tokenizer are pulled.
|
||||
|
|
@ -584,6 +631,12 @@ class VideoBackend:
|
|||
# body, so a bare acquire is the exit barrier (never while holding _lock).
|
||||
with self._generate_lock:
|
||||
pass
|
||||
# The barrier wait can outlive this load: an unload or a newer load may
|
||||
# have superseded it while blocked, and tearing down now would destroy
|
||||
# the model that should remain current (or waste minutes building a
|
||||
# pipeline nobody wants). Recheck before touching shared state.
|
||||
if _load_token is not None and _load_token != self._load_token:
|
||||
raise RuntimeError("Video load was cancelled or superseded.")
|
||||
self._teardown_state()
|
||||
|
||||
target = resolve_diffusion_device_target()
|
||||
|
|
@ -594,13 +647,23 @@ class VideoBackend:
|
|||
dtype = target.dtype
|
||||
if fam.fp16_incompatible and dtype is torch.float16:
|
||||
dtype = torch.float32
|
||||
# The size tables below are bf16 (2-byte) figures. When the promotion
|
||||
# above lands fp32 weights on an accelerator (a pre-bf16 GPU), every
|
||||
# dense estimate doubles; budgeting the 2-byte figure would let auto
|
||||
# pick a resident plan that OOMs inside from_pretrained. GGUF weights
|
||||
# stay quantised on disk and in memory, so only dense estimates scale.
|
||||
dtype_scale = 2.0 if device != "cpu" and dtype is torch.float32 else 1.0
|
||||
|
||||
# ── memory plan: family-table resident estimate + frames-aware headroom.
|
||||
device_memory = snapshot_device_memory(target)
|
||||
components = fam.bf16_components_gb
|
||||
mib_per_gb = 1000.0**3 / (1024.0 * 1024.0)
|
||||
if kind == "pipeline":
|
||||
model_dense_mib = int(sum(components) * mib_per_gb) if components is not None else None
|
||||
model_dense_mib = (
|
||||
int(sum(components) * mib_per_gb * dtype_scale)
|
||||
if components is not None
|
||||
else None
|
||||
)
|
||||
companion_mib = None
|
||||
else:
|
||||
checkpoint_path = self._resolve_checkpoint_path(repo_id, gguf_filename, hf_token)
|
||||
|
|
@ -610,8 +673,10 @@ class VideoBackend:
|
|||
transformer_mib = estimate_gguf_resident_mib(size_mib)
|
||||
else:
|
||||
transformer_mib = estimate_safetensors_dense_mib(size_mib)
|
||||
if transformer_mib is not None:
|
||||
transformer_mib = int(transformer_mib * dtype_scale)
|
||||
companion_mib = (
|
||||
int((components[1] + components[2]) * mib_per_gb)
|
||||
int((components[1] + components[2]) * mib_per_gb * dtype_scale)
|
||||
if components is not None
|
||||
else None
|
||||
)
|
||||
|
|
@ -691,6 +756,11 @@ class VideoBackend:
|
|||
# placement/offload last.
|
||||
effective_speed = resolve_speed_mode(speed_mode, is_gguf = kind == "gguf")
|
||||
backend_flags = snapshot_backend_flags()
|
||||
# Until the state commit below transfers ownership to _teardown_state, a
|
||||
# failure or cancellation must restore these process-wide globals itself
|
||||
# (_run_load's error handler calls _rollback_precommit_globals with this
|
||||
# token). Registered BEFORE the first mutating call.
|
||||
self._precommit_globals = (_load_token, backend_flags)
|
||||
cache_engaged = apply_step_cache(
|
||||
pipe,
|
||||
mode = normalize_transformer_cache(transformer_cache),
|
||||
|
|
@ -782,6 +852,8 @@ class VideoBackend:
|
|||
transformer_cache = cache_engaged,
|
||||
resolved = resolved,
|
||||
)
|
||||
# Ownership of the globals transferred to _state / _teardown_state.
|
||||
self._precommit_globals = None
|
||||
logger.info(
|
||||
"video.loaded: %s (%s, %s, offload=%s, speed=%s)",
|
||||
repo_id,
|
||||
|
|
@ -1019,6 +1091,14 @@ class VideoBackend:
|
|||
self._loading = None
|
||||
if self._active_generate_cancel is not None:
|
||||
self._active_generate_cancel.set()
|
||||
# Wait for the signalled generation to actually exit before freeing the
|
||||
# pipeline: the denoise loop holds its own pipe reference until the next
|
||||
# step callback, so tearing down under it would report the VRAM free (and
|
||||
# let the GPU arbiter start another multi-GB load) while this clip still
|
||||
# occupies it. generate() holds _generate_lock for its full body, so a
|
||||
# bare acquire is the exit barrier (never taken while holding _lock).
|
||||
with self._generate_lock:
|
||||
pass
|
||||
self._teardown_state()
|
||||
logger.info("video.unloaded")
|
||||
return self.status()
|
||||
|
|
|
|||
|
|
@ -40,15 +40,19 @@ from core.training.diffusion_train_common import (
|
|||
DEFAULT_LORA_TARGETS,
|
||||
DiffusionLoraConfig,
|
||||
EventCb,
|
||||
LATENT_CACHE_OVER_BUDGET,
|
||||
StopCb,
|
||||
_apply_perf_flags,
|
||||
_assert_trusted_base_model,
|
||||
_emit,
|
||||
_latent_cache_forced,
|
||||
_latent_cache_over_budget,
|
||||
_plan_cache_variants,
|
||||
_publish_to_lora_catalog,
|
||||
_restore_perf_flags,
|
||||
discover_image_caption_pairs,
|
||||
has_functional_torchao,
|
||||
PermutationBatchSampler,
|
||||
repo_is_prequantized,
|
||||
resolve_train_steps,
|
||||
)
|
||||
|
|
@ -1012,6 +1016,9 @@ def _build_latent_cache(spec, vae, image_paths, cfg, device, weight_dtype, on_ev
|
|||
|
||||
cache: list[list[tuple]] = []
|
||||
total = len(image_paths)
|
||||
total_variants = sum(len(v) for v in plan)
|
||||
forced = _latent_cache_forced()
|
||||
gated = False
|
||||
for i, path in enumerate(image_paths):
|
||||
variants = []
|
||||
for u_left, u_top, flip in plan[i]:
|
||||
|
|
@ -1023,7 +1030,30 @@ def _build_latent_cache(spec, vae, image_paths, cfg, device, weight_dtype, on_ev
|
|||
.to(device)
|
||||
)
|
||||
a, b = spec.encode_latent_stats(vae, px)
|
||||
variants.append((_hold(a), _hold(b)))
|
||||
a, b = _hold(a), _hold(b)
|
||||
if not forced and not gated:
|
||||
# Size-gate the automatic cache off the first REAL encoded variant, before
|
||||
# building the rest: packed 16-channel DiT latents x variants x images of two
|
||||
# fp32 tensors can exhaust host/pinned RAM. Over budget we bail with the VAE
|
||||
# still resident so the loop encodes latents per step instead. ``b`` is None
|
||||
# for a deterministic-latent family, so only ``a`` contributes bytes there.
|
||||
per_variant = a.numel() * a.element_size()
|
||||
if b is not None:
|
||||
per_variant += b.numel() * b.element_size()
|
||||
if _latent_cache_over_budget(per_variant, total_variants):
|
||||
_emit(
|
||||
on_event,
|
||||
"warning",
|
||||
message = (
|
||||
"Latent cache disabled: estimated "
|
||||
f"{per_variant * total_variants / 1024 ** 3:.1f} GiB over the "
|
||||
"budget; encoding latents per step instead. Set "
|
||||
"UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE=1 to keep it."
|
||||
),
|
||||
)
|
||||
return LATENT_CACHE_OVER_BUDGET
|
||||
gated = True
|
||||
variants.append((a, b))
|
||||
cache.append(variants)
|
||||
if (i + 1) % 4 == 0 or i + 1 == total:
|
||||
_emit(on_event, "preparing", stage = "cache_latents", done = i + 1, total = total)
|
||||
|
|
@ -1258,7 +1288,11 @@ def _train_dit(cfg, spec, pairs, rng, device, weight_dtype, on_event, _check_sto
|
|||
latent_cache = _build_latent_cache(
|
||||
spec, vae, image_paths, cfg, device, weight_dtype, on_event, _check_stop
|
||||
)
|
||||
if latent_cache is None: # stopped during the cache build; nothing trained yet
|
||||
if latent_cache is LATENT_CACHE_OVER_BUDGET:
|
||||
# The estimated cache exceeded the host-memory budget; keep the VAE resident and
|
||||
# fall through to the in-loop encode path (latent_cache stays None).
|
||||
latent_cache = None
|
||||
elif latent_cache is None: # stopped during the cache build; nothing trained yet
|
||||
_emit(
|
||||
on_event,
|
||||
"complete",
|
||||
|
|
@ -1268,15 +1302,16 @@ def _train_dit(cfg, spec, pairs, rng, device, weight_dtype, on_event, _check_sto
|
|||
steps_run = 0,
|
||||
)
|
||||
return str(out_dir)
|
||||
try:
|
||||
pipe.vae = None
|
||||
except Exception: # noqa: BLE001 -- a pipeline without a settable vae keeps it
|
||||
pass
|
||||
del vae
|
||||
vae = None
|
||||
gc.collect()
|
||||
if device == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
else:
|
||||
try:
|
||||
pipe.vae = None
|
||||
except Exception: # noqa: BLE001 -- a pipeline without a settable vae keeps it
|
||||
pass
|
||||
del vae
|
||||
vae = None
|
||||
gc.collect()
|
||||
if device == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
# Variant picks use their own stream so the training loop's index/noise draws stay on
|
||||
# the same seed-deterministic sequence whether or not the cache is enabled.
|
||||
variant_rng = random.Random(cfg.seed + 1)
|
||||
|
|
@ -1349,6 +1384,10 @@ def _train_dit(cfg, spec, pairs, rng, device, weight_dtype, on_event, _check_sto
|
|||
transformer.train()
|
||||
n_images = len(image_paths)
|
||||
batch_size = cfg.train_batch_size
|
||||
# Permutation-cycle index sampler (shared with the SDXL trainer): visits every image once
|
||||
# per cycle before repeating, so a short run covers the whole dataset instead of the old
|
||||
# with-replacement draw. Uses the loop's own rng to stay seed-deterministic.
|
||||
index_sampler = PermutationBatchSampler(n_images, rng)
|
||||
stopped = False
|
||||
running_loss = 0.0
|
||||
peak_gb = 0.0
|
||||
|
|
@ -1368,7 +1407,7 @@ def _train_dit(cfg, spec, pairs, rng, device, weight_dtype, on_event, _check_sto
|
|||
optimizer.zero_grad(set_to_none = True)
|
||||
step_loss = 0.0
|
||||
for _ in range(cfg.gradient_accumulation_steps):
|
||||
idxs = [rng.randrange(n_images) for _ in range(batch_size)]
|
||||
idxs = index_sampler.next_batch(batch_size)
|
||||
if latent_cache is not None:
|
||||
latents = _sample_cached_latents(
|
||||
latent_cache, idxs, variant_rng, device, weight_dtype
|
||||
|
|
|
|||
|
|
@ -50,17 +50,21 @@ from core.training.diffusion_train_common import ( # noqa: F401
|
|||
EventCb,
|
||||
StopCb,
|
||||
DiffusionLoraConfig,
|
||||
LATENT_CACHE_OVER_BUDGET,
|
||||
_apply_perf_flags,
|
||||
_assert_trusted_base_model,
|
||||
_coerce_gradient_checkpointing,
|
||||
_config_from_dict,
|
||||
_CONFIG_ALIASES,
|
||||
_emit,
|
||||
_latent_cache_forced,
|
||||
_latent_cache_over_budget,
|
||||
_plan_cache_variants,
|
||||
_publish_to_lora_catalog,
|
||||
_restore_perf_flags,
|
||||
discover_image_caption_pairs,
|
||||
get_trainer,
|
||||
PermutationBatchSampler,
|
||||
resolve_train_steps,
|
||||
)
|
||||
|
||||
|
|
@ -206,6 +210,9 @@ def _build_sdxl_latent_cache(
|
|||
|
||||
cache: list[list[tuple]] = []
|
||||
total = len(image_paths)
|
||||
total_variants = sum(len(v) for v in plan)
|
||||
forced = _latent_cache_forced()
|
||||
gated = False
|
||||
for i, path in enumerate(image_paths):
|
||||
variants = []
|
||||
for u_left, u_top, flip in plan[i]:
|
||||
|
|
@ -217,6 +224,25 @@ def _build_sdxl_latent_cache(
|
|||
dist = vae.encode(pixel_values).latent_dist
|
||||
a = _hold(dist.mean * vae_scale)
|
||||
b = _hold(dist.std * vae_scale)
|
||||
if not forced and not gated:
|
||||
# Size-gate the automatic cache off the first REAL encoded variant, before
|
||||
# building the rest: thousands of images x variants of two fp32 tensors can
|
||||
# exhaust host/pinned RAM with no fallback. Over budget we bail with the VAE
|
||||
# still resident so the loop encodes latents per step instead.
|
||||
per_variant = a.numel() * a.element_size() + b.numel() * b.element_size()
|
||||
if _latent_cache_over_budget(per_variant, total_variants):
|
||||
_emit(
|
||||
on_event,
|
||||
"warning",
|
||||
message = (
|
||||
"Latent cache disabled: estimated "
|
||||
f"{per_variant * total_variants / 1024 ** 3:.1f} GiB over the "
|
||||
"budget; encoding latents per step instead. Set "
|
||||
"UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE=1 to keep it."
|
||||
),
|
||||
)
|
||||
return LATENT_CACHE_OVER_BUDGET
|
||||
gated = True
|
||||
variants.append((a, b, tuple(time_ids)))
|
||||
cache.append(variants)
|
||||
if (i + 1) % 4 == 0 or i + 1 == total:
|
||||
|
|
@ -422,7 +448,11 @@ def run_diffusion_lora_training(
|
|||
on_event,
|
||||
_check_stop,
|
||||
)
|
||||
if latent_cache is None: # stopped during the cache build; nothing trained yet
|
||||
if latent_cache is LATENT_CACHE_OVER_BUDGET:
|
||||
# The estimated cache exceeded the host-memory budget; keep the VAE resident
|
||||
# and fall through to the in-loop encode path (latent_cache stays None).
|
||||
latent_cache = None
|
||||
elif latent_cache is None: # stopped during the cache build; nothing trained yet
|
||||
out_dir = Path(cfg.output_dir).expanduser()
|
||||
_emit(
|
||||
on_event,
|
||||
|
|
@ -433,23 +463,30 @@ def run_diffusion_lora_training(
|
|||
steps_run = 0,
|
||||
)
|
||||
return str(out_dir)
|
||||
try:
|
||||
pipe.vae = None
|
||||
except Exception: # noqa: BLE001 -- a pipeline without a settable vae keeps it
|
||||
pass
|
||||
del vae
|
||||
vae = None
|
||||
gc.collect()
|
||||
if device == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
else:
|
||||
try:
|
||||
pipe.vae = None
|
||||
except Exception: # noqa: BLE001 -- a pipeline without a settable vae keeps it
|
||||
pass
|
||||
del vae
|
||||
vae = None
|
||||
gc.collect()
|
||||
if device == "cuda":
|
||||
torch.cuda.empty_cache()
|
||||
# Variant picks use their own stream so the loop's index/noise draws stay on the same
|
||||
# seed-deterministic sequence whether or not the cache is enabled.
|
||||
variant_rng = random.Random(cfg.seed + 1)
|
||||
|
||||
_emit(on_event, "model_load_completed", compiled = compiled)
|
||||
|
||||
# Permutation-cycle index sampler (shared with the DiT trainer): each dataset image is
|
||||
# visited once per cycle before any repeat, so a short run does not leave part of a
|
||||
# small dataset unseen. Draws from the loop's own rng so the sequence stays
|
||||
# seed-deterministic.
|
||||
index_sampler = PermutationBatchSampler(len(pairs), rng)
|
||||
|
||||
def _next_batch() -> tuple[list[int], list[str], list[str]]:
|
||||
idx = rng.sample(range(len(pairs)), k = min(cfg.train_batch_size, len(pairs)))
|
||||
idx = index_sampler.next_batch(min(cfg.train_batch_size, len(pairs)))
|
||||
chosen = [pairs[i] for i in idx]
|
||||
return idx, [c[0] for c in chosen], [c[1] for c in chosen]
|
||||
|
||||
|
|
|
|||
|
|
@ -444,6 +444,47 @@ def resolve_train_steps(cfg: "DiffusionLoraConfig", n_images: int) -> int:
|
|||
return cfg.train_steps
|
||||
|
||||
|
||||
class PermutationBatchSampler:
|
||||
"""Yields batch indices as consecutive slices of a reshuffled permutation of
|
||||
``range(n)``, so every index is visited exactly once per cycle before any repeats --
|
||||
an epoch-style full pass instead of the with-replacement draw that leaves part of a
|
||||
small dataset unseen at low step counts (num_epochs converts to a step budget, but the
|
||||
per-batch index draw is what decides coverage). When a cycle is exhausted the order is
|
||||
reshuffled from the run's own ``rng`` so the index stream stays seed-deterministic and
|
||||
each cycle differs.
|
||||
|
||||
Both trainers share this so the SDXL ``_next_batch`` path and the DiT per-sample draw
|
||||
select indices the same way. Only the index selection changes (with-replacement ->
|
||||
permutation cycles); step count and batch shapes are unchanged.
|
||||
"""
|
||||
|
||||
def __init__(self, n: int, rng: random.Random) -> None:
|
||||
if n <= 0:
|
||||
raise ValueError("PermutationBatchSampler needs at least one item")
|
||||
self._n = n
|
||||
self._rng = rng
|
||||
self._order: list[int] = []
|
||||
self._pos = 0
|
||||
|
||||
def _reshuffle(self) -> None:
|
||||
self._order = list(range(self._n))
|
||||
self._rng.shuffle(self._order)
|
||||
self._pos = 0
|
||||
|
||||
def next_batch(self, k: int) -> list[int]:
|
||||
# k may exceed n (batch larger than the dataset): the permutation is refilled across
|
||||
# as many cycles as needed so the caller always gets exactly k indices and the batch
|
||||
# never shrinks, matching the old sampler's fixed batch shape.
|
||||
out: list[int] = []
|
||||
while len(out) < k:
|
||||
if self._pos >= len(self._order):
|
||||
self._reshuffle()
|
||||
take = min(k - len(out), len(self._order) - self._pos)
|
||||
out.extend(self._order[self._pos : self._pos + take])
|
||||
self._pos += take
|
||||
return out
|
||||
|
||||
|
||||
def discover_image_caption_pairs(
|
||||
data_dir: str | os.PathLike[str],
|
||||
*,
|
||||
|
|
@ -549,6 +590,42 @@ def _plan_cache_variants(
|
|||
return plan
|
||||
|
||||
|
||||
# Host-memory budget for the AUTOMATIC latent cache. The cache holds two fp32 posterior
|
||||
# tensors (mean/std, VAE scale folded in) per crop/flip variant per image, pinned on a CUDA
|
||||
# host. At 1024px an SDXL variant is ~0.5 MiB and a 16-channel DiT variant several times
|
||||
# that, so a few thousand images x cache_variants can exhaust host or pinned RAM with no
|
||||
# fallback. Over this budget the default falls back to per-step VAE encoding. A fixed
|
||||
# constant (rather than a psutil RAM fraction) keeps the gate dependency-free and identical
|
||||
# across hosts; it is deliberately conservative, well under a typical training host's RAM.
|
||||
_LATENT_CACHE_BUDGET_BYTES = 4 * 1024 ** 3 # 4 GiB
|
||||
|
||||
# Returned by the cache builders when the estimated cache exceeds the budget: the caller
|
||||
# keeps the VAE resident and encodes each step's latents in-loop. A distinct sentinel from
|
||||
# ``None`` (which means a stop was requested mid-build) so the two are not conflated.
|
||||
LATENT_CACHE_OVER_BUDGET: Any = object()
|
||||
|
||||
|
||||
def _latent_cache_forced() -> bool:
|
||||
"""The user explicitly forced the latent cache on, bypassing the size gate. This is the
|
||||
explicit opt-in counterpart to ``UNSLOTH_DIFFUSION_NO_LATENT_CACHE`` (the explicit
|
||||
opt-out); only the automatic default is size-gated, so an explicit choice is honoured
|
||||
verbatim in either direction."""
|
||||
return os.environ.get("UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE", "") in ("1", "true")
|
||||
|
||||
|
||||
def _latent_cache_over_budget(
|
||||
per_variant_bytes: int, total_variants: int, budget_bytes: Optional[int] = None
|
||||
) -> bool:
|
||||
"""True when a cache of ``total_variants`` entries, each two fp32 tensors totalling
|
||||
``per_variant_bytes``, is estimated to exceed ``budget_bytes``. ``per_variant_bytes`` is
|
||||
measured from a real encoded latent, so the estimate tracks the actual per-family tensor
|
||||
shape (SDXL 4-channel vs. a packed 16-channel DiT latent) rather than a guess. The budget
|
||||
is read from the module constant at call time when not given, so tests can override it."""
|
||||
if budget_bytes is None:
|
||||
budget_bytes = _LATENT_CACHE_BUDGET_BYTES
|
||||
return per_variant_bytes * max(0, total_variants) > budget_bytes
|
||||
|
||||
|
||||
def _apply_perf_flags(
|
||||
cfg: "DiffusionLoraConfig",
|
||||
device: str,
|
||||
|
|
|
|||
|
|
@ -1337,7 +1337,11 @@ async def get_diffusion_training_run(
|
|||
from core.training.diffusion_training_service import get_diffusion_run
|
||||
|
||||
rec = get_diffusion_run(job_id)
|
||||
if rec is None:
|
||||
# A valid-JSON file that is not an object (a truncated / hand-edited [] record) would make
|
||||
# DiffusionTrainingRunDetail(**rec) raise TypeError -- not the ValidationError caught below
|
||||
# -- and 500 the endpoint. Treat any non-dict record as absent, matching the list route's
|
||||
# shape check.
|
||||
if not isinstance(rec, dict):
|
||||
raise HTTPException(status_code = 404, detail = "No such training run.")
|
||||
try:
|
||||
return DiffusionTrainingRunDetail(**rec)
|
||||
|
|
|
|||
|
|
@ -91,6 +91,7 @@ async def load_video_model(
|
|||
backend.validate_load_request,
|
||||
request.model_path,
|
||||
gguf_filename = request.gguf_filename,
|
||||
base_repo = request.base_repo,
|
||||
family_override = request.family_override,
|
||||
model_kind = request.model_kind,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -32,11 +32,16 @@ from core.training.diffusion_dit_trainer import (
|
|||
)
|
||||
from core.training.diffusion_train_common import (
|
||||
DiffusionLoraConfig,
|
||||
LATENT_CACHE_OVER_BUDGET,
|
||||
_apply_perf_flags,
|
||||
_config_from_dict,
|
||||
_latent_cache_forced,
|
||||
_latent_cache_over_budget,
|
||||
_plan_cache_variants,
|
||||
_restore_perf_flags,
|
||||
)
|
||||
import core.training.diffusion_lora_trainer as sdxl_trainer
|
||||
import core.training.diffusion_train_common as train_common
|
||||
from core.training.diffusion_training_service import DiffusionTrainingService
|
||||
from models.training import DiffusionTrainingStartRequest, DiffusionTrainingStopRequest
|
||||
from routes.training import router as training_router
|
||||
|
|
@ -375,3 +380,90 @@ def test_perf_flags_tf32_off_clears_flags():
|
|||
torch.get_float32_matmul_precision(),
|
||||
)
|
||||
assert after == before
|
||||
|
||||
|
||||
# ── latent cache size gate ────────────────────────────────────────────────────
|
||||
class _FakeLatentDist:
|
||||
def __init__(self, shape):
|
||||
self.mean = torch.zeros(shape, dtype = torch.float32)
|
||||
self.std = torch.ones(shape, dtype = torch.float32)
|
||||
|
||||
|
||||
class _FakeEncoded:
|
||||
def __init__(self, shape):
|
||||
self.latent_dist = _FakeLatentDist(shape)
|
||||
|
||||
|
||||
class _FakeVae:
|
||||
# Minimal VAE stand-in: encode() returns a posterior of the requested latent shape so the
|
||||
# builder measures a real per-variant byte size without a model load or image files.
|
||||
def __init__(self, shape):
|
||||
self._shape = shape
|
||||
|
||||
def encode(self, pixel_values):
|
||||
return _FakeEncoded(self._shape)
|
||||
|
||||
|
||||
def _fake_planned_loader(path, resolution, center_crop, u_left, u_top, flip):
|
||||
# The fake VAE ignores pixels; return a valid tensor + square SDXL time_ids.
|
||||
tensor = torch.zeros(3, resolution, resolution, dtype = torch.float32)
|
||||
return tensor, (resolution, resolution, 0, 0, resolution, resolution)
|
||||
|
||||
|
||||
def _build_fake_sdxl_cache(monkeypatch, num_images, latent_shape):
|
||||
# center_crop + no flip collapses to one variant per image, so total_variants == num_images.
|
||||
monkeypatch.setattr(sdxl_trainer, "_load_image_tensor_planned", _fake_planned_loader)
|
||||
cfg = _cfg(cache_variants = 1, center_crop = True, random_flip = False).normalized()
|
||||
return sdxl_trainer._build_sdxl_latent_cache(
|
||||
_FakeVae(latent_shape),
|
||||
1.0,
|
||||
[f"img{i}.png" for i in range(num_images)],
|
||||
cfg,
|
||||
"cpu",
|
||||
torch.float32,
|
||||
None,
|
||||
lambda: False,
|
||||
)
|
||||
|
||||
|
||||
def test_latent_cache_over_budget_boundary():
|
||||
# 32 bytes per variant x 4 variants = 128 bytes; exactly at budget is not "over".
|
||||
assert _latent_cache_over_budget(32, 4, budget_bytes = 200) is False
|
||||
assert _latent_cache_over_budget(32, 4, budget_bytes = 128) is False
|
||||
assert _latent_cache_over_budget(32, 4, budget_bytes = 127) is True
|
||||
# An empty plan can never overflow.
|
||||
assert _latent_cache_over_budget(1_000_000, 0, budget_bytes = 1) is False
|
||||
|
||||
|
||||
def test_latent_cache_forced_env(monkeypatch):
|
||||
monkeypatch.delenv("UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE", raising = False)
|
||||
assert _latent_cache_forced() is False
|
||||
monkeypatch.setenv("UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE", "1")
|
||||
assert _latent_cache_forced() is True
|
||||
|
||||
|
||||
def test_sdxl_cache_built_under_budget(monkeypatch):
|
||||
# Default (4 GiB) budget: a handful of tiny latents fits, so the full cache is returned.
|
||||
monkeypatch.delenv("UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE", raising = False)
|
||||
cache = _build_fake_sdxl_cache(monkeypatch, num_images = 3, latent_shape = (1, 4, 8, 8))
|
||||
assert cache is not LATENT_CACHE_OVER_BUDGET and cache is not None
|
||||
assert len(cache) == 3
|
||||
assert all(len(variants) == 1 for variants in cache)
|
||||
|
||||
|
||||
def test_sdxl_cache_gated_over_budget(monkeypatch):
|
||||
# A budget below one variant forces the gate on the first encode: the sentinel is returned
|
||||
# so the caller keeps the VAE resident and encodes per step.
|
||||
monkeypatch.delenv("UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE", raising = False)
|
||||
monkeypatch.setattr(train_common, "_LATENT_CACHE_BUDGET_BYTES", 8)
|
||||
cache = _build_fake_sdxl_cache(monkeypatch, num_images = 3, latent_shape = (1, 4, 8, 8))
|
||||
assert cache is LATENT_CACHE_OVER_BUDGET
|
||||
|
||||
|
||||
def test_sdxl_cache_force_bypasses_gate(monkeypatch):
|
||||
# An explicit force-on must be honoured verbatim even when the estimate is over budget.
|
||||
monkeypatch.setenv("UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE", "1")
|
||||
monkeypatch.setattr(train_common, "_LATENT_CACHE_BUDGET_BYTES", 8)
|
||||
cache = _build_fake_sdxl_cache(monkeypatch, num_images = 3, latent_shape = (1, 4, 8, 8))
|
||||
assert cache is not LATENT_CACHE_OVER_BUDGET and cache is not None
|
||||
assert len(cache) == 3
|
||||
|
|
|
|||
|
|
@ -441,6 +441,45 @@ def test_config_from_dict_epoch_mode_drops_max_steps_sentinel():
|
|||
assert cfg_explicit.train_steps == 25
|
||||
|
||||
|
||||
def test_permutation_sampler_covers_dataset_once_per_cycle():
|
||||
# Every index must appear exactly once per cycle before any repeat (epoch-style pass),
|
||||
# so a short run over a small dataset never leaves images unseen the way the old
|
||||
# with-replacement draw did. Consecutive cycles must be reshuffled (differ).
|
||||
import random
|
||||
|
||||
from core.training.diffusion_train_common import PermutationBatchSampler
|
||||
|
||||
n = 100
|
||||
sampler = PermutationBatchSampler(n, random.Random(0))
|
||||
|
||||
# Draw exactly one cycle in batches of 3 (n not divisible by the batch, so a batch spans
|
||||
# the cycle boundary); the first n indices must be a permutation of range(n).
|
||||
drawn: list[int] = []
|
||||
while len(drawn) < n:
|
||||
drawn.extend(sampler.next_batch(3))
|
||||
first_cycle = drawn[:n]
|
||||
assert sorted(first_cycle) == list(range(n)) # each index once, none missing
|
||||
|
||||
# The next full cycle is also a permutation, and it is reshuffled (order differs).
|
||||
fresh = PermutationBatchSampler(n, random.Random(0))
|
||||
cycle_a = fresh.next_batch(n)
|
||||
cycle_b = fresh.next_batch(n)
|
||||
assert sorted(cycle_a) == list(range(n))
|
||||
assert sorted(cycle_b) == list(range(n))
|
||||
assert cycle_a != cycle_b # cycles are reshuffled, not repeated in the same order
|
||||
|
||||
# A seed replays the exact index stream (determinism for reproducible runs).
|
||||
replay = PermutationBatchSampler(n, random.Random(0))
|
||||
assert replay.next_batch(n) == cycle_a
|
||||
|
||||
# A batch larger than the dataset refills across cycles so it never shrinks (batch shape
|
||||
# preserved), even though it must then repeat indices within the batch.
|
||||
big = PermutationBatchSampler(4, random.Random(1))
|
||||
batch = big.next_batch(10)
|
||||
assert len(batch) == 10
|
||||
assert set(batch) == {0, 1, 2, 3}
|
||||
|
||||
|
||||
def test_route_start_accepts_zero_max_grad_norm(client):
|
||||
# 0 is the documented "disable clipping" value (the trainer skips clip_grad_norm_);
|
||||
# the request model must not reject it.
|
||||
|
|
@ -951,3 +990,16 @@ def test_runs_route_tolerates_bad_field_record(client, _isolated_runs_dir):
|
|||
assert r.status_code == 200, r.text
|
||||
adapters = [x["adapter"] for x in r.json()["runs"]]
|
||||
assert adapters == ["good"] # the bad-field record was skipped, the good one remained
|
||||
|
||||
|
||||
def test_run_detail_route_non_object_record_is_404(client, _isolated_runs_dir):
|
||||
# A valid-JSON but non-object record (a truncated / hand-edited [] file named with a real
|
||||
# job id) makes DiffusionTrainingRunDetail(**rec) raise TypeError, not ValidationError; the
|
||||
# detail route must shape-check like the list path and 404 instead of 500.
|
||||
import json
|
||||
|
||||
job_id = "a" * 32
|
||||
(_isolated_runs_dir / f"{job_id}.json").write_text(json.dumps([]))
|
||||
|
||||
r = client.get(f"/api/train/diffusion/runs/{job_id}")
|
||||
assert r.status_code == 404, r.text
|
||||
|
|
|
|||
|
|
@ -12,7 +12,12 @@ import types
|
|||
|
||||
import pytest
|
||||
|
||||
from core.inference.video import VideoBackend, get_video_backend, resolve_video_model_kind
|
||||
from core.inference.video import (
|
||||
VideoBackend,
|
||||
_detect_load_family,
|
||||
get_video_backend,
|
||||
resolve_video_model_kind,
|
||||
)
|
||||
from core.inference.video_families import VIDEO_NOT_LOADED_MSG
|
||||
|
||||
|
||||
|
|
@ -211,6 +216,33 @@ def test_validate_gates_base_repo_and_local_paths(tmp_path):
|
|||
)
|
||||
|
||||
|
||||
def test_validate_rejects_gguf_repo_as_pipeline():
|
||||
backend = VideoBackend()
|
||||
# A -GGUF repo with no quant filename resolves to the pipeline kind and would
|
||||
# only fail minutes later in from_pretrained, AFTER evicting the GPU owner.
|
||||
with pytest.raises(ValueError, match = "pick one of its .gguf files"):
|
||||
backend.validate_load_request("unsloth/LTX-2.3-GGUF")
|
||||
with pytest.raises(ValueError, match = "pick one of its .gguf files"):
|
||||
backend.validate_load_request("unsloth/Wan2.2-TI2V-5B-GGUF/")
|
||||
|
||||
|
||||
def test_detect_load_family_filename_fallback():
|
||||
# Repo id alone carries the family.
|
||||
fam = _detect_load_family("Lightricks/LTX-2", None, None)
|
||||
assert fam is not None and fam.name == "ltx-2"
|
||||
# Repo id is opaque but the picked filename carries it: fall back to the
|
||||
# combined path so validate and _run_load agree on the family.
|
||||
fam = _detect_load_family("someorg/quants", "ltx-2-19b-Q4_K_M.gguf", None)
|
||||
assert fam is not None and fam.name == "ltx-2"
|
||||
# No filename and no recognisable repo id: no family.
|
||||
assert _detect_load_family("someorg/quants", None, None) is None
|
||||
# An explicit override resolves by name/alias and skips the filename fallback:
|
||||
# a bogus override stays None even when the filename would have matched.
|
||||
fam = _detect_load_family("someorg/quants", "ltx-2-19b-Q4_K_M.gguf", "ltxv")
|
||||
assert fam is not None and fam.name == "ltx-2"
|
||||
assert _detect_load_family("someorg/quants", "ltx-2-19b-Q4_K_M.gguf", "bogus") is None
|
||||
|
||||
|
||||
def test_load_generate_unload_gguf(fake_runtime, tmp_path):
|
||||
backend = VideoBackend()
|
||||
status = _load_gguf(backend, tmp_path)
|
||||
|
|
|
|||
|
|
@ -68,6 +68,7 @@ class _FakeBackend:
|
|||
model_path,
|
||||
*,
|
||||
gguf_filename = None,
|
||||
base_repo = None,
|
||||
family_override = None,
|
||||
model_kind = None,
|
||||
):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue