Fix Wan previewing, add symmetric ortho blend modes

This commit is contained in:
blepping
2026-01-28 05:43:03 -07:00
parent 57f289d807
commit 006767817b
4 changed files with 160 additions and 62 deletions
+5
View File
@@ -4,6 +4,8 @@ from typing import TYPE_CHECKING, NamedTuple
from comfy import latent_formats
from .tae_vid import TAEVid, TAEVidBase, TAEVidLTX2
if TYPE_CHECKING:
from pathlib import Path
@@ -17,6 +19,7 @@ class VideoModelInfo(NamedTuple):
patch_size: int = 1
nested_tensor_index: int = 0
tae_model: str | Path | None = None
tae_class: TAEVidBase | None = TAEVid
VIDEO_FORMATS = {
@@ -69,6 +72,7 @@ VIDEO_FORMATS = {
patch_size=4,
temporal_layers=3,
tae_model="taeltx_2.pth",
tae_class=TAEVidLTX2,
),
VideoModelInfo(
"ltxav",
@@ -77,6 +81,7 @@ VIDEO_FORMATS = {
patch_size=4,
temporal_layers=3,
tae_model="taeltx_2.pth",
tae_class=TAEVidLTX2,
),
)
}
+7 -7
View File
@@ -748,21 +748,21 @@ def bleh_get_previewer(
)
tae_model = None
if preview_method in {LatentPreviewMethod.TAESD, LatentPreviewMethod.Auto}:
if vid_info is not None and vid_info.tae_model is not None:
if (
vid_info is not None
and vid_info.tae_model is not None
and vid_info.tae_class is not None
):
tae_model_path = folder_paths.get_full_path(
"vae_approx",
vid_info.tae_model,
)
tupscale_limit = SETTINGS.btp_video_temporal_upscale_level
decoder_time_upscale = tuple(
i < tupscale_limit for i in range(TAEVid.temporal_upscale_blocks)
)
tae_model = (
TAEVid(
vid_info.tae_class(
checkpoint_path=tae_model_path,
vmi=vid_info,
device=torch.device("cpu"),
decoder_time_upscale=decoder_time_upscale,
decoder_time_upscale_level=SETTINGS.btp_video_temporal_upscale_level,
)
if tae_model_path is not None
else None
+122 -55
View File
@@ -53,6 +53,10 @@ class MemBlock(nn.Module):
return self.act(self.conv(torch.cat((x, past), 1)) + self.skip(x))
def make_memblocks(n: int, *, count: int = 3) -> tuple[MemBlock, ...]:
return tuple(MemBlock(n, n) for _ in range(count))
class TPool(nn.Module):
def __init__(self, n_f, stride):
super().__init__()
@@ -183,7 +187,7 @@ class TAEVidContext:
return torch.stack(out, 1)
class TAEVid(nn.Module):
class TAEVidBase(nn.Module):
temporal_upscale_blocks = 3
spatial_upscale_blocks = 3
_nf = (256, 128, 64, 64)
@@ -195,68 +199,29 @@ class TAEVid(nn.Module):
vmi: VideoModelInfo,
image_channels: int = 3,
device="cpu",
encoder_time_downscale=(True, True, False),
decoder_time_upscale=(False, True, True),
decoder_space_upscale=(True, True, True),
encoder_time_downscale_level: int = 3,
decoder_time_upscale_level: int = 3,
decoder_space_upscale_level: int = 3,
):
n_f = self._nf
super().__init__()
if len(decoder_time_upscale) == 2:
decoder_time_upscale = (True, *decoder_time_upscale)
self.vmi = vmi
self.latent_channels = vmi.latent_format.latent_channels
self.image_channels = image_channels
self.latent_channels = vmi.latent_format.latent_channels
self.patch_size = vmi.patch_size
if vmi.name in {"ltxv", "ltxav"}:
encoder_time_downscale = (True, True, True)
decoder_time_upscale = (True, True, True)
encoder_time_downscale = self._get_encoder_flags(
time_level=encoder_time_downscale_level,
)
decoder_time_upscale, decoder_space_upscale = self._get_decoder_flags(
time_level=decoder_time_upscale_level,
space_level=decoder_space_upscale_level,
)
encoder_strides = tuple(1 + int(flag) for flag in encoder_time_downscale)
decoder_strides = tuple(1 + int(flag) for flag in decoder_time_upscale)
decoder_scale_factors = tuple(1 + int(flag) for flag in decoder_space_upscale)
self.encoder = nn.Sequential(
conv(image_channels * self.patch_size**2, 64),
nn.ReLU(inplace=True),
TPool(64, encoder_strides[0]),
conv(64, 64, stride=2, bias=False),
MemBlock(64, 64),
MemBlock(64, 64),
MemBlock(64, 64),
TPool(64, encoder_strides[1]),
conv(64, 64, stride=2, bias=False),
MemBlock(64, 64),
MemBlock(64, 64),
MemBlock(64, 64),
TPool(64, encoder_strides[2]),
conv(64, 64, stride=2, bias=False),
MemBlock(64, 64),
MemBlock(64, 64),
MemBlock(64, 64),
conv(64, vmi.latent_format.latent_channels),
)
self.decoder = nn.Sequential(
Clamp(),
conv(vmi.latent_format.latent_channels, n_f[0]),
nn.ReLU(inplace=True),
MemBlock(n_f[0], n_f[0]),
MemBlock(n_f[0], n_f[0]),
MemBlock(n_f[0], n_f[0]),
nn.Upsample(scale_factor=decoder_scale_factors[0]),
TGrow(n_f[0], decoder_strides[0]),
conv(n_f[0], n_f[1], bias=False),
MemBlock(n_f[1], n_f[1]),
MemBlock(n_f[1], n_f[1]),
MemBlock(n_f[1], n_f[1]),
nn.Upsample(scale_factor=decoder_scale_factors[1]),
TGrow(n_f[1], decoder_strides[1]),
conv(n_f[1], n_f[2], bias=False),
MemBlock(n_f[2], n_f[2]),
MemBlock(n_f[2], n_f[2]),
MemBlock(n_f[2], n_f[2]),
nn.Upsample(scale_factor=decoder_scale_factors[2]),
TGrow(n_f[2], decoder_strides[2]),
conv(n_f[2], n_f[3], bias=False),
nn.ReLU(inplace=True),
conv(n_f[3], image_channels * self.patch_size**2),
self.encoder = self._build_encoder(strides=encoder_strides)
self.decoder = self._build_decoder(
strides=decoder_strides,
scale_factors=decoder_scale_factors,
)
self.t_upscale = 2 ** sum(decoder_time_upscale)
self.t_downscale = 2 ** sum(encoder_time_downscale)
@@ -269,6 +234,66 @@ class TAEVid(nn.Module):
),
)
def _get_decoder_flags(
self,
*,
time_level: int = 3,
space_level: int = 3,
) -> tuple[tuple[bool, ...], tuple[bool, ...]]:
decoder_time_upscale = tuple(i < time_level for i in range(3))
decoder_space_upscale = tuple(i < space_level for i in range(3))
return decoder_time_upscale, decoder_space_upscale
def _get_encoder_flags(
self,
*,
time_level: int = 3,
) -> tuple[bool, ...]:
return tuple(i < time_level for i in range(3))
def _build_decoder(
self,
*,
strides: tuple[int, ...],
scale_factors: tuple[int, ...],
) -> nn.Module:
n_f = self._nf
return nn.Sequential(
Clamp(),
conv(self.latent_channels, n_f[0]),
nn.ReLU(inplace=True),
*make_memblocks(n_f[0]),
nn.Upsample(scale_factor=scale_factors[0]),
TGrow(n_f[0], strides[0]),
conv(n_f[0], n_f[1], bias=False),
*make_memblocks(n_f[1]),
nn.Upsample(scale_factor=scale_factors[1]),
TGrow(n_f[1], strides[1]),
conv(n_f[1], n_f[2], bias=False),
*make_memblocks(n_f[2]),
nn.Upsample(scale_factor=scale_factors[2]),
TGrow(n_f[2], strides[2]),
conv(n_f[2], n_f[3], bias=False),
nn.ReLU(inplace=True),
conv(n_f[3], self.image_channels * self.patch_size**2),
)
def _build_encoder(self, *, strides: tuple[int, ...]) -> nn.Module:
return nn.Sequential(
conv(self.image_channels * self.patch_size**2, 64),
nn.ReLU(inplace=True),
TPool(64, strides[0]),
conv(64, 64, stride=2, bias=False),
*make_memblocks(64),
TPool(64, strides[1]),
conv(64, 64, stride=2, bias=False),
*make_memblocks(64),
TPool(64, strides[2]),
conv(64, 64, stride=2, bias=False),
*make_memblocks(64),
conv(64, self.latent_channels),
)
def patch_tgrow_layers(self, sd: dict) -> dict:
new_sd = self.state_dict()
for i, layer in enumerate(self.decoder):
@@ -342,3 +367,45 @@ class TAEVid(nn.Module):
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.c(x)
class TAEVid(TAEVidBase):
def _get_decoder_flags(
self,
*,
time_level: int = 3,
space_level: int = 3,
) -> tuple[tuple[bool, ...], tuple[bool, ...]]:
tu, su = super()._get_decoder_flags(
time_level=time_level,
space_level=space_level,
)
return (False, *tu[:2]), su
def _get_encoder_flags(
self,
*,
time_level: int = 3,
) -> tuple[bool, ...]:
return (*super()._get_encoder_flags(time_level=time_level)[:2], False)
class TAEVidLTX2(TAEVidBase):
def _get_decoder_flags(
self,
*,
time_level: int = 3,
space_level: int = 3,
) -> tuple[tuple[bool, ...], tuple[bool, ...]]:
_tu, su = super()._get_decoder_flags(
time_level=time_level,
space_level=space_level,
)
return (True, True, True), su
def _get_encoder_flags(
self,
*,
time_level: int = 3, # noqa: ARG002
) -> tuple[bool, ...]:
return (True, True, True)
+26
View File
@@ -944,6 +944,30 @@ def ortho_blend(
return ortho_result.reshape(orig_shape)
def symmetric_ortho_blend(
a: torch.Tensor,
b: torch.Tensor,
t: torch.Tensor,
*,
symmetric_strength: float = 1.0,
symmetric_deduce_mode: bool = False,
**kwargs: dict,
) -> torch.Tensor:
blended = ortho_blend(a, b, t, **kwargs)
if symmetric_strength == 0.0:
return blended
b_ortho = blended.sub_(a)
if symmetric_deduce_mode:
b_proj = b - b_ortho
# Projection would theoretically be the same for both, in the simple case at least?
# Actually, probably not. Oh well, this is here as an option now.
a_ortho = a - b_proj
else:
a_ortho = ortho_blend(b, a, a.new_tensor(1.0), **kwargs) - b
a_proj = a - a_ortho
return a_proj.mul_(1.0 - symmetric_strength).add_(a_ortho).add_(b_ortho)
class WaveletBlend:
wavelet: wavef.Wavelet | None = None
use_float64: bool = False
@@ -1738,6 +1762,8 @@ BLENDING_MODES = {
rescale_result_mode="blend",
rescale_limit=2.0,
),
"symmetric_ortho": BlendMode(symmetric_ortho_blend),
"symmetric_ortho_rescaled": BlendMode(symmetric_ortho_blend, rescale_limit=2.0),
}
BLENDING_MODES |= {