feat(hiflow): add reference trajectory + base stage
Step 3 of plan 2026-09-03. Entries park on CPU (a 30-step 4K trajectory is ~1.5 GB fp32); time matching uses nearest-sigma within tolerance, robust across differently-spaced schedules.
This commit is contained in:
+143
@@ -20,6 +20,8 @@ never sees model-space tensors.
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable
|
||||
|
||||
import torch
|
||||
import torch.fft as fft
|
||||
@@ -175,3 +177,144 @@ def alignment_scales(
|
||||
scales = sigmas / entry
|
||||
scales = torch.clamp(scales, min=0.0)
|
||||
return scales, scales.clone()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Reference trajectory storage + base (Stage A) sampling
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _sigma_key(sigma: float) -> float:
|
||||
"""Round a sigma to a stable dict key (float equality is fragile)."""
|
||||
return round(float(sigma), 6)
|
||||
|
||||
|
||||
class TrajectoryDict:
|
||||
"""sigma -> predicted-clean-x0 dictionary, the HiFlow reference flow.
|
||||
|
||||
Stores every per-step clean prediction from a sampling run keyed by its
|
||||
sigma (6-decimal rounded). Tensors are parked on CPU to keep a 30-step
|
||||
4K trajectory (~1.5 GB fp32) off the GPU (plan D8); ``get``/``nearest``
|
||||
move them back to the caller's device.
|
||||
|
||||
Time matching (plan D8): exact sigma first, then the nearest stored sigma
|
||||
within ``tol`` — robust across differently-spaced schedules, unlike the
|
||||
reference's exact-timestep lookup.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._entries: dict[float, torch.Tensor] = {}
|
||||
|
||||
def put(self, sigma: float, x0: torch.Tensor) -> None:
|
||||
self._entries[_sigma_key(sigma)] = x0.detach().to("cpu")
|
||||
|
||||
def get(self, sigma: float) -> torch.Tensor:
|
||||
key = _sigma_key(sigma)
|
||||
if key not in self._entries:
|
||||
raise KeyError(
|
||||
f"no trajectory entry at sigma={sigma!r}; stored sigmas: "
|
||||
f"{sorted(self._entries)[:8]}{'...' if len(self._entries) > 8 else ''}"
|
||||
)
|
||||
return self._entries[key]
|
||||
|
||||
def nearest(
|
||||
self,
|
||||
sigma: float,
|
||||
tol: float = 1e-3,
|
||||
device: torch.device | None = None,
|
||||
) -> tuple[float, torch.Tensor]:
|
||||
"""Return (stored_sigma, x0) closest to ``sigma`` within ``tol``."""
|
||||
if not self._entries:
|
||||
raise ValueError(
|
||||
"no reference trajectory — run the base stage first"
|
||||
)
|
||||
keys = sorted(self._entries)
|
||||
target = float(sigma)
|
||||
best = min(keys, key=lambda k: abs(k - target))
|
||||
if abs(best - target) > tol:
|
||||
raise ValueError(
|
||||
f"no reference entry within {tol} of sigma={target!r} "
|
||||
f"(nearest stored: {best!r})"
|
||||
)
|
||||
x0 = self._entries[best]
|
||||
if device is not None:
|
||||
x0 = x0.to(device)
|
||||
return best, x0
|
||||
|
||||
def sigmas(self) -> list[float]:
|
||||
return sorted(self._entries)
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self._entries)
|
||||
|
||||
def as_dict(self) -> dict[float, torch.Tensor]:
|
||||
return dict(self._entries)
|
||||
|
||||
|
||||
@dataclass
|
||||
class HiFlowConfig:
|
||||
"""HiFlow hyperparameters (paper values; plan D7).
|
||||
|
||||
tau stage-entry noise level; paper cascade [0.6, 0.3, 0.3]
|
||||
steps base-stage sampling steps (paper: 30)
|
||||
steps_per_stage guided-stage transitions per cascade stage (repo: 16/10)
|
||||
cfg base-stage classifier-free guidance (FLUX-dev: 3.5)
|
||||
guidance_high guided-stage CFG (repo: 4.5-6)
|
||||
filter_ratio normalized Butterworth cutoff D (paper 0.4 / repo 0.2)
|
||||
alpha_scale direction-alignment multiplier (repo first stage 1.0)
|
||||
beta_scale acceleration-alignment multiplier (repo 0.5)
|
||||
upsampling "latent" (bicubic on latents, repo default) | "pixel"
|
||||
eps_sigma sigma floor in (x - x0)/sigma (reference uses 1e-6)
|
||||
"""
|
||||
|
||||
tau: float = 0.6
|
||||
steps: int = 30
|
||||
steps_per_stage: int = 16
|
||||
cfg: float = 3.5
|
||||
guidance_high: float = 4.5
|
||||
filter_ratio: float = 0.2
|
||||
alpha_scale: float = 1.0
|
||||
beta_scale: float = 0.5
|
||||
upsampling: str = "latent"
|
||||
eps_sigma: float = 1e-6
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def base_trajectory(
|
||||
initial_latent: Tensor,
|
||||
sigmas: torch.Tensor,
|
||||
predict_x0: Callable[[Tensor, float], Tensor],
|
||||
cfg: HiFlowConfig,
|
||||
progress_callback: Callable[[int, int, int], None] | None = None,
|
||||
) -> tuple[Tensor, TrajectoryDict]:
|
||||
"""Stage A — ordinary rectified-flow sampling that records every step.
|
||||
|
||||
Euler rule (paper Sec. 1): ``v = (x_t - x0_pred) / t``,
|
||||
``x_{t-1} = x_t + v * (t_{t-1} - t)`` with the velocity derived from the
|
||||
predicted clean sample (``sigma`` clamped at ``cfg.eps_sigma`` — the
|
||||
reference divides by ``sigma + 1e-6``). The final sample is recorded at
|
||||
sigma 0 as the endpoint prediction, mirroring the reference's
|
||||
``clean_predictions[0] = X``.
|
||||
|
||||
Runs in fp32 regardless of the input dtype; returns the final latent in
|
||||
the input's dtype and the per-step trajectory (VAE space throughout —
|
||||
the predict_x0 adapter owns model-space conversions).
|
||||
"""
|
||||
x = initial_latent.float()
|
||||
traj = TrajectoryDict()
|
||||
n_transitions = sigmas.numel() - 1
|
||||
|
||||
for i in range(n_transitions):
|
||||
sigma = float(sigmas[i])
|
||||
x0 = predict_x0(x, sigma).float()
|
||||
traj.put(sigma, x0)
|
||||
|
||||
sigma_safe = max(sigma, cfg.eps_sigma)
|
||||
v = (x - x0) / sigma_safe
|
||||
x = x + v * (float(sigmas[i + 1]) - sigma)
|
||||
|
||||
if progress_callback is not None:
|
||||
progress_callback(i, n_transitions, -1)
|
||||
|
||||
# Endpoint: the last state IS the sigma-0 clean prediction.
|
||||
traj.put(0.0, x)
|
||||
return x.to(initial_latent.dtype), traj
|
||||
|
||||
@@ -316,3 +316,176 @@ class TestAlignmentScales:
|
||||
def test_zero_entry_raises(self):
|
||||
with pytest.raises(ValueError, match="positive"):
|
||||
alignment_scales(torch.tensor([0.0, 0.0]))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Step 3 — TrajectoryDict + base_trajectory
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
from src.hiflow import HiFlowConfig, TrajectoryDict, base_trajectory # noqa: E402
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestTrajectoryDict:
|
||||
def test_put_get_roundtrip(self):
|
||||
traj = TrajectoryDict()
|
||||
x0 = torch.randn(1, 4, 8, 8)
|
||||
traj.put(0.5679, x0)
|
||||
out = traj.get(0.5679)
|
||||
assert torch.equal(out, x0)
|
||||
|
||||
def test_stored_on_cpu(self):
|
||||
"""D8: entries park on CPU so a long 4K trajectory stays off-GPU."""
|
||||
traj = TrajectoryDict()
|
||||
traj.put(0.5, torch.randn(1, 2, 4, 4))
|
||||
stored = traj.as_dict()[0.5]
|
||||
assert stored.device.type == "cpu"
|
||||
|
||||
def test_get_missing_key_message(self):
|
||||
traj = TrajectoryDict()
|
||||
traj.put(0.5, torch.zeros(1))
|
||||
with pytest.raises(KeyError, match="no trajectory entry"):
|
||||
traj.get(0.123)
|
||||
|
||||
def test_nearest_exact(self):
|
||||
traj = TrajectoryDict()
|
||||
x0 = torch.randn(1, 2, 4, 4)
|
||||
traj.put(0.4321, x0)
|
||||
key, out = traj.nearest(0.4321)
|
||||
assert key == 0.4321
|
||||
assert torch.equal(out, x0)
|
||||
|
||||
def test_nearest_off_key_within_tol(self):
|
||||
traj = TrajectoryDict()
|
||||
traj.put(0.4, torch.zeros(1))
|
||||
traj.put(0.2, torch.ones(1))
|
||||
# 0.0004 from 0.4 — within the 1e-3 tolerance, far from 0.2.
|
||||
key, out = traj.nearest(0.4004)
|
||||
assert key == 0.4
|
||||
|
||||
def test_nearest_tol_exceeded_raises(self):
|
||||
traj = TrajectoryDict()
|
||||
traj.put(0.4, torch.zeros(1))
|
||||
with pytest.raises(ValueError, match="within"):
|
||||
traj.nearest(0.6, tol=1e-3)
|
||||
|
||||
def test_nearest_empty_dict_raises(self):
|
||||
traj = TrajectoryDict()
|
||||
with pytest.raises(ValueError, match="reference trajectory"):
|
||||
traj.nearest(0.5)
|
||||
|
||||
def test_nearest_moves_to_device(self):
|
||||
traj = TrajectoryDict()
|
||||
traj.put(0.5, torch.zeros(1))
|
||||
_, out = traj.nearest(0.5, device=torch.device("cpu"))
|
||||
assert out.device.type == "cpu"
|
||||
|
||||
def test_sigma_rounding(self):
|
||||
"""Keys are 6-decimal rounded: 0.499999999 vs 0.5 share one slot."""
|
||||
traj = TrajectoryDict()
|
||||
traj.put(0.499999999, torch.zeros(1))
|
||||
assert torch.equal(traj.get(0.5), torch.zeros(1))
|
||||
assert len(traj) == 1
|
||||
|
||||
def test_sigmas_sorted(self):
|
||||
traj = TrajectoryDict()
|
||||
for s in (0.3, 0.9, 0.1):
|
||||
traj.put(s, torch.zeros(1))
|
||||
assert traj.sigmas() == [0.1, 0.3, 0.9]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestBaseTrajectory:
|
||||
CFG = HiFlowConfig(steps=5)
|
||||
|
||||
def _sigmas(self):
|
||||
return torch.tensor([1.0, 0.8, 0.6, 0.4, 0.2, 0.0])
|
||||
|
||||
def test_records_every_sigma_plus_endpoint(self):
|
||||
sigmas = self._sigmas()
|
||||
traj_out = {}
|
||||
|
||||
def predict(x, s):
|
||||
traj_out[s] = True
|
||||
return torch.zeros_like(x)
|
||||
|
||||
_, traj = base_trajectory(
|
||||
torch.randn(1, 4, 8, 8), sigmas, predict, self.CFG)
|
||||
for s in sigmas[:-1].tolist():
|
||||
assert traj.get(s) is not None, f"missing recorded sigma {s}"
|
||||
assert traj.get(0.0) is not None, "endpoint at sigma 0 must be stored"
|
||||
assert len(traj) == 6
|
||||
|
||||
def test_final_matches_reference_euler(self):
|
||||
"""Euler fidelity tripwire: analytic model, hand-computed walk.
|
||||
|
||||
predict_x0(x, s) = 0.5 * x => v = (x - 0.5x)/s = x/(2s), so
|
||||
x_next = x + (s_next - s) * x/(2s) — deterministic closed form.
|
||||
"""
|
||||
sigmas = self._sigmas()
|
||||
x0 = torch.randn(1, 2, 4, 4)
|
||||
|
||||
x = x0.clone()
|
||||
for i in range(len(sigmas) - 1):
|
||||
s = float(sigmas[i])
|
||||
v = x / (2.0 * s)
|
||||
x = x + v * (float(sigmas[i + 1]) - s)
|
||||
|
||||
_, traj = base_trajectory(
|
||||
x0.clone(), sigmas, lambda x, s: 0.5 * x, self.CFG)
|
||||
final = traj.get(0.0)
|
||||
assert torch.allclose(final, x, atol=1e-5), (
|
||||
"base_trajectory must reproduce the hand-rolled Euler walk"
|
||||
)
|
||||
|
||||
def test_calls_predict_in_sigma_order(self):
|
||||
calls = []
|
||||
|
||||
def predict(x, s):
|
||||
calls.append(s)
|
||||
return torch.zeros_like(x)
|
||||
|
||||
base_trajectory(
|
||||
torch.randn(1, 2, 4, 4), self._sigmas(), predict, self.CFG)
|
||||
# float32 sigmas round-trip with tiny artifacts — compare 6-decimal.
|
||||
assert [round(s, 6) for s in calls] == [1.0, 0.8, 0.6, 0.4, 0.2]
|
||||
|
||||
def test_records_time_matched_x0(self):
|
||||
"""Each stored entry must be the x0 predicted AT that sigma, not a
|
||||
later state (time matching is the whole point of the reference)."""
|
||||
seen = {}
|
||||
|
||||
def predict(x, s):
|
||||
x0 = x + s # sigma-identifiable values
|
||||
seen[s] = x0.clone()
|
||||
return x0
|
||||
|
||||
x_start = torch.zeros(1, 2, 2, 2)
|
||||
_, traj = base_trajectory(x_start, self._sigmas(), predict, self.CFG)
|
||||
for s, x0_at in seen.items():
|
||||
assert torch.equal(traj.get(s), x0_at)
|
||||
|
||||
def test_progress_callback_events(self):
|
||||
events = []
|
||||
base_trajectory(
|
||||
torch.randn(1, 2, 4, 4), self._sigmas(),
|
||||
lambda x, s: torch.zeros_like(x), self.CFG,
|
||||
progress_callback=lambda i, total, stage: events.append((i, total, stage)),
|
||||
)
|
||||
assert events == [
|
||||
(0, 5, -1), (1, 5, -1), (2, 5, -1), (3, 5, -1), (4, 5, -1),
|
||||
]
|
||||
|
||||
def test_no_nan_fp16_input(self):
|
||||
x = torch.randn(1, 2, 4, 4, dtype=torch.float16)
|
||||
out, traj = base_trajectory(
|
||||
x, self._sigmas(), lambda x_, s: 0.5 * x_, self.CFG)
|
||||
assert out.dtype == torch.float16
|
||||
assert torch.isfinite(out).all()
|
||||
assert torch.isfinite(traj.get(0.0)).all()
|
||||
|
||||
def test_output_shape_preserved(self):
|
||||
x = torch.randn(2, 3, 16, 12)
|
||||
out, _ = base_trajectory(
|
||||
x, self._sigmas(), lambda x_, s: 0.5 * x_, self.CFG)
|
||||
assert out.shape == x.shape
|
||||
|
||||
Reference in New Issue
Block a user