diff --git a/py/better_previews/base.py b/py/better_previews/base.py index 8c3178e..f08df3c 100644 --- a/py/better_previews/base.py +++ b/py/better_previews/base.py @@ -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, ), ) } diff --git a/py/better_previews/previewer.py b/py/better_previews/previewer.py index 9d41c35..00547de 100644 --- a/py/better_previews/previewer.py +++ b/py/better_previews/previewer.py @@ -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 diff --git a/py/better_previews/tae_vid.py b/py/better_previews/tae_vid.py index 8bba95e..f676f11 100644 --- a/py/better_previews/tae_vid.py +++ b/py/better_previews/tae_vid.py @@ -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) diff --git a/py/latent_utils.py b/py/latent_utils.py index 0c7b26e..873ee02 100644 --- a/py/latent_utils.py +++ b/py/latent_utils.py @@ -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 |= {