Merge pull request #8 from xmarre/codex/wan-aspect-adaptive-base
[codex] Add WAN aspect-adaptive base resolution
This commit is contained in:
@@ -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
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user