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:
WildAi
2026-09-03 10:34:40 +03:00
parent 19d0e258e4
commit f2b2cfd276
2 changed files with 390 additions and 0 deletions
+137
View File
@@ -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
+253
View File
@@ -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)