Merge pull request #8 from xmarre/codex/wan-aspect-adaptive-base

[codex] Add WAN aspect-adaptive base resolution
This commit is contained in:
xmarre
2026-05-15 19:35:59 +02:00
committed by GitHub
4 changed files with 137 additions and 8 deletions
+38 -7
View File
@@ -3,15 +3,15 @@ from __future__ import annotations
from typing import Any
try:
from .tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper, install_tide_wan_patch
from .tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper, aspect_adaptive_base_resolution, install_tide_wan_patch
except ModuleNotFoundError as exc:
if exc.name not in {f"{__package__}.tide_core", "tide_core"}:
raise
from tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper, install_tide_wan_patch
from tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper, aspect_adaptive_base_resolution, install_tide_wan_patch
except ImportError as exc:
if "attempted relative import with no known parent package" not in str(exc):
raise
from tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper, install_tide_wan_patch
from tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper, aspect_adaptive_base_resolution, install_tide_wan_patch
class TIDEHighResolutionExtrapolation:
@@ -110,6 +110,24 @@ class TIDEHighResolutionExtrapolation:
return (patched,)
def _resolve_wan_base_resolution(
*,
width: int,
height: int,
base_width: int,
base_height: int,
base_resolution_mode: str,
) -> tuple[int, int]:
mode = str(base_resolution_mode)
if mode == "manual":
return int(base_width), int(base_height)
if mode == "aspect_adaptive_720p":
return aspect_adaptive_base_resolution(width, height, 1280, 720)
if mode == "aspect_adaptive_480p":
return aspect_adaptive_base_resolution(width, height, 832, 480)
raise ValueError(f"Unsupported WAN base_resolution_mode: {base_resolution_mode!r}")
class TIDEWANHighResolutionExtrapolation:
"""Patch a WAN 2.1/2.2-style DiT model with TIDE Dynamic Temperature Control."""
@@ -126,8 +144,12 @@ class TIDEWANHighResolutionExtrapolation:
),
},
"optional": {
"base_width": ("INT", {"default": 640, "min": 16, "max": 16384, "step": 16}),
"base_height": ("INT", {"default": 640, "min": 16, "max": 16384, "step": 16}),
"base_resolution_mode": (
["aspect_adaptive_720p", "aspect_adaptive_480p", "manual"],
{"default": "aspect_adaptive_720p", "tooltip": "WAN native reference. Adaptive modes preserve target aspect and keep a 1280x720 or 832x480 native pixel budget; manual uses base_width/base_height exactly."},
),
"base_width": ("INT", {"default": 1280, "min": 16, "max": 16384, "step": 16}),
"base_height": ("INT", {"default": 720, "min": 16, "max": 16384, "step": 16}),
"alpha_low": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 8.0, "step": 0.05}),
"alpha_high": ("FLOAT", {"default": 0.2, "min": 0.0, "max": 8.0, "step": 0.05}),
"tau_max": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 4.0, "step": 0.01}),
@@ -149,8 +171,9 @@ class TIDEWANHighResolutionExtrapolation:
width: int,
height: int,
temperature_strength: float,
base_width: int = 640,
base_height: int = 640,
base_resolution_mode: str = "aspect_adaptive_720p",
base_width: int = 1280,
base_height: int = 720,
alpha_low: float = 0.6,
alpha_high: float = 0.2,
tau_max: float = 1.0,
@@ -159,6 +182,14 @@ class TIDEWANHighResolutionExtrapolation:
preserve_existing_wrapper: bool = True,
debug: bool = False,
):
base_width, base_height = _resolve_wan_base_resolution(
width=int(width),
height=int(height),
base_width=int(base_width),
base_height=int(base_height),
base_resolution_mode=str(base_resolution_mode),
)
config = TIDEConfig(
width=int(width),
height=int(height),
+31 -1
View File
@@ -2,13 +2,14 @@ import math
import pathlib
import sys
import pytest
import torch
ROOT = pathlib.Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from tide_core.config import TIDEConfig
from tide_core.math import adaptive_text_bias, get_default_temperature, rope_temperature_scale
from tide_core.math import adaptive_text_bias, aspect_adaptive_base_resolution, get_default_temperature, rope_temperature_scale
def test_adaptive_text_bias_matches_paper_and_official_flux_script():
@@ -43,3 +44,32 @@ def test_temperature_strength_zero_disables_scaling():
cfg = TIDEConfig(width=4096, height=4096, temperature_strength=0.0)
scale = rope_temperature_scale(cfg, timestep=1.0, device=torch.device("cpu"))
assert torch.allclose(scale, torch.ones_like(scale), atol=1e-6)
def test_aspect_adaptive_base_resolution_preserves_720p_area_budget_for_square_targets():
assert aspect_adaptive_base_resolution(960, 960, 1280, 720) == (960, 960)
def test_aspect_adaptive_base_resolution_preserves_native_landscape_and_portrait():
assert aspect_adaptive_base_resolution(1280, 720, 1280, 720) == (1280, 720)
assert aspect_adaptive_base_resolution(720, 1280, 1280, 720) == (720, 1280)
def test_aspect_adaptive_base_resolution_uses_target_aspect_for_nonstandard_i2v_size():
assert aspect_adaptive_base_resolution(896, 656, 1280, 720) == (1120, 816)
def test_aspect_adaptive_base_resolution_keeps_snapped_result_within_native_budget():
base_width, base_height = aspect_adaptive_base_resolution(1000, 777, 1280, 720)
assert base_width * base_height <= 1280 * 720
def test_aspect_adaptive_base_resolution_preserves_wan_480p_budget():
assert aspect_adaptive_base_resolution(832, 480, 832, 480) == (832, 480)
base_width, base_height = aspect_adaptive_base_resolution(480, 832, 832, 480)
assert base_width * base_height <= 832 * 480
def test_aspect_adaptive_base_resolution_rejects_unsnappable_budget():
with pytest.raises(ValueError, match="No grid-aligned base resolution"):
aspect_adaptive_base_resolution(100_000_000, 1, 1280, 720)
+2
View File
@@ -1,6 +1,7 @@
from .config import TIDEConfig
from .math import (
adaptive_text_bias,
aspect_adaptive_base_resolution,
get_default_temperature,
get_mscale,
rope_temperature_scale,
@@ -11,6 +12,7 @@ from .wan import TIDEWanDiffusionWrapper, install_tide_wan_patch
__all__ = [
"TIDEConfig",
"adaptive_text_bias",
"aspect_adaptive_base_resolution",
"get_default_temperature",
"get_mscale",
"rope_temperature_scale",
+66
View File
@@ -26,6 +26,72 @@ def get_default_temperature(scale: float) -> float:
return 1.0 / (mscale * mscale)
def aspect_adaptive_base_resolution(
width: int,
height: int,
native_width: int,
native_height: int,
*,
grid: int = 16,
) -> tuple[int, int]:
"""Return native-budget base dimensions matched to the target aspect.
WAN 480p/720p checkpoints are better represented as a native pixel budget
than as one fixed landscape rectangle. For an arbitrary I2V aspect, preserve
the target aspect while keeping native_width * native_height as the maximum
native reference area used by TIDE's extrapolation gates and RoPE-axis scale
factors.
"""
width = int(width)
height = int(height)
native_width = int(native_width)
native_height = int(native_height)
grid = int(grid)
if width <= 0 or height <= 0:
raise ValueError(f"width and height must be > 0, got {width}x{height}")
if native_width <= 0 or native_height <= 0:
raise ValueError(f"native base must be > 0, got {native_width}x{native_height}")
if grid <= 0:
raise ValueError(f"grid must be > 0, got {grid}")
native_area = float(native_width) * float(native_height)
aspect = float(width) / float(height)
raw_width = math.sqrt(native_area * aspect)
raw_height = math.sqrt(native_area / aspect)
def snapped_neighbors(value: float) -> tuple[int, int]:
scaled = value / grid
lower = max(grid, int(math.floor(scaled)) * grid)
upper = max(grid, int(math.ceil(scaled)) * grid)
return lower, upper
candidates = {
(candidate_width, candidate_height)
for candidate_width in snapped_neighbors(raw_width)
for candidate_height in snapped_neighbors(raw_height)
}
under_budget = [
candidate
for candidate in candidates
if candidate[0] * candidate[1] <= native_width * native_height
]
if not under_budget:
raise ValueError(
"No grid-aligned base resolution satisfies the native area budget "
f"for {width}x{height} against {native_width}x{native_height} on grid {grid}"
)
candidates = set(under_budget)
def score(candidate: tuple[int, int]) -> tuple[float, float]:
candidate_width, candidate_height = candidate
area_error = abs((candidate_width * candidate_height) - native_area) / native_area
aspect_error = abs((candidate_width / candidate_height) - aspect) / aspect
return area_error, aspect_error
return min(candidates, key=score)
def adaptive_text_bias(config: TIDEConfig) -> float:
"""Paper Eq. 17/18: beta = log(lambda), with lambda = pixel ratio.