diff --git a/src/hiflow.py b/src/hiflow.py index 40d1623..4092f0a 100644 --- a/src/hiflow.py +++ b/src/hiflow.py @@ -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 diff --git a/tests/test_hiflow.py b/tests/test_hiflow.py index 0278dde..cafb43e 100644 --- a/tests/test_hiflow.py +++ b/tests/test_hiflow.py @@ -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