From f2b2cfd2765bdccb2fdb297f2bfd151341901f0a Mon Sep 17 00:00:00 2001 From: WildAi <2853742+wildminder@users.noreply.github.com> Date: Thu, 3 Sep 2026 10:34:40 +0300 Subject: [PATCH] feat(hiflow): add latent upsample + cascade driver MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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. --- src/hiflow.py | 137 +++++++++++++++++++++++ tests/test_hiflow.py | 253 +++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 390 insertions(+) diff --git a/src/hiflow.py b/src/hiflow.py index 4d685f6..a1268a4 100644 --- a/src/hiflow.py +++ b/src/hiflow.py @@ -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 diff --git a/tests/test_hiflow.py b/tests/test_hiflow.py index 615ec84..f2ae84a 100644 --- a/tests/test_hiflow.py +++ b/tests/test_hiflow.py @@ -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)