feat(hiflow): add latent upsample + cascade driver
Step 5 of plan 2026-09-03. Stage sizes double per stage until the pixel target (16px-multiple snap for odd bases); steps_per_stage is an upper bound — the stage only walks schedule sigmas below tau.
This commit is contained in:
+137
@@ -25,6 +25,7 @@ from typing import Callable
|
||||
|
||||
import torch
|
||||
import torch.fft as fft
|
||||
import torch.nn.functional as F
|
||||
|
||||
Tensor = torch.Tensor
|
||||
|
||||
@@ -471,3 +472,139 @@ def guided_stage(
|
||||
|
||||
traj.put(0.0, x)
|
||||
return x.to(dtype), traj
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Latent upsampling + cascade driver
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def upsample_latent(x0: Tensor, target_h: int, target_w: int) -> Tensor:
|
||||
"""Antialiased bicubic upscale of a VAE-space x0 to (target_h, target_w).
|
||||
|
||||
fp32 round trip for half dtypes (antialias is not implemented for fp16 —
|
||||
the PixelRush precedent).
|
||||
"""
|
||||
orig_dtype = x0.dtype
|
||||
up = F.interpolate(
|
||||
x0.float(), size=(target_h, target_w),
|
||||
mode="bicubic", align_corners=False, antialias=True,
|
||||
)
|
||||
return up.to(orig_dtype)
|
||||
|
||||
|
||||
def _stage_latent_sizes(
|
||||
base_h: int, base_w: int, target_resolution: int, vae_downscale: int = 8,
|
||||
latent_multiple: int = 2,
|
||||
) -> list[tuple[int, int]]:
|
||||
"""Cascade stage sizes (latent H, W), doubling per stage until the pixel
|
||||
resolution reaches the target (plan D6). Sizes snap so pixel dimensions
|
||||
are multiples of 16 and latent dims of ``latent_multiple`` (FLUX packs
|
||||
2x2 latent patches -> even latent dims).
|
||||
"""
|
||||
def snap_latent(dim: int) -> int:
|
||||
px = dim * vae_downscale
|
||||
px = max(vae_downscale * latent_multiple, round(px / 16) * 16)
|
||||
snapped = px // vae_downscale
|
||||
# keep the latent dim a multiple of latent_multiple (round UP)
|
||||
return ((snapped + latent_multiple - 1) // latent_multiple) * latent_multiple
|
||||
|
||||
sizes = []
|
||||
h, w = snap_latent(base_h), snap_latent(base_w)
|
||||
while h * vae_downscale < target_resolution or w * vae_downscale < target_resolution:
|
||||
h, w = snap_latent(h * 2), snap_latent(w * 2)
|
||||
sizes.append((h, w))
|
||||
return sizes
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def hiflow_cascade(
|
||||
initial_latent: Tensor,
|
||||
base_sigmas: torch.Tensor,
|
||||
predict_x0_base: Callable[[Tensor, float], Tensor],
|
||||
predict_x0_stage: Callable[[Tensor, float], Tensor],
|
||||
target_resolution: int,
|
||||
cfg: HiFlowConfig,
|
||||
vae_decode: Callable[[Tensor], Tensor] | None = None,
|
||||
vae_encode: Callable[[Tensor], Tensor] | None = None,
|
||||
sharpen: Callable[[Tensor], Tensor] | None = None,
|
||||
vae_downscale: int = 8,
|
||||
progress_callback: Callable[[int, int, int], None] | None = None,
|
||||
) -> Tensor:
|
||||
"""Full HiFlow cascade: base trajectory, then guided upscale stages.
|
||||
|
||||
Stage sizes double per stage until the pixel resolution reaches the
|
||||
target (plan D6). The base trajectory is recorded once at the native
|
||||
size; each guided stage consumes the PREVIOUS stage's corrected-x0
|
||||
trajectory (R9), upsampled to the new latent size — latent bicubic by
|
||||
default, or the pixel decode->(sharpen)->encode round trip when
|
||||
``cfg.upsampling == "pixel"`` (needs vae_decode/vae_encode/sharpen).
|
||||
|
||||
Returns the final VAE-space latent at the last stage's size (or the
|
||||
base latent unchanged when the input is already at the target).
|
||||
"""
|
||||
if cfg.upsampling not in ("latent", "pixel"):
|
||||
raise ValueError(
|
||||
f"upsampling must be 'latent' or 'pixel'; got {cfg.upsampling!r}"
|
||||
)
|
||||
if cfg.upsampling == "pixel" and (
|
||||
vae_decode is None or vae_encode is None
|
||||
):
|
||||
raise ValueError(
|
||||
"upsampling='pixel' needs vae_decode and vae_encode adapters"
|
||||
)
|
||||
|
||||
# ---- Stage A: base trajectory at native size (guidance cfg). ----------
|
||||
final_base, ref_traj = base_trajectory(
|
||||
initial_latent, base_sigmas, predict_x0_base, cfg,
|
||||
progress_callback=progress_callback,
|
||||
)
|
||||
|
||||
base_h, base_w = initial_latent.shape[-2], initial_latent.shape[-1]
|
||||
sizes = _stage_latent_sizes(
|
||||
base_h, base_w, target_resolution, vae_downscale,
|
||||
)
|
||||
if not sizes:
|
||||
logger.info(
|
||||
"HiFlow: base %dx%d already at target %d — returning base output",
|
||||
base_h * vae_downscale, base_w * vae_downscale, target_resolution,
|
||||
)
|
||||
return final_base
|
||||
|
||||
def make_upsample(target_h: int, target_w: int) -> Callable[[Tensor], Tensor]:
|
||||
if cfg.upsampling == "latent":
|
||||
return lambda x: upsample_latent(x, target_h, target_w)
|
||||
|
||||
def pixel_up(x: Tensor) -> Tensor:
|
||||
image = vae_decode(x)
|
||||
image_up = F.interpolate(
|
||||
image.float(), size=(
|
||||
target_h * vae_downscale, target_w * vae_downscale),
|
||||
mode="bicubic", align_corners=False, antialias=True,
|
||||
).to(image.dtype)
|
||||
if sharpen is not None:
|
||||
image_up = sharpen(image_up)
|
||||
return vae_encode(image_up)
|
||||
|
||||
return pixel_up
|
||||
|
||||
x = final_base
|
||||
for stage_idx, (t_h, t_w) in enumerate(sizes):
|
||||
stage_sigmas = build_stage_sigmas(
|
||||
base_sigmas, cfg.tau, cfg.steps_per_stage,
|
||||
)
|
||||
logger.info(
|
||||
"HiFlow: stage %d/%d — latent %dx%d -> %dx%d, entry sigma %.4f",
|
||||
stage_idx + 1, len(sizes), x.shape[-2], x.shape[-1], t_h, t_w,
|
||||
float(stage_sigmas[0]),
|
||||
)
|
||||
seed = torch.zeros(
|
||||
initial_latent.shape[0], initial_latent.shape[1], t_h, t_w,
|
||||
device=initial_latent.device, dtype=initial_latent.dtype,
|
||||
)
|
||||
x, ref_traj = guided_stage(
|
||||
seed, ref_traj, stage_sigmas, predict_x0_stage,
|
||||
make_upsample(t_h, t_w), cfg,
|
||||
progress_callback=progress_callback,
|
||||
stage_index=stage_idx,
|
||||
)
|
||||
return x
|
||||
|
||||
@@ -824,3 +824,256 @@ class TestGuidedStage:
|
||||
latent, TrajectoryDict(), self.STAGE_SIGMAS,
|
||||
lambda x, s: 0.5 * x, _identity_upsample, _plain_cfg(),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Step 5 — upsample + cascade driver
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
from src.hiflow import ( # noqa: E402
|
||||
_stage_latent_sizes,
|
||||
hiflow_cascade,
|
||||
upsample_latent,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestUpsampleLatent:
|
||||
def test_doubles_shape(self):
|
||||
x = torch.randn(1, 4, 16, 16)
|
||||
up = upsample_latent(x, 32, 32)
|
||||
assert up.shape == (1, 4, 32, 32)
|
||||
|
||||
def test_antialias_smooths_high_freq(self):
|
||||
"""Upscaled (antialiased bicubic) must smooth a checkerboard more
|
||||
than the raw input scaled to the same footprint."""
|
||||
x = torch.zeros(1, 1, 16, 16)
|
||||
x[0, 0, ::2, ::2] = 1.0
|
||||
x[0, 0, 1::2, 1::2] = 1.0
|
||||
up = upsample_latent(x, 32, 32)
|
||||
|
||||
def rough(t):
|
||||
return (t[0, 0, 1:, :] - t[0, 0, :-1, :]).abs().mean().item()
|
||||
|
||||
assert rough(up) < rough(x), (
|
||||
"bicubic antialiasing must reduce checkerboard roughness"
|
||||
)
|
||||
|
||||
def test_fp16_roundtrip_dtype(self):
|
||||
x = torch.randn(1, 4, 8, 8, dtype=torch.float16)
|
||||
up = upsample_latent(x, 16, 16)
|
||||
assert up.dtype == torch.float16
|
||||
assert torch.isfinite(up).all()
|
||||
|
||||
def test_constant_image_preserved(self):
|
||||
x = torch.full((1, 2, 8, 8), 0.7)
|
||||
up = upsample_latent(x, 16, 16)
|
||||
assert torch.allclose(up, torch.full_like(up, 0.7), atol=1e-3)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestStageLatentSizes:
|
||||
def test_flux_1024_targets(self):
|
||||
assert _stage_latent_sizes(128, 128, 1024) == []
|
||||
assert _stage_latent_sizes(128, 128, 2048) == [(256, 256)]
|
||||
assert _stage_latent_sizes(128, 128, 4096) == [(256, 256), (512, 512)]
|
||||
|
||||
def test_odd_base_snapped(self):
|
||||
"""1088px base (latent 136) doubles to a 16-px-multiple target."""
|
||||
sizes = _stage_latent_sizes(136, 136, 2048)
|
||||
assert sizes == [(272, 272)]
|
||||
for h, w in sizes:
|
||||
assert (h * 8) % 16 == 0 and (w * 8) % 16 == 0
|
||||
assert h % 2 == 0 and w % 2 == 0 # FLUX 2x2 packing
|
||||
|
||||
def test_terminates_at_target(self):
|
||||
sizes = _stage_latent_sizes(64, 64, 4096)
|
||||
h, w = sizes[-1]
|
||||
assert h * 8 >= 4096 or w * 8 >= 4096
|
||||
assert len(sizes) <= 6 # 512px base -> bounded stage count
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestHiflowCascade:
|
||||
SIGMAS = torch.tensor([1.0, 0.75, 0.5, 0.25, 0.0])
|
||||
|
||||
def _cfg(self, **kw):
|
||||
base = dict(
|
||||
tau=0.5, steps_per_stage=3, filter_ratio=0.2, upsampling="latent",
|
||||
)
|
||||
base.update(kw)
|
||||
return HiFlowConfig(**base)
|
||||
|
||||
def test_two_stages_resolution_progression(self):
|
||||
"""Base 32x32 latent (256px); target 1024px -> stages 64, 128."""
|
||||
torch.manual_seed(0)
|
||||
z = torch.randn(1, 4, 32, 32)
|
||||
out = hiflow_cascade(
|
||||
z, self.SIGMAS,
|
||||
lambda x, s: 0.5 * x, # base
|
||||
lambda x, s: 0.5 * x, # stage
|
||||
target_resolution=1024,
|
||||
cfg=self._cfg(),
|
||||
vae_downscale=8,
|
||||
)
|
||||
assert out.shape == (1, 4, 128, 128)
|
||||
|
||||
def test_single_stage_doubles(self):
|
||||
torch.manual_seed(1)
|
||||
z = torch.randn(1, 4, 32, 32)
|
||||
out = hiflow_cascade(
|
||||
z, self.SIGMAS,
|
||||
lambda x, s: 0.5 * x, lambda x, s: 0.5 * x,
|
||||
target_resolution=512, # 256px base -> one 512px stage
|
||||
cfg=self._cfg(),
|
||||
)
|
||||
assert out.shape == (1, 4, 64, 64)
|
||||
|
||||
def test_second_stage_references_first_stage_trajectory(self):
|
||||
"""Stage-2's model inputs must be at stage-2 size (the reference
|
||||
chain feeds through the previous stage's upsampled trajectory)."""
|
||||
torch.manual_seed(2)
|
||||
z = torch.randn(1, 4, 16, 16)
|
||||
seen_sizes = []
|
||||
|
||||
def stage_predict(x, s):
|
||||
seen_sizes.append(tuple(x.shape[-2:]))
|
||||
return 0.5 * x
|
||||
|
||||
hiflow_cascade(
|
||||
z, self.SIGMAS,
|
||||
lambda x, s: 0.5 * x, stage_predict,
|
||||
target_resolution=512, # 128px base -> stages 256, 512
|
||||
cfg=self._cfg(),
|
||||
)
|
||||
# Unique sizes, order-preserving: stage 1 at 32x32, stage 2 at 64x64.
|
||||
uniq = list(dict.fromkeys(seen_sizes))
|
||||
assert uniq == [(32, 32), (64, 64)]
|
||||
|
||||
def test_pixel_mode_calls_vae_adapters(self):
|
||||
torch.manual_seed(3)
|
||||
z = torch.randn(1, 4, 16, 16)
|
||||
calls = {"decode": 0, "encode": 0}
|
||||
|
||||
def vae_decode(latent):
|
||||
"""Real-VAE contract: latent [B,C,h,w] -> image [B,3,8h,8w]."""
|
||||
calls["decode"] += 1
|
||||
b, c, h, w = latent.shape
|
||||
img = latent[:, :3].repeat_interleave(8, -2).repeat_interleave(8, -1)
|
||||
assert img.shape == (b, 3, h * 8, w * 8)
|
||||
return img
|
||||
|
||||
def vae_encode(image):
|
||||
"""Real-VAE contract: image [B,3,H,W] -> latent [B,4,H/8,W/8]."""
|
||||
calls["encode"] += 1
|
||||
small = image[:, :1, ::8, ::8]
|
||||
return torch.cat([small] * 4, dim=1)
|
||||
|
||||
hiflow_cascade(
|
||||
z, self.SIGMAS,
|
||||
lambda x, s: 0.5 * x, lambda x, s: 0.5 * x,
|
||||
target_resolution=256, # one stage
|
||||
cfg=self._cfg(upsampling="pixel"),
|
||||
vae_decode=vae_decode, vae_encode=vae_encode,
|
||||
sharpen=lambda im: im,
|
||||
)
|
||||
assert calls["decode"] >= 1 and calls["encode"] >= 1
|
||||
|
||||
def test_latent_mode_never_calls_vae_adapters(self):
|
||||
torch.manual_seed(4)
|
||||
z = torch.randn(1, 4, 16, 16)
|
||||
|
||||
def fail_decode(x):
|
||||
raise AssertionError("latent mode must not call vae_decode")
|
||||
|
||||
hiflow_cascade(
|
||||
z, self.SIGMAS,
|
||||
lambda x, s: 0.5 * x, lambda x, s: 0.5 * x,
|
||||
target_resolution=256,
|
||||
cfg=self._cfg(upsampling="latent"),
|
||||
vae_decode=fail_decode, vae_encode=fail_decode,
|
||||
)
|
||||
|
||||
def test_pixel_mode_without_vae_rejected(self):
|
||||
torch.manual_seed(5)
|
||||
z = torch.randn(1, 4, 8, 8)
|
||||
with pytest.raises(ValueError, match="vae_decode and vae_encode"):
|
||||
hiflow_cascade(
|
||||
z, self.SIGMAS,
|
||||
lambda x, s: 0.5 * x, lambda x, s: 0.5 * x,
|
||||
target_resolution=128,
|
||||
cfg=self._cfg(upsampling="pixel"),
|
||||
)
|
||||
|
||||
def test_invalid_upsampling_rejected(self):
|
||||
torch.manual_seed(6)
|
||||
z = torch.randn(1, 4, 8, 8)
|
||||
with pytest.raises(ValueError, match="latent.*pixel"):
|
||||
hiflow_cascade(
|
||||
z, self.SIGMAS,
|
||||
lambda x, s: 0.5 * x, lambda x, s: 0.5 * x,
|
||||
target_resolution=128,
|
||||
cfg=self._cfg(upsampling="bogus"),
|
||||
)
|
||||
|
||||
def test_base_at_target_returns_base_output(self):
|
||||
"""No stages: the base trajectory's final latent is returned."""
|
||||
from src.hiflow import base_trajectory
|
||||
torch.manual_seed(7)
|
||||
z = torch.randn(1, 4, 32, 32)
|
||||
cfg = self._cfg()
|
||||
out = hiflow_cascade(
|
||||
z, self.SIGMAS,
|
||||
lambda x, s: 0.5 * x, lambda x, s: 0.5 * x,
|
||||
target_resolution=256, # 32*8 == 256 -> no stages
|
||||
cfg=cfg,
|
||||
)
|
||||
expected, _ = base_trajectory(
|
||||
z, self.SIGMAS, lambda x, s: 0.5 * x, cfg)
|
||||
assert torch.allclose(out, expected, atol=1e-6)
|
||||
|
||||
def test_cascade_no_nan(self):
|
||||
torch.manual_seed(8)
|
||||
z = torch.randn(2, 4, 32, 32)
|
||||
out = hiflow_cascade(
|
||||
z, self.SIGMAS,
|
||||
lambda x, s: 0.5 * x + 0.05 * torch.randn_like(x),
|
||||
lambda x, s: 0.5 * x + 0.05 * torch.randn_like(x),
|
||||
target_resolution=1024,
|
||||
cfg=self._cfg(alpha_scale=1.0, beta_scale=0.5),
|
||||
)
|
||||
assert torch.isfinite(out).all()
|
||||
|
||||
def test_progress_total_events(self):
|
||||
"""Base steps + per-stage transitions fire progress events.
|
||||
|
||||
steps_per_stage is an UPPER bound: the stage walks only schedule
|
||||
sigmas below tau (tau=0.5 leaves 0.25 -> 0: 2 transitions here).
|
||||
"""
|
||||
events = []
|
||||
torch.manual_seed(9)
|
||||
z = torch.randn(1, 4, 32, 32)
|
||||
hiflow_cascade(
|
||||
z, self.SIGMAS,
|
||||
lambda x, s: 0.5 * x, lambda x, s: 0.5 * x,
|
||||
target_resolution=512, # one stage
|
||||
cfg=self._cfg(),
|
||||
progress_callback=lambda i, total, stage: events.append(stage),
|
||||
)
|
||||
assert events.count(-1) == len(self.SIGMAS) - 1
|
||||
assert events.count(0) == 2 # 0.25 -> 0 only (tau=0.5 entry)
|
||||
|
||||
def test_tau_above_schedule_clamps_logged(self, caplog):
|
||||
"""tau above sigma_max clamps the entry to sigma_max with a warning
|
||||
(schedule [1.0, 0.5, 0.0]: tau=0.99 enters at 1.0)."""
|
||||
import logging
|
||||
torch.manual_seed(10)
|
||||
z = torch.randn(1, 4, 32, 32)
|
||||
with caplog.at_level(logging.WARNING, logger="ComfyUI-DyPE"):
|
||||
hiflow_cascade(
|
||||
z, torch.tensor([0.9, 0.5, 0.0]),
|
||||
lambda x, s: 0.5 * x, lambda x, s: 0.5 * x,
|
||||
target_resolution=512,
|
||||
cfg=self._cfg(tau=0.99, steps_per_stage=2),
|
||||
)
|
||||
assert any("above schedule" in r.message for r in caplog.records)
|
||||
|
||||
Reference in New Issue
Block a user