Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a7f7225b3a | ||
|
|
5dc5c2ce8b | ||
|
|
83645811f4 | ||
|
|
216bc05762 | ||
|
|
49e099ee7f | ||
|
|
32a16645b5 | ||
|
|
7bb0b0abea | ||
|
|
0b189dcf7a |
+161
-26
@@ -1,6 +1,8 @@
|
|||||||
|
from contextlib import contextmanager
|
||||||
from functools import lru_cache
|
from functools import lru_cache
|
||||||
import gc
|
import gc
|
||||||
import glob
|
import glob
|
||||||
|
import logging
|
||||||
import math
|
import math
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
@@ -40,9 +42,12 @@ LATENTS_STD = [
|
|||||||
]
|
]
|
||||||
|
|
||||||
_LATENT_UPSCALE_FOLDER = "latent_upscale_models"
|
_LATENT_UPSCALE_FOLDER = "latent_upscale_models"
|
||||||
|
LOGGER = logging.getLogger(__name__)
|
||||||
MP_UNIT = 1024 * 1024
|
MP_UNIT = 1024 * 1024
|
||||||
RES_MULTIPLE = 32
|
RES_MULTIPLE = 32
|
||||||
CUDA_MODEL_CHUNK_LENGTH = 17
|
CUDA_MODEL_CHUNK_LENGTH = 17
|
||||||
|
CUDA_MODEL_MIN_TOTAL_VRAM_BYTES = 10 * 1024 * 1024 * 1024
|
||||||
|
_MODEL_HOLD_DEPTH = 0
|
||||||
|
|
||||||
|
|
||||||
def _uses_cuda_model_upscale(param):
|
def _uses_cuda_model_upscale(param):
|
||||||
@@ -53,15 +58,83 @@ def _uses_cuda_model_upscale(param):
|
|||||||
return device == "cuda" and (mode == "model" or has_model_name)
|
return device == "cuda" and (mode == "model" or has_model_name)
|
||||||
|
|
||||||
|
|
||||||
def _effective_temporal_params(param):
|
def _effective_temporal_params(param, frame_count=None):
|
||||||
chunk_length = int(param.get("chunk_length", 0) or 0)
|
chunk_length = int(param.get("chunk_length", 0) or 0)
|
||||||
temporal_overlap = int(param.get("temporal_overlap", 0) or 0)
|
temporal_overlap = int(param.get("temporal_overlap", 0) or 0)
|
||||||
if _uses_cuda_model_upscale(param):
|
if _uses_cuda_model_upscale(param):
|
||||||
chunk_length = CUDA_MODEL_CHUNK_LENGTH if chunk_length <= 0 else min(chunk_length, CUDA_MODEL_CHUNK_LENGTH)
|
if chunk_length <= 0:
|
||||||
temporal_overlap = min(max(0, temporal_overlap), max(0, chunk_length - 17))
|
chunk_length = 85
|
||||||
|
temporal_overlap = 17
|
||||||
return chunk_length, temporal_overlap
|
return chunk_length, temporal_overlap
|
||||||
|
|
||||||
|
|
||||||
|
def _is_oom_error(exc):
|
||||||
|
text = str(exc).lower()
|
||||||
|
return "out of memory" in text or "exhausted its gpu spatial fallbacks" in text
|
||||||
|
|
||||||
|
|
||||||
|
def _should_retry_temporal_before_spatial(param):
|
||||||
|
return _uses_cuda_model_upscale(param) and int(param.get("chunk_length", 0) or 0) > CUDA_MODEL_CHUNK_LENGTH
|
||||||
|
|
||||||
|
|
||||||
|
def _cuda_total_memory_bytes():
|
||||||
|
try:
|
||||||
|
if not torch.cuda.is_available():
|
||||||
|
return 0
|
||||||
|
if hasattr(torch.cuda, "mem_get_info"):
|
||||||
|
_free, total = torch.cuda.mem_get_info()
|
||||||
|
return int(total)
|
||||||
|
current_device = torch.cuda.current_device() if hasattr(torch.cuda, "current_device") else 0
|
||||||
|
props = torch.cuda.get_device_properties(current_device)
|
||||||
|
return int(getattr(props, "total_memory", 0) or 0)
|
||||||
|
except Exception:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
def _should_skip_cuda_model_upscale(param):
|
||||||
|
total = _cuda_total_memory_bytes()
|
||||||
|
return _uses_cuda_model_upscale(param) and 0 < total < CUDA_MODEL_MIN_TOTAL_VRAM_BYTES
|
||||||
|
|
||||||
|
|
||||||
|
def _fallback_to_interp(video, param, reason):
|
||||||
|
method = param.get("method", "bilinear")
|
||||||
|
LOGGER.warning("H3 latent upscale model %s; using %s interpolation instead", reason, method)
|
||||||
|
try:
|
||||||
|
gc.collect()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
try:
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
interp_param = dict(param)
|
||||||
|
interp_param["mode"] = "interp"
|
||||||
|
return upscale_video_interp(video, interp_param)
|
||||||
|
|
||||||
|
|
||||||
|
def _retry_with_smaller_temporal(video, param, upscaler, exc):
|
||||||
|
if not _is_oom_error(exc):
|
||||||
|
raise exc
|
||||||
|
smaller = _shrink_temporal_param(param)
|
||||||
|
if smaller is None:
|
||||||
|
raise exc
|
||||||
|
LOGGER.info(
|
||||||
|
"H3 latent upscale temporal OOM: retrying with chunk_length=%s overlap=%s",
|
||||||
|
smaller.get("chunk_length"), smaller.get("temporal_overlap"),
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
gc.collect()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
try:
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return _upscale_video_temporal_chunks(video, smaller, upscaler)
|
||||||
|
|
||||||
|
|
||||||
def _models_dir():
|
def _models_dir():
|
||||||
try:
|
try:
|
||||||
if _LATENT_UPSCALE_FOLDER not in folder_paths.folder_names_and_paths:
|
if _LATENT_UPSCALE_FOLDER not in folder_paths.folder_names_and_paths:
|
||||||
@@ -349,7 +422,7 @@ def load_upscale_model(name, device, precision):
|
|||||||
return model
|
return model
|
||||||
|
|
||||||
|
|
||||||
def unload_upscale_model(name, device, precision):
|
def _unload_upscale_model_now(name, device, precision):
|
||||||
cache_key = f"{name}::{device}::{precision}"
|
cache_key = f"{name}::{device}::{precision}"
|
||||||
model = _MODEL_CACHE.get(cache_key)
|
model = _MODEL_CACHE.get(cache_key)
|
||||||
if model is not None and str(next(model.parameters()).device) != "cpu":
|
if model is not None and str(next(model.parameters()).device) != "cpu":
|
||||||
@@ -361,6 +434,22 @@ def unload_upscale_model(name, device, precision):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def unload_upscale_model(name, device, precision):
|
||||||
|
if _MODEL_HOLD_DEPTH > 0:
|
||||||
|
return
|
||||||
|
_unload_upscale_model_now(name, device, precision)
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def _hold_upscale_model_loaded():
|
||||||
|
global _MODEL_HOLD_DEPTH
|
||||||
|
_MODEL_HOLD_DEPTH += 1
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
_MODEL_HOLD_DEPTH -= 1
|
||||||
|
|
||||||
|
|
||||||
def _compute_upscale_target(width, height, h_in, w_in):
|
def _compute_upscale_target(width, height, h_in, w_in):
|
||||||
ds = 16
|
ds = 16
|
||||||
w_px = float(width)
|
w_px = float(width)
|
||||||
@@ -617,12 +706,25 @@ def upscale_video_model(video, param):
|
|||||||
except RuntimeError as exc:
|
except RuntimeError as exc:
|
||||||
if "out of memory" not in str(exc).lower():
|
if "out of memory" not in str(exc).lower():
|
||||||
raise
|
raise
|
||||||
|
if _should_retry_temporal_before_spatial(param):
|
||||||
|
raise RuntimeError(
|
||||||
|
"out of memory: retry H3 latent upscale with a smaller temporal chunk "
|
||||||
|
"before spatial fallback"
|
||||||
|
) from exc
|
||||||
smaller = _shrink_model_tile_param(param)
|
smaller = _shrink_model_tile_param(param)
|
||||||
if smaller is None:
|
if smaller is None:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"H3 latent upscale exhausted its GPU spatial fallbacks. "
|
"H3 latent upscale exhausted its GPU spatial fallbacks. "
|
||||||
"Reduce the target size, tile size, or split the shot earlier."
|
"Reduce the target size, tile size, or split the shot earlier."
|
||||||
) from exc
|
) from exc
|
||||||
|
LOGGER.info(
|
||||||
|
"H3 latent upscale model OOM: retrying with tile_size_mode=%s rows=%s cols=%s tile=%sx%s",
|
||||||
|
smaller.get("tile_size_mode"),
|
||||||
|
smaller.get("grid_rows"),
|
||||||
|
smaller.get("grid_cols"),
|
||||||
|
smaller.get("tile_width"),
|
||||||
|
smaller.get("tile_height"),
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
gc.collect()
|
gc.collect()
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -683,12 +785,22 @@ def _upscale_video_model_tiled(video, param):
|
|||||||
# If the requested tile is not smaller than the target on either axis,
|
# If the requested tile is not smaller than the target on either axis,
|
||||||
# the tiled path would just duplicate work.
|
# the tiled path would just duplicate work.
|
||||||
if len(rows) == 1 and len(cols) == 1:
|
if len(rows) == 1 and len(cols) == 1:
|
||||||
|
LOGGER.info(
|
||||||
|
"H3 latent upscale model: single core pass, tokens=%s target=%sx%s",
|
||||||
|
int(video.shape[2]), w_out, h_out,
|
||||||
|
)
|
||||||
return _upscale_video_model_core(video, param)
|
return _upscale_video_model_core(video, param)
|
||||||
|
|
||||||
scale_h = h_out / float(h_in)
|
scale_h = h_out / float(h_in)
|
||||||
scale_w = w_out / float(w_in)
|
scale_w = w_out / float(w_in)
|
||||||
orig_dtype = video.dtype
|
orig_dtype = video.dtype
|
||||||
out = torch.zeros((video.shape[0], video.shape[1], video.shape[2], h_out, w_out), device="cpu", dtype=orig_dtype)
|
out = torch.zeros((video.shape[0], video.shape[1], video.shape[2], h_out, w_out), device="cpu", dtype=orig_dtype)
|
||||||
|
LOGGER.info(
|
||||||
|
"H3 latent upscale model: %s spatial tiles, tokens=%s target=%sx%s tile_mode=%s tile=%sx%s overlap=%sx%s",
|
||||||
|
len(rows) * len(cols), int(video.shape[2]), w_out, h_out, mode, tile_w, tile_h,
|
||||||
|
spatial_w_overlap if mode == "rows_cols" else overlap,
|
||||||
|
spatial_h_overlap if mode == "rows_cols" else overlap,
|
||||||
|
)
|
||||||
|
|
||||||
for i, r0 in enumerate(rows):
|
for i, r0 in enumerate(rows):
|
||||||
tr = trows[i]
|
tr = trows[i]
|
||||||
@@ -753,43 +865,51 @@ def upscale_video_interp(video, param):
|
|||||||
def _upscale_video_temporal_chunks(video, param, upscaler):
|
def _upscale_video_temporal_chunks(video, param, upscaler):
|
||||||
if video.device.type != "cpu":
|
if video.device.type != "cpu":
|
||||||
video = video.to(device="cpu", copy=True)
|
video = video.to(device="cpu", copy=True)
|
||||||
chunk_length, temporal_overlap = _effective_temporal_params(param)
|
t = int(video.shape[2])
|
||||||
|
frame_count = _frames_for_tokens(t)
|
||||||
|
chunk_length, temporal_overlap = _effective_temporal_params(param, frame_count)
|
||||||
chunk_param = dict(param)
|
chunk_param = dict(param)
|
||||||
chunk_param["chunk_length"] = chunk_length
|
chunk_param["chunk_length"] = chunk_length
|
||||||
chunk_param["temporal_overlap"] = temporal_overlap
|
chunk_param["temporal_overlap"] = temporal_overlap
|
||||||
anchor_strength = float(param.get("anchor_strength", 0.999) or 0.999)
|
anchor_strength = float(param.get("anchor_strength", 0.999) or 0.999)
|
||||||
t = int(video.shape[2])
|
|
||||||
frame_count = _frames_for_tokens(t)
|
|
||||||
if chunk_length <= 0 or frame_count <= chunk_length:
|
if chunk_length <= 0 or frame_count <= chunk_length:
|
||||||
|
LOGGER.info(
|
||||||
|
"H3 latent upscale: single temporal chunk, tokens=%s frames=%s chunk_length=%s overlap=%s",
|
||||||
|
t, frame_count, chunk_length, temporal_overlap,
|
||||||
|
)
|
||||||
|
try:
|
||||||
return upscaler(video, chunk_param)
|
return upscaler(video, chunk_param)
|
||||||
|
except RuntimeError as exc:
|
||||||
|
return _retry_with_smaller_temporal(video, chunk_param, upscaler, exc)
|
||||||
|
|
||||||
bounds = _temporal_segments(t, chunk_length, temporal_overlap)
|
bounds = _temporal_segments(t, chunk_length, temporal_overlap)
|
||||||
if len(bounds) <= 1:
|
if len(bounds) <= 1:
|
||||||
|
LOGGER.info(
|
||||||
|
"H3 latent upscale: single temporal segment, tokens=%s frames=%s chunk_length=%s overlap=%s",
|
||||||
|
t, frame_count, chunk_length, temporal_overlap,
|
||||||
|
)
|
||||||
|
try:
|
||||||
return upscaler(video, chunk_param)
|
return upscaler(video, chunk_param)
|
||||||
|
except RuntimeError as exc:
|
||||||
|
return _retry_with_smaller_temporal(video, chunk_param, upscaler, exc)
|
||||||
|
|
||||||
orig_dtype = video.dtype
|
orig_dtype = video.dtype
|
||||||
out = None
|
out = None
|
||||||
out_h = out_w = None
|
out_h = out_w = None
|
||||||
|
LOGGER.info(
|
||||||
|
"H3 latent upscale: %s temporal chunks, tokens=%s frames=%s chunk_length=%s overlap=%s",
|
||||||
|
len(bounds), t, frame_count, chunk_length, temporal_overlap,
|
||||||
|
)
|
||||||
for i, (k0, f0, k1, f1) in enumerate(bounds):
|
for i, (k0, f0, k1, f1) in enumerate(bounds):
|
||||||
chunk = video[:, :, k0:k1].contiguous()
|
chunk = video[:, :, k0:k1].contiguous()
|
||||||
|
LOGGER.info(
|
||||||
|
"H3 latent upscale: temporal chunk %s/%s tokens %s:%s frames %s:%s",
|
||||||
|
i + 1, len(bounds), k0, k1, f0, f1,
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
chunk_out, chunk_h, chunk_w = upscaler(chunk, chunk_param)
|
chunk_out, chunk_h, chunk_w = upscaler(chunk, chunk_param)
|
||||||
except RuntimeError as exc:
|
except RuntimeError as exc:
|
||||||
if "out of memory" not in str(exc).lower():
|
return _retry_with_smaller_temporal(video, chunk_param, upscaler, exc)
|
||||||
raise
|
|
||||||
smaller = _shrink_temporal_param(chunk_param)
|
|
||||||
if smaller is None:
|
|
||||||
raise
|
|
||||||
try:
|
|
||||||
gc.collect()
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
if torch.cuda.is_available():
|
|
||||||
try:
|
|
||||||
torch.cuda.empty_cache()
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
return _upscale_video_temporal_chunks(video, smaller, upscaler)
|
|
||||||
chunk_out = chunk_out.to(device="cpu", dtype=orig_dtype)
|
chunk_out = chunk_out.to(device="cpu", dtype=orig_dtype)
|
||||||
if out is None:
|
if out is None:
|
||||||
out_h, out_w = chunk_h, chunk_w
|
out_h, out_w = chunk_h, chunk_w
|
||||||
@@ -819,7 +939,22 @@ def upscale_latent_video(video, param):
|
|||||||
if mode == "off":
|
if mode == "off":
|
||||||
return video, video.shape[-2], video.shape[-1]
|
return video, video.shape[-2], video.shape[-1]
|
||||||
if mode == "model":
|
if mode == "model":
|
||||||
|
if _should_skip_cuda_model_upscale(param):
|
||||||
|
return _fallback_to_interp(video, param, "requires more than this card's VRAM")
|
||||||
|
model_name = param.get("model_name")
|
||||||
|
device = param.get("device", "cuda")
|
||||||
|
precision = param.get("precision", "fp16")
|
||||||
|
dev = torch.device(device if (device == "cpu" or torch.cuda.is_available()) else "cpu")
|
||||||
|
try:
|
||||||
|
with _hold_upscale_model_loaded():
|
||||||
|
try:
|
||||||
return _upscale_video_temporal_chunks(video, param, upscale_video_model)
|
return _upscale_video_temporal_chunks(video, param, upscale_video_model)
|
||||||
|
finally:
|
||||||
|
_unload_upscale_model_now(model_name, dev, precision)
|
||||||
|
except RuntimeError as exc:
|
||||||
|
if not _is_oom_error(exc):
|
||||||
|
raise
|
||||||
|
return _fallback_to_interp(video, param, "exhausted GPU memory")
|
||||||
return _upscale_video_temporal_chunks(video, param, upscale_video_interp)
|
return _upscale_video_temporal_chunks(video, param, upscale_video_interp)
|
||||||
|
|
||||||
|
|
||||||
@@ -892,10 +1027,10 @@ class H3LatentUpscaleParams:
|
|||||||
"tooltip": "Temporal fade schedule over each tile's sampling. Off keeps the fade fixed; narrowing shrinks it over steps; widening grows it over steps."}),
|
"tooltip": "Temporal fade schedule over each tile's sampling. Off keeps the fade fixed; narrowing shrinks it over steps; widening grows it over steps."}),
|
||||||
"dynamic_fade_min": ("INT", {"default": 32, "min": 0, "max": 4096, "step": 32,
|
"dynamic_fade_min": ("INT", {"default": 32, "min": 0, "max": 4096, "step": 32,
|
||||||
"tooltip": "Minimum fade width used by dynamic_fade when it is enabled."}),
|
"tooltip": "Minimum fade width used by dynamic_fade when it is enabled."}),
|
||||||
"chunk_length": ("INT", {"default": 17, "min": 17, "max": 100000, "step": 17,
|
"chunk_length": ("INT", {"default": 85, "min": 17, "max": 100000, "step": 17,
|
||||||
"tooltip": "Temporal chunk length for latent upscale. CUDA model upscale is capped to 17 internally so short long-video shots do not bypass splitting and OOM."}),
|
"tooltip": "Temporal chunk length for latent upscale. CUDA model upscale keeps this when it splits the shot, but uses 17/0 if this would otherwise process the whole shot as one OOM-prone batch."}),
|
||||||
"temporal_overlap": ("INT", {"default": 0, "min": 0, "max": 100000, "step": 17,
|
"temporal_overlap": ("INT", {"default": 17, "min": 0, "max": 100000, "step": 17,
|
||||||
"tooltip": "Temporal overlap between latent chunks. CUDA model upscale uses 0 when capped to one H3 block to minimize peak VRAM."}),
|
"tooltip": "Temporal overlap between latent chunks. 17 matches the upstream split example; CUDA model upscale drops overlap only for the emergency 17-frame guard path."}),
|
||||||
"resize_conditioning": ("BOOLEAN", {"default": False,
|
"resize_conditioning": ("BOOLEAN", {"default": False,
|
||||||
"tooltip": "Reserved for upstream split compatibility. Leave OFF unless you need the original fallback behavior."}),
|
"tooltip": "Reserved for upstream split compatibility. Leave OFF unless you need the original fallback behavior."}),
|
||||||
"anchor_strength": ("FLOAT", {"default": 0.999, "min": 0.0, "max": 1.0, "step": 0.01,
|
"anchor_strength": ("FLOAT", {"default": 0.999, "min": 0.0, "max": 1.0, "step": 0.01,
|
||||||
|
|||||||
+98
-9
@@ -4040,6 +4040,88 @@ def _copy_sample_latent(out_latent):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _retarget_conditioning_spatial(cond, latent_h, latent_w):
|
||||||
|
"""Resize H3 keyframe latents in existing conditioning to a new latent grid."""
|
||||||
|
latent_h = int(latent_h)
|
||||||
|
latent_w = int(latent_w)
|
||||||
|
if latent_h <= 0 or latent_w <= 0:
|
||||||
|
raise RuntimeError("conditioning target latent size must be positive")
|
||||||
|
out = []
|
||||||
|
for item in cond:
|
||||||
|
try:
|
||||||
|
tensor, data = item
|
||||||
|
except Exception:
|
||||||
|
out.append(item)
|
||||||
|
continue
|
||||||
|
nd = dict(data)
|
||||||
|
keyframes = nd.get("minimax_keyframes")
|
||||||
|
if keyframes:
|
||||||
|
resized_keyframes = []
|
||||||
|
for keyframe in keyframes:
|
||||||
|
nkf = dict(keyframe)
|
||||||
|
latent_value = nkf.get("latent")
|
||||||
|
if latent_value is not None and len(getattr(latent_value, "shape", ())) >= 5:
|
||||||
|
if latent_value.shape[3] != latent_h or latent_value.shape[4] != latent_w:
|
||||||
|
b, c, t, h, w = latent_value.shape
|
||||||
|
resized = torch.nn.functional.interpolate(
|
||||||
|
latent_value.to(torch.float32).reshape(b * t, c, h, w),
|
||||||
|
size=(latent_h, latent_w),
|
||||||
|
mode="bilinear",
|
||||||
|
align_corners=False,
|
||||||
|
).reshape(b, c, t, latent_h, latent_w)
|
||||||
|
nkf["latent"] = resized.to(device=latent_value.device, dtype=latent_value.dtype)
|
||||||
|
resized_keyframes.append(nkf)
|
||||||
|
nd["minimax_keyframes"] = resized_keyframes
|
||||||
|
out.append([tensor, nd])
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def _pad_to_h3_patch_size(tensor):
|
||||||
|
try:
|
||||||
|
import comfy.ldm.common_dit as common_dit
|
||||||
|
return common_dit.pad_to_patch_size(tensor, (1, 2, 2))
|
||||||
|
except Exception:
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
|
||||||
|
def _crop_conditioning_to_tile(cond, source_h, source_w, row, col, tile_h, tile_w):
|
||||||
|
"""Crop H3 keyframe latents in existing conditioning for a spatial tile."""
|
||||||
|
out = []
|
||||||
|
for item in cond:
|
||||||
|
try:
|
||||||
|
tensor, data = item
|
||||||
|
except Exception:
|
||||||
|
out.append(item)
|
||||||
|
continue
|
||||||
|
nd = dict(data)
|
||||||
|
keyframes = nd.get("minimax_keyframes")
|
||||||
|
if keyframes:
|
||||||
|
cropped_keyframes = []
|
||||||
|
for keyframe in keyframes:
|
||||||
|
nkf = dict(keyframe)
|
||||||
|
latent_value = nkf.get("latent")
|
||||||
|
if latent_value is not None and len(getattr(latent_value, "shape", ())) >= 5:
|
||||||
|
kh, kw = latent_value.shape[3], latent_value.shape[4]
|
||||||
|
if kh != source_h or kw != source_w:
|
||||||
|
b, c, t, h, w = latent_value.shape
|
||||||
|
latent_value = torch.nn.functional.interpolate(
|
||||||
|
latent_value.to(torch.float32).reshape(b * t, c, h, w),
|
||||||
|
size=(source_h, source_w),
|
||||||
|
mode="bilinear",
|
||||||
|
align_corners=False,
|
||||||
|
).reshape(b, c, t, source_h, source_w).to(
|
||||||
|
device=latent_value.device,
|
||||||
|
dtype=latent_value.dtype,
|
||||||
|
)
|
||||||
|
nkf["latent"] = _pad_to_h3_patch_size(
|
||||||
|
latent_value[:, :, :, row:row + tile_h, col:col + tile_w].contiguous()
|
||||||
|
)
|
||||||
|
cropped_keyframes.append(nkf)
|
||||||
|
nd["minimax_keyframes"] = cropped_keyframes
|
||||||
|
out.append([tensor, nd])
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
def _latent_with_replaced_samples(template_latent, sampled_latent):
|
def _latent_with_replaced_samples(template_latent, sampled_latent):
|
||||||
"""Reuse the original latent payload, but swap in freshly sampled tensors."""
|
"""Reuse the original latent payload, but swap in freshly sampled tensors."""
|
||||||
if not isinstance(template_latent, dict):
|
if not isinstance(template_latent, dict):
|
||||||
@@ -6693,18 +6775,22 @@ class H3LongVideos:
|
|||||||
# pass; otherwise the 12-step base latent and the upscale latent sit
|
# pass; otherwise the 12-step base latent and the upscale latent sit
|
||||||
# in memory together and can trigger a retry loop.
|
# in memory together and can trigger a retry loop.
|
||||||
out["samples"] = comfy.nested_tensor.NestedTensor((upscaled_video, full_audio))
|
out["samples"] = comfy.nested_tensor.NestedTensor((upscaled_video, full_audio))
|
||||||
del out_samples, positive, latent, parts
|
del out_samples, parts
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
target_w = int(up_w) * 16
|
target_w = int(up_w) * 16
|
||||||
target_h = int(up_h) * 16
|
target_h = int(up_h) * 16
|
||||||
if target_w <= 0 or target_h <= 0:
|
if target_w <= 0 or target_h <= 0:
|
||||||
raise RuntimeError("latent upscale target size must be positive")
|
raise RuntimeError("latent upscale target size must be positive")
|
||||||
|
try:
|
||||||
|
upscale_cond = _retarget_conditioning_spatial(positive, int(up_h), int(up_w))
|
||||||
|
upscale_latent = dict(latent) if isinstance(latent, dict) else {}
|
||||||
|
except Exception:
|
||||||
upscale_cond, upscale_latent = _build_shot_conditioning(
|
upscale_cond, upscale_latent = _build_shot_conditioning(
|
||||||
clip, vae, prompt, target_w, target_h, ln, fps, handoff,
|
clip, vae, prompt, target_w, target_h, ln, fps, handoff,
|
||||||
ref_images=refs, ref_image_size=ref_image_size,
|
ref_images=refs, ref_image_size=ref_image_size,
|
||||||
ref_noise_aug=ref_noise_aug, audio_vae=audio_vae, silent=silent)
|
ref_noise_aug=ref_noise_aug, audio_vae=audio_vae, silent=silent)
|
||||||
upscale_latent["samples"] = comfy.nested_tensor.NestedTensor(
|
upscale_latent["samples"] = comfy.nested_tensor.NestedTensor((upscaled_video, full_audio))
|
||||||
(upscaled_video, full_audio))
|
del positive, latent
|
||||||
refine_steps = int(latent_upscale_param.get("steps", 2) or 2)
|
refine_steps = int(latent_upscale_param.get("steps", 2) or 2)
|
||||||
refine_sampler = latent_upscale_param.get("sampler_name", sn)
|
refine_sampler = latent_upscale_param.get("sampler_name", sn)
|
||||||
refine_scheduler = latent_upscale_param.get("scheduler", sch)
|
refine_scheduler = latent_upscale_param.get("scheduler", sch)
|
||||||
@@ -6771,6 +6857,11 @@ class H3LongVideos:
|
|||||||
rows, cols, trows, tcols, row_ovl, col_ovl = compute_spatial_grid(
|
rows, cols, trows, tcols, row_ovl, col_ovl = compute_spatial_grid(
|
||||||
int(up_h), int(up_w), tile_th, tile_tw, ol_th, ol_tw, min_tile_tw, min_tile_tw
|
int(up_h), int(up_w), tile_th, tile_tw, ol_th, ol_tw, min_tile_tw, min_tile_tw
|
||||||
)
|
)
|
||||||
|
logging.info(
|
||||||
|
"H3 latent refine: %s spatial sampler tiles, target=%sx%s tile_mode=%s tile=%sx%s overlap=%sx%s",
|
||||||
|
len(rows) * len(cols), target_w, target_h, tile_size_mode, tile_w_px, tile_h_px,
|
||||||
|
spatial_w_overlap_px, spatial_h_overlap_px,
|
||||||
|
)
|
||||||
if len(rows) == 1 and len(cols) == 1:
|
if len(rows) == 1 and len(cols) == 1:
|
||||||
(refined_out,) = nodes.common_ksampler(
|
(refined_out,) = nodes.common_ksampler(
|
||||||
model, seed, refine_steps, cfg, refine_sampler, refine_scheduler, upscale_cond, negative, upscale_latent,
|
model, seed, refine_steps, cfg, refine_sampler, refine_scheduler, upscale_cond, negative, upscale_latent,
|
||||||
@@ -6783,12 +6874,10 @@ class H3LongVideos:
|
|||||||
for col_index, c0 in enumerate(cols):
|
for col_index, c0 in enumerate(cols):
|
||||||
tc = tcols[col_index]
|
tc = tcols[col_index]
|
||||||
ovw = col_ovl[col_index]
|
ovw = col_ovl[col_index]
|
||||||
tile_target_w = int(tc) * 16
|
tile_cond = _crop_conditioning_to_tile(
|
||||||
tile_target_h = int(tr) * 16
|
upscale_cond, int(up_h), int(up_w), r0, c0, tr, tc
|
||||||
tile_cond, tile_latent = _build_shot_conditioning(
|
)
|
||||||
clip, vae, prompt, tile_target_w, tile_target_h, ln, fps, handoff,
|
tile_latent = dict(upscale_latent) if isinstance(upscale_latent, dict) else {}
|
||||||
ref_images=refs, ref_image_size=ref_image_size,
|
|
||||||
ref_noise_aug=ref_noise_aug, audio_vae=audio_vae, silent=silent)
|
|
||||||
tile_video = upscaled_video[:, :, :, r0:r0 + tr, c0:c0 + tc].contiguous()
|
tile_video = upscaled_video[:, :, :, r0:r0 + tr, c0:c0 + tc].contiguous()
|
||||||
tr_s = tr + (tr % 2)
|
tr_s = tr + (tr % 2)
|
||||||
tc_s = tc + (tc % 2)
|
tc_s = tc + (tc % 2)
|
||||||
|
|||||||
@@ -362,6 +362,7 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
|
|||||||
return self._parts
|
return self._parts
|
||||||
|
|
||||||
order = []
|
order = []
|
||||||
|
build_calls = []
|
||||||
first_out = {"samples": FakeNestedTensor((FakeTensor("v1"), FakeTensor("a1")))}
|
first_out = {"samples": FakeNestedTensor((FakeTensor("v1"), FakeTensor("a1")))}
|
||||||
second_out = {"samples": FakeNestedTensor((FakeTensor("v2"), FakeTensor("a2")))}
|
second_out = {"samples": FakeNestedTensor((FakeTensor("v2"), FakeTensor("a2")))}
|
||||||
|
|
||||||
@@ -408,11 +409,15 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
|
|||||||
order.append("upscale")
|
order.append("upscale")
|
||||||
return FakeTensor("upv"), 8, 16
|
return FakeTensor("upv"), 8, 16
|
||||||
|
|
||||||
self.module.nodes.common_ksampler = common_ksampler
|
def build_conditioning(*_args, **_kwargs):
|
||||||
self.module._build_shot_conditioning = lambda *_args, **_kwargs: (
|
build_calls.append(True)
|
||||||
"cond",
|
return (
|
||||||
|
[["cond", {}]],
|
||||||
{"samples": FakeNestedTensor((FakeTensor("basev"), FakeTensor("basea")))},
|
{"samples": FakeNestedTensor((FakeTensor("basev"), FakeTensor("basea")))},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self.module.nodes.common_ksampler = common_ksampler
|
||||||
|
self.module._build_shot_conditioning = build_conditioning
|
||||||
self.module._evict_all_but = lambda *_args, **_kwargs: None
|
self.module._evict_all_but = lambda *_args, **_kwargs: None
|
||||||
self.module.mm.unload_model_and_clones = unload_model_and_clones
|
self.module.mm.unload_model_and_clones = unload_model_and_clones
|
||||||
self.module.mm.unload_all_models = unload_all_models
|
self.module.mm.unload_all_models = unload_all_models
|
||||||
@@ -457,6 +462,7 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
|
|||||||
self.assertLess(order.index("upscale"), order.index("latent_upscale_sample"))
|
self.assertLess(order.index("upscale"), order.index("latent_upscale_sample"))
|
||||||
self.assertLess(order.index("audio"), order.index("video"))
|
self.assertLess(order.index("audio"), order.index("video"))
|
||||||
self.assertEqual(order[-1], "cleanup")
|
self.assertEqual(order[-1], "cleanup")
|
||||||
|
self.assertEqual(len(build_calls), 1)
|
||||||
finally:
|
finally:
|
||||||
self.module.nodes.common_ksampler = original_common_ksampler
|
self.module.nodes.common_ksampler = original_common_ksampler
|
||||||
self.module._build_shot_conditioning = original_build
|
self.module._build_shot_conditioning = original_build
|
||||||
@@ -476,6 +482,12 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
|
|||||||
else:
|
else:
|
||||||
self.module.comfy.nested_tensor.NestedTensor = original_nested
|
self.module.comfy.nested_tensor.NestedTensor = original_nested
|
||||||
|
|
||||||
|
def test_latent_refine_tiles_do_not_rebuild_conditioning(self):
|
||||||
|
source = inspect.getsource(self.module.H3LongVideos._render)
|
||||||
|
tile_branch = source[source.index("for col_index, c0 in enumerate(cols):"):]
|
||||||
|
self.assertIn("_crop_conditioning_to_tile", tile_branch)
|
||||||
|
self.assertNotIn("_build_shot_conditioning(", tile_branch)
|
||||||
|
|
||||||
def test_latent_upscale_off_skips_second_pass(self):
|
def test_latent_upscale_off_skips_second_pass(self):
|
||||||
calls = []
|
calls = []
|
||||||
original_common_ksampler = self.module.nodes.common_ksampler
|
original_common_ksampler = self.module.nodes.common_ksampler
|
||||||
@@ -1165,8 +1177,8 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
|
|||||||
self.assertFalse(required["brightness_match"][1]["default"])
|
self.assertFalse(required["brightness_match"][1]["default"])
|
||||||
self.assertEqual(required["dynamic_fade"][1]["default"], "off")
|
self.assertEqual(required["dynamic_fade"][1]["default"], "off")
|
||||||
self.assertEqual(required["dynamic_fade_min"][1]["default"], 32)
|
self.assertEqual(required["dynamic_fade_min"][1]["default"], 32)
|
||||||
self.assertEqual(required["chunk_length"][1]["default"], 17)
|
self.assertEqual(required["chunk_length"][1]["default"], 85)
|
||||||
self.assertEqual(required["temporal_overlap"][1]["default"], 0)
|
self.assertEqual(required["temporal_overlap"][1]["default"], 17)
|
||||||
self.assertFalse(required["resize_conditioning"][1]["default"])
|
self.assertFalse(required["resize_conditioning"][1]["default"])
|
||||||
self.assertEqual(required["anchor_strength"][1]["default"], 0.999)
|
self.assertEqual(required["anchor_strength"][1]["default"], 0.999)
|
||||||
|
|
||||||
@@ -1256,16 +1268,61 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
|
|||||||
self.assertEqual(smaller["chunk_length"], 17)
|
self.assertEqual(smaller["chunk_length"], 17)
|
||||||
self.assertEqual(smaller["temporal_overlap"], 0)
|
self.assertEqual(smaller["temporal_overlap"], 0)
|
||||||
|
|
||||||
def test_cuda_model_temporal_params_cap_saved_workflows(self):
|
def test_cuda_model_temporal_params_keep_splitting_saved_workflows(self):
|
||||||
latent = importlib.import_module("dumas_h3_latent_upscale")
|
latent = importlib.import_module("dumas_h3_latent_upscale")
|
||||||
chunk_length, temporal_overlap = latent._effective_temporal_params({
|
chunk_length, temporal_overlap = latent._effective_temporal_params({
|
||||||
"mode": "model",
|
"mode": "model",
|
||||||
"device": "cuda",
|
"device": "cuda",
|
||||||
"chunk_length": 85,
|
"chunk_length": 85,
|
||||||
"temporal_overlap": 17,
|
"temporal_overlap": 17,
|
||||||
})
|
}, frame_count=124)
|
||||||
self.assertEqual(chunk_length, 17)
|
self.assertEqual(chunk_length, 85)
|
||||||
self.assertEqual(temporal_overlap, 0)
|
self.assertEqual(temporal_overlap, 17)
|
||||||
|
|
||||||
|
def test_cuda_model_temporal_params_keep_short_saved_workflows_until_oom(self):
|
||||||
|
latent = importlib.import_module("dumas_h3_latent_upscale")
|
||||||
|
chunk_length, temporal_overlap = latent._effective_temporal_params({
|
||||||
|
"mode": "model",
|
||||||
|
"device": "cuda",
|
||||||
|
"chunk_length": 85,
|
||||||
|
"temporal_overlap": 17,
|
||||||
|
}, frame_count=85)
|
||||||
|
self.assertEqual(chunk_length, 85)
|
||||||
|
self.assertEqual(temporal_overlap, 17)
|
||||||
|
|
||||||
|
def test_cuda_model_oom_retries_temporal_before_spatial_fallback(self):
|
||||||
|
latent = importlib.import_module("dumas_h3_latent_upscale")
|
||||||
|
|
||||||
|
calls = []
|
||||||
|
original_tiled = latent._upscale_video_model_tiled
|
||||||
|
original_shrink_model = latent._shrink_model_tile_param
|
||||||
|
try:
|
||||||
|
def tiled(_video, param):
|
||||||
|
calls.append(("tiled", param.get("chunk_length"), param.get("tile_size_mode")))
|
||||||
|
raise RuntimeError("out of memory")
|
||||||
|
|
||||||
|
latent._upscale_video_model_tiled = tiled
|
||||||
|
latent._shrink_model_tile_param = (
|
||||||
|
lambda param: calls.append(("shrink_spatial", param.get("tile_size_mode"))) or None
|
||||||
|
)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "smaller temporal chunk"):
|
||||||
|
latent.upscale_video_model(
|
||||||
|
"video",
|
||||||
|
{
|
||||||
|
"mode": "model",
|
||||||
|
"device": "cuda",
|
||||||
|
"model_name": "upscale.safetensors",
|
||||||
|
"chunk_length": 85,
|
||||||
|
"temporal_overlap": 17,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(calls[0], ("tiled", 85, None))
|
||||||
|
self.assertNotIn(("shrink_spatial", None), calls)
|
||||||
|
finally:
|
||||||
|
latent._upscale_video_model_tiled = original_tiled
|
||||||
|
latent._shrink_model_tile_param = original_shrink_model
|
||||||
|
|
||||||
def test_interp_temporal_params_preserve_upstream_defaults(self):
|
def test_interp_temporal_params_preserve_upstream_defaults(self):
|
||||||
latent = importlib.import_module("dumas_h3_latent_upscale")
|
latent = importlib.import_module("dumas_h3_latent_upscale")
|
||||||
@@ -1278,6 +1335,182 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
|
|||||||
self.assertEqual(chunk_length, 85)
|
self.assertEqual(chunk_length, 85)
|
||||||
self.assertEqual(temporal_overlap, 17)
|
self.assertEqual(temporal_overlap, 17)
|
||||||
|
|
||||||
|
def test_unload_upscale_model_defers_while_held(self):
|
||||||
|
latent = importlib.import_module("dumas_h3_latent_upscale")
|
||||||
|
|
||||||
|
class FakeParam:
|
||||||
|
device = "cuda"
|
||||||
|
|
||||||
|
class FakeModel:
|
||||||
|
def __init__(self):
|
||||||
|
self.moves = []
|
||||||
|
|
||||||
|
def parameters(self):
|
||||||
|
return iter((FakeParam(),))
|
||||||
|
|
||||||
|
def to(self, device):
|
||||||
|
self.moves.append(device)
|
||||||
|
return self
|
||||||
|
|
||||||
|
cache_key = "upscale.safetensors::cuda::fp16"
|
||||||
|
original_cache_value = latent._MODEL_CACHE.get(cache_key)
|
||||||
|
original_hold_depth = latent._MODEL_HOLD_DEPTH
|
||||||
|
fake_model = FakeModel()
|
||||||
|
try:
|
||||||
|
latent._MODEL_CACHE[cache_key] = fake_model
|
||||||
|
latent._MODEL_HOLD_DEPTH = 0
|
||||||
|
with latent._hold_upscale_model_loaded():
|
||||||
|
latent.unload_upscale_model("upscale.safetensors", "cuda", "fp16")
|
||||||
|
self.assertEqual(fake_model.moves, [])
|
||||||
|
|
||||||
|
latent.unload_upscale_model("upscale.safetensors", "cuda", "fp16")
|
||||||
|
self.assertEqual(fake_model.moves, ["cpu"])
|
||||||
|
finally:
|
||||||
|
latent._MODEL_HOLD_DEPTH = original_hold_depth
|
||||||
|
if original_cache_value is None:
|
||||||
|
latent._MODEL_CACHE.pop(cache_key, None)
|
||||||
|
else:
|
||||||
|
latent._MODEL_CACHE[cache_key] = original_cache_value
|
||||||
|
|
||||||
|
def test_model_upscale_releases_cached_model_after_pass(self):
|
||||||
|
latent = importlib.import_module("dumas_h3_latent_upscale")
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
original_temporal = latent._upscale_video_temporal_chunks
|
||||||
|
original_unload_now = latent._unload_upscale_model_now
|
||||||
|
original_cuda = latent.torch.cuda
|
||||||
|
original_device = getattr(latent.torch, "device", None)
|
||||||
|
try:
|
||||||
|
latent.torch.cuda = types.SimpleNamespace(is_available=lambda: True)
|
||||||
|
latent.torch.device = lambda value: value
|
||||||
|
|
||||||
|
def temporal(video, param, upscaler):
|
||||||
|
calls.append(("temporal", latent._MODEL_HOLD_DEPTH))
|
||||||
|
return "video", 8, 16
|
||||||
|
|
||||||
|
def unload_now(name, device, precision):
|
||||||
|
calls.append(("unload", name, device, precision, latent._MODEL_HOLD_DEPTH))
|
||||||
|
|
||||||
|
latent._upscale_video_temporal_chunks = temporal
|
||||||
|
latent._unload_upscale_model_now = unload_now
|
||||||
|
|
||||||
|
result = latent.upscale_latent_video("source", {
|
||||||
|
"mode": "model",
|
||||||
|
"model_name": "upscale.safetensors",
|
||||||
|
"device": "cuda",
|
||||||
|
"precision": "fp16",
|
||||||
|
})
|
||||||
|
|
||||||
|
self.assertEqual(result, ("video", 8, 16))
|
||||||
|
self.assertEqual(calls[0], ("temporal", 1))
|
||||||
|
self.assertEqual(calls[1], ("unload", "upscale.safetensors", "cuda", "fp16", 1))
|
||||||
|
self.assertEqual(latent._MODEL_HOLD_DEPTH, 0)
|
||||||
|
finally:
|
||||||
|
latent._upscale_video_temporal_chunks = original_temporal
|
||||||
|
latent._unload_upscale_model_now = original_unload_now
|
||||||
|
latent.torch.cuda = original_cuda
|
||||||
|
if original_device is None:
|
||||||
|
delattr(latent.torch, "device")
|
||||||
|
else:
|
||||||
|
latent.torch.device = original_device
|
||||||
|
|
||||||
|
def test_model_upscale_oom_falls_back_to_interp(self):
|
||||||
|
latent = importlib.import_module("dumas_h3_latent_upscale")
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
original_temporal = latent._upscale_video_temporal_chunks
|
||||||
|
original_interp = latent.upscale_video_interp
|
||||||
|
original_unload_now = latent._unload_upscale_model_now
|
||||||
|
original_cuda = latent.torch.cuda
|
||||||
|
original_device = getattr(latent.torch, "device", None)
|
||||||
|
try:
|
||||||
|
latent.torch.cuda = types.SimpleNamespace(
|
||||||
|
is_available=lambda: True,
|
||||||
|
empty_cache=lambda: calls.append(("empty_cache",)),
|
||||||
|
)
|
||||||
|
latent.torch.device = lambda value: value
|
||||||
|
|
||||||
|
def temporal(_video, _param, _upscaler):
|
||||||
|
calls.append(("temporal", latent._MODEL_HOLD_DEPTH))
|
||||||
|
raise RuntimeError("H3 latent upscale exhausted its GPU spatial fallbacks")
|
||||||
|
|
||||||
|
def interp(video, param):
|
||||||
|
calls.append(("interp", video, param.get("mode"), param.get("method")))
|
||||||
|
return "interp_video", 8, 16
|
||||||
|
|
||||||
|
def unload_now(name, device, precision):
|
||||||
|
calls.append(("unload", name, device, precision, latent._MODEL_HOLD_DEPTH))
|
||||||
|
|
||||||
|
latent._upscale_video_temporal_chunks = temporal
|
||||||
|
latent.upscale_video_interp = interp
|
||||||
|
latent._unload_upscale_model_now = unload_now
|
||||||
|
|
||||||
|
result = latent.upscale_latent_video("source", {
|
||||||
|
"mode": "model",
|
||||||
|
"model_name": "upscale.safetensors",
|
||||||
|
"method": "bilinear",
|
||||||
|
"device": "cuda",
|
||||||
|
"precision": "fp16",
|
||||||
|
})
|
||||||
|
|
||||||
|
self.assertEqual(result, ("interp_video", 8, 16))
|
||||||
|
self.assertEqual(calls[0], ("temporal", 1))
|
||||||
|
self.assertEqual(calls[1], ("unload", "upscale.safetensors", "cuda", "fp16", 1))
|
||||||
|
self.assertIn(("empty_cache",), calls)
|
||||||
|
self.assertEqual(calls[-1], ("interp", "source", "interp", "bilinear"))
|
||||||
|
self.assertEqual(latent._MODEL_HOLD_DEPTH, 0)
|
||||||
|
finally:
|
||||||
|
latent._upscale_video_temporal_chunks = original_temporal
|
||||||
|
latent.upscale_video_interp = original_interp
|
||||||
|
latent._unload_upscale_model_now = original_unload_now
|
||||||
|
latent.torch.cuda = original_cuda
|
||||||
|
if original_device is None:
|
||||||
|
delattr(latent.torch, "device")
|
||||||
|
else:
|
||||||
|
latent.torch.device = original_device
|
||||||
|
|
||||||
|
def test_model_upscale_skips_learned_model_on_8gb_cuda(self):
|
||||||
|
latent = importlib.import_module("dumas_h3_latent_upscale")
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
original_temporal = latent._upscale_video_temporal_chunks
|
||||||
|
original_interp = latent.upscale_video_interp
|
||||||
|
original_cuda = latent.torch.cuda
|
||||||
|
try:
|
||||||
|
latent.torch.cuda = types.SimpleNamespace(
|
||||||
|
is_available=lambda: True,
|
||||||
|
mem_get_info=lambda: (1 * 1024 * 1024 * 1024, 8 * 1024 * 1024 * 1024),
|
||||||
|
empty_cache=lambda: calls.append(("empty_cache",)),
|
||||||
|
)
|
||||||
|
|
||||||
|
def temporal(_video, _param, _upscaler):
|
||||||
|
calls.append(("temporal",))
|
||||||
|
raise AssertionError("learned model path should be skipped on 8GB CUDA")
|
||||||
|
|
||||||
|
def interp(video, param):
|
||||||
|
calls.append(("interp", video, param.get("mode"), param.get("method")))
|
||||||
|
return "interp_video", 8, 16
|
||||||
|
|
||||||
|
latent._upscale_video_temporal_chunks = temporal
|
||||||
|
latent.upscale_video_interp = interp
|
||||||
|
|
||||||
|
result = latent.upscale_latent_video("source", {
|
||||||
|
"mode": "model",
|
||||||
|
"model_name": "upscale.safetensors",
|
||||||
|
"method": "bilinear",
|
||||||
|
"device": "cuda",
|
||||||
|
"precision": "fp16",
|
||||||
|
})
|
||||||
|
|
||||||
|
self.assertEqual(result, ("interp_video", 8, 16))
|
||||||
|
self.assertNotIn(("temporal",), calls)
|
||||||
|
self.assertIn(("empty_cache",), calls)
|
||||||
|
self.assertEqual(calls[-1], ("interp", "source", "interp", "bilinear"))
|
||||||
|
finally:
|
||||||
|
latent._upscale_video_temporal_chunks = original_temporal
|
||||||
|
latent.upscale_video_interp = original_interp
|
||||||
|
latent.torch.cuda = original_cuda
|
||||||
|
|
||||||
def test_upscale_video_model_raises_when_gpu_cannot_shrink(self):
|
def test_upscale_video_model_raises_when_gpu_cannot_shrink(self):
|
||||||
latent = importlib.import_module("dumas_h3_latent_upscale")
|
latent = importlib.import_module("dumas_h3_latent_upscale")
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user