Squashed commit of the following:
commit 73dd1a06d33953912f5dd684f168028b14e42a36 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Mon Oct 13 19:47:38 2025 +0300 cleanup commit 39bc2cecf493e2eb176b55e8841d933f0da1ec39 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Mon Oct 13 19:24:20 2025 +0300 Allow scheduling ovi cfg commit 2c153c5f324dbd59670ad9c51a7995459504a3cd Merge: dba766732eb6b4Author: kijai <40791699+kijai@users.noreply.github.com> Date: Mon Oct 13 17:48:20 2025 +0300 Merge branch 'main' into ovi commit dba76674c71af7bf94c82834a0b0e40d94043c99 Merge: 0f11a435a0456eAuthor: kijai <40791699+kijai@users.noreply.github.com> Date: Sun Oct 12 22:45:43 2025 +0300 Merge branch 'main' into ovi commit 0f11a439622799ad8070f8a2b8cc8e6a041b761d Merge: 0999f50e2d8c9bAuthor: kijai <40791699+kijai@users.noreply.github.com> Date: Sat Oct 11 07:48:06 2025 +0300 Merge branch 'main' into ovi commit 0999f50cfe025290cd7ce88a8dd1acff0b38d9bd Merge: d45df1ff1d1c83Author: kijai <40791699+kijai@users.noreply.github.com> Date: Fri Oct 10 22:16:09 2025 +0300 Merge branch 'main' into ovi commit d45df1fb5b7c629b15eabc197357d62bdc232aaf Author: kijai <40791699+kijai@users.noreply.github.com> Date: Thu Oct 9 20:21:37 2025 +0300 Remove dependency for librosa commit d8e7533fdf7eab1d2489c3e025a908c02d997444 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Thu Oct 9 19:57:28 2025 +0300 Remove omegaconf dependency commit f4e27ff018e98cb5b09655dceda399baea36b240 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Thu Oct 9 19:31:06 2025 +0300 Fix VACE commit 35d3df39294831e5e7568b6f7e16d2ecf2d790a0 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Thu Oct 9 00:26:40 2025 +0300 small update commit 96f8ea1d26869ab7e49e12a07f19d5d5a2023253 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Wed Oct 8 22:32:57 2025 +0300 Create wanvideo_2_2_5B_ovi_testing.json commit a2511be73b9da7019fd21aeb0b521af941c09150 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Wed Oct 8 22:32:54 2025 +0300 Update nodes_sampler.py commit d3688b8db71452ea1f7c9a2bc0216441d524e56c Author: kijai <40791699+kijai@users.noreply.github.com> Date: Wed Oct 8 21:43:02 2025 +0300 Allow EasyCache to work with ovi commit 586d9148a0306ef5d30e9a971a9c3be4cd3ecc97 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Wed Oct 8 19:09:06 2025 +0300 Update model.py commit 61eedd2839decdb7d4c2ddd5f1310fdaf49d36ad Author: kijai <40791699+kijai@users.noreply.github.com> Date: Wed Oct 8 19:09:02 2025 +0300 I2V fix commit a97fcb1b9ae9fb7bbfdf668c24816e014a1b58d1 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Wed Oct 8 17:57:28 2025 +0300 Add nodes to set audio latent size commit d41e42a697f3d561dabbc22566f633b5f1bbd952 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Wed Oct 8 16:42:04 2025 +0300 Support loading mmaudio vae from .safetensors commit 1b0e28ec41e3c97fe1f2f057fef9b9bbcb87bca7 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Wed Oct 8 16:19:53 2025 +0300 Update nodes_sampler.py commit fbd18f45fe85ede8edcb5aebaea7ceb5b6eab5a2 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Wed Oct 8 10:16:44 2025 +0300 Fixes for other workflows commit b06993b637198f7fad92208f3b3dc9a7d7f57c7f Author: kijai <40791699+kijai@users.noreply.github.com> Date: Wed Oct 8 09:46:27 2025 +0300 initial commit T2V works
This commit is contained in:
@@ -0,0 +1,48 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
class ChannelLastConv1d(nn.Conv1d):
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = x.permute(0, 2, 1)
|
||||
x = super().forward(x)
|
||||
x = x.permute(0, 2, 1)
|
||||
return x
|
||||
|
||||
|
||||
class ConvMLP(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
hidden_dim: int,
|
||||
multiple_of: int = 256,
|
||||
kernel_size: int = 3,
|
||||
padding: int = 1,
|
||||
):
|
||||
"""
|
||||
Initialize the FeedForward module.
|
||||
|
||||
Args:
|
||||
dim (int): Input dimension.
|
||||
hidden_dim (int): Hidden dimension of the feedforward layer.
|
||||
multiple_of (int): Value to ensure hidden dimension is a multiple of this value.
|
||||
|
||||
Attributes:
|
||||
w1 (ColumnParallelLinear): Linear transformation for the first layer.
|
||||
w2 (RowParallelLinear): Linear transformation for the second layer.
|
||||
w3 (ColumnParallelLinear): Linear transformation for the third layer.
|
||||
|
||||
"""
|
||||
super().__init__()
|
||||
hidden_dim = int(2 * hidden_dim / 3)
|
||||
hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
|
||||
|
||||
self.w1 = ChannelLastConv1d(dim, hidden_dim, bias=False, kernel_size=kernel_size, padding=padding)
|
||||
self.w2 = ChannelLastConv1d(hidden_dim, dim, bias=False, kernel_size=kernel_size, padding=padding)
|
||||
self.w3 = ChannelLastConv1d(dim, hidden_dim, bias=False, kernel_size=kernel_size, padding=padding)
|
||||
|
||||
def forward(self, x):
|
||||
return self.w2(F.silu(self.w1(x)) * self.w3(x))
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2022 NVIDIA CORPORATION.
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
@@ -0,0 +1 @@
|
||||
from .bigvgan import BigVGAN
|
||||
@@ -0,0 +1,120 @@
|
||||
# Implementation adapted from https://github.com/EdwardDixon/snake under the MIT license.
|
||||
# LICENSE is in incl_licenses directory.
|
||||
|
||||
import torch
|
||||
from torch import nn, sin, pow
|
||||
from torch.nn import Parameter
|
||||
|
||||
|
||||
class Snake(nn.Module):
|
||||
'''
|
||||
Implementation of a sine-based periodic activation function
|
||||
Shape:
|
||||
- Input: (B, C, T)
|
||||
- Output: (B, C, T), same shape as the input
|
||||
Parameters:
|
||||
- alpha - trainable parameter
|
||||
References:
|
||||
- This activation function is from this paper by Liu Ziyin, Tilman Hartwig, Masahito Ueda:
|
||||
https://arxiv.org/abs/2006.08195
|
||||
Examples:
|
||||
>>> a1 = snake(256)
|
||||
>>> x = torch.randn(256)
|
||||
>>> x = a1(x)
|
||||
'''
|
||||
def __init__(self, in_features, alpha=1.0, alpha_trainable=True, alpha_logscale=False):
|
||||
'''
|
||||
Initialization.
|
||||
INPUT:
|
||||
- in_features: shape of the input
|
||||
- alpha: trainable parameter
|
||||
alpha is initialized to 1 by default, higher values = higher-frequency.
|
||||
alpha will be trained along with the rest of your model.
|
||||
'''
|
||||
super(Snake, self).__init__()
|
||||
self.in_features = in_features
|
||||
|
||||
# initialize alpha
|
||||
self.alpha_logscale = alpha_logscale
|
||||
if self.alpha_logscale: # log scale alphas initialized to zeros
|
||||
self.alpha = Parameter(torch.zeros(in_features) * alpha)
|
||||
else: # linear scale alphas initialized to ones
|
||||
self.alpha = Parameter(torch.ones(in_features) * alpha)
|
||||
|
||||
self.alpha.requires_grad = alpha_trainable
|
||||
|
||||
self.no_div_by_zero = 0.000000001
|
||||
|
||||
def forward(self, x):
|
||||
'''
|
||||
Forward pass of the function.
|
||||
Applies the function to the input elementwise.
|
||||
Snake ∶= x + 1/a * sin^2 (xa)
|
||||
'''
|
||||
alpha = self.alpha.unsqueeze(0).unsqueeze(-1) # line up with x to [B, C, T]
|
||||
if self.alpha_logscale:
|
||||
alpha = torch.exp(alpha)
|
||||
x = x + (1.0 / (alpha + self.no_div_by_zero)) * pow(sin(x * alpha), 2)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class SnakeBeta(nn.Module):
|
||||
'''
|
||||
A modified Snake function which uses separate parameters for the magnitude of the periodic components
|
||||
Shape:
|
||||
- Input: (B, C, T)
|
||||
- Output: (B, C, T), same shape as the input
|
||||
Parameters:
|
||||
- alpha - trainable parameter that controls frequency
|
||||
- beta - trainable parameter that controls magnitude
|
||||
References:
|
||||
- This activation function is a modified version based on this paper by Liu Ziyin, Tilman Hartwig, Masahito Ueda:
|
||||
https://arxiv.org/abs/2006.08195
|
||||
Examples:
|
||||
>>> a1 = snakebeta(256)
|
||||
>>> x = torch.randn(256)
|
||||
>>> x = a1(x)
|
||||
'''
|
||||
def __init__(self, in_features, alpha=1.0, alpha_trainable=True, alpha_logscale=False):
|
||||
'''
|
||||
Initialization.
|
||||
INPUT:
|
||||
- in_features: shape of the input
|
||||
- alpha - trainable parameter that controls frequency
|
||||
- beta - trainable parameter that controls magnitude
|
||||
alpha is initialized to 1 by default, higher values = higher-frequency.
|
||||
beta is initialized to 1 by default, higher values = higher-magnitude.
|
||||
alpha will be trained along with the rest of your model.
|
||||
'''
|
||||
super(SnakeBeta, self).__init__()
|
||||
self.in_features = in_features
|
||||
|
||||
# initialize alpha
|
||||
self.alpha_logscale = alpha_logscale
|
||||
if self.alpha_logscale: # log scale alphas initialized to zeros
|
||||
self.alpha = Parameter(torch.zeros(in_features) * alpha)
|
||||
self.beta = Parameter(torch.zeros(in_features) * alpha)
|
||||
else: # linear scale alphas initialized to ones
|
||||
self.alpha = Parameter(torch.ones(in_features) * alpha)
|
||||
self.beta = Parameter(torch.ones(in_features) * alpha)
|
||||
|
||||
self.alpha.requires_grad = alpha_trainable
|
||||
self.beta.requires_grad = alpha_trainable
|
||||
|
||||
self.no_div_by_zero = 0.000000001
|
||||
|
||||
def forward(self, x):
|
||||
'''
|
||||
Forward pass of the function.
|
||||
Applies the function to the input elementwise.
|
||||
SnakeBeta ∶= x + 1/b * sin^2 (xa)
|
||||
'''
|
||||
alpha = self.alpha.unsqueeze(0).unsqueeze(-1) # line up with x to [B, C, T]
|
||||
beta = self.beta.unsqueeze(0).unsqueeze(-1)
|
||||
if self.alpha_logscale:
|
||||
alpha = torch.exp(alpha)
|
||||
beta = torch.exp(beta)
|
||||
x = x + (1.0 / (beta + self.no_div_by_zero)) * pow(sin(x * alpha), 2)
|
||||
|
||||
return x
|
||||
@@ -0,0 +1,6 @@
|
||||
# Adapted from https://github.com/junjun3518/alias-free-torch under the Apache License 2.0
|
||||
# LICENSE is in incl_licenses directory.
|
||||
|
||||
from .filter import *
|
||||
from .resample import *
|
||||
from .act import *
|
||||
@@ -0,0 +1,28 @@
|
||||
# Adapted from https://github.com/junjun3518/alias-free-torch under the Apache License 2.0
|
||||
# LICENSE is in incl_licenses directory.
|
||||
|
||||
import torch.nn as nn
|
||||
from .resample import UpSample1d, DownSample1d
|
||||
|
||||
|
||||
class Activation1d(nn.Module):
|
||||
def __init__(self,
|
||||
activation,
|
||||
up_ratio: int = 2,
|
||||
down_ratio: int = 2,
|
||||
up_kernel_size: int = 12,
|
||||
down_kernel_size: int = 12):
|
||||
super().__init__()
|
||||
self.up_ratio = up_ratio
|
||||
self.down_ratio = down_ratio
|
||||
self.act = activation
|
||||
self.upsample = UpSample1d(up_ratio, up_kernel_size)
|
||||
self.downsample = DownSample1d(down_ratio, down_kernel_size)
|
||||
|
||||
# x: [B,C,T]
|
||||
def forward(self, x):
|
||||
x = self.upsample(x)
|
||||
x = self.act(x)
|
||||
x = self.downsample(x)
|
||||
|
||||
return x
|
||||
@@ -0,0 +1,95 @@
|
||||
# Adapted from https://github.com/junjun3518/alias-free-torch under the Apache License 2.0
|
||||
# LICENSE is in incl_licenses directory.
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import math
|
||||
|
||||
if 'sinc' in dir(torch):
|
||||
sinc = torch.sinc
|
||||
else:
|
||||
# This code is adopted from adefossez's julius.core.sinc under the MIT License
|
||||
# https://adefossez.github.io/julius/julius/core.html
|
||||
# LICENSE is in incl_licenses directory.
|
||||
def sinc(x: torch.Tensor):
|
||||
"""
|
||||
Implementation of sinc, i.e. sin(pi * x) / (pi * x)
|
||||
__Warning__: Different to julius.sinc, the input is multiplied by `pi`!
|
||||
"""
|
||||
return torch.where(x == 0,
|
||||
torch.tensor(1., device=x.device, dtype=x.dtype),
|
||||
torch.sin(math.pi * x) / math.pi / x)
|
||||
|
||||
|
||||
# This code is adopted from adefossez's julius.lowpass.LowPassFilters under the MIT License
|
||||
# https://adefossez.github.io/julius/julius/lowpass.html
|
||||
# LICENSE is in incl_licenses directory.
|
||||
def kaiser_sinc_filter1d(cutoff, half_width, kernel_size): # return filter [1,1,kernel_size]
|
||||
even = (kernel_size % 2 == 0)
|
||||
half_size = kernel_size // 2
|
||||
|
||||
#For kaiser window
|
||||
delta_f = 4 * half_width
|
||||
A = 2.285 * (half_size - 1) * math.pi * delta_f + 7.95
|
||||
if A > 50.:
|
||||
beta = 0.1102 * (A - 8.7)
|
||||
elif A >= 21.:
|
||||
beta = 0.5842 * (A - 21)**0.4 + 0.07886 * (A - 21.)
|
||||
else:
|
||||
beta = 0.
|
||||
window = torch.kaiser_window(kernel_size, beta=beta, periodic=False)
|
||||
|
||||
# ratio = 0.5/cutoff -> 2 * cutoff = 1 / ratio
|
||||
if even:
|
||||
time = (torch.arange(-half_size, half_size) + 0.5)
|
||||
else:
|
||||
time = torch.arange(kernel_size) - half_size
|
||||
if cutoff == 0:
|
||||
filter_ = torch.zeros_like(time)
|
||||
else:
|
||||
filter_ = 2 * cutoff * window * sinc(2 * cutoff * time)
|
||||
# Normalize filter to have sum = 1, otherwise we will have a small leakage
|
||||
# of the constant component in the input signal.
|
||||
filter_ /= filter_.sum()
|
||||
filter = filter_.view(1, 1, kernel_size)
|
||||
|
||||
return filter
|
||||
|
||||
|
||||
class LowPassFilter1d(nn.Module):
|
||||
def __init__(self,
|
||||
cutoff=0.5,
|
||||
half_width=0.6,
|
||||
stride: int = 1,
|
||||
padding: bool = True,
|
||||
padding_mode: str = 'replicate',
|
||||
kernel_size: int = 12):
|
||||
# kernel_size should be even number for stylegan3 setup,
|
||||
# in this implementation, odd number is also possible.
|
||||
super().__init__()
|
||||
if cutoff < -0.:
|
||||
raise ValueError("Minimum cutoff must be larger than zero.")
|
||||
if cutoff > 0.5:
|
||||
raise ValueError("A cutoff above 0.5 does not make sense.")
|
||||
self.kernel_size = kernel_size
|
||||
self.even = (kernel_size % 2 == 0)
|
||||
self.pad_left = kernel_size // 2 - int(self.even)
|
||||
self.pad_right = kernel_size // 2
|
||||
self.stride = stride
|
||||
self.padding = padding
|
||||
self.padding_mode = padding_mode
|
||||
filter = kaiser_sinc_filter1d(cutoff, half_width, kernel_size)
|
||||
self.register_buffer("filter", filter)
|
||||
|
||||
#input [B, C, T]
|
||||
def forward(self, x):
|
||||
_, C, _ = x.shape
|
||||
|
||||
if self.padding:
|
||||
x = F.pad(x, (self.pad_left, self.pad_right),
|
||||
mode=self.padding_mode)
|
||||
out = F.conv1d(x, self.filter.expand(C, -1, -1),
|
||||
stride=self.stride, groups=C)
|
||||
|
||||
return out
|
||||
@@ -0,0 +1,49 @@
|
||||
# Adapted from https://github.com/junjun3518/alias-free-torch under the Apache License 2.0
|
||||
# LICENSE is in incl_licenses directory.
|
||||
|
||||
import torch.nn as nn
|
||||
from torch.nn import functional as F
|
||||
from .filter import LowPassFilter1d
|
||||
from .filter import kaiser_sinc_filter1d
|
||||
|
||||
|
||||
class UpSample1d(nn.Module):
|
||||
def __init__(self, ratio=2, kernel_size=None):
|
||||
super().__init__()
|
||||
self.ratio = ratio
|
||||
self.kernel_size = int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
|
||||
self.stride = ratio
|
||||
self.pad = self.kernel_size // ratio - 1
|
||||
self.pad_left = self.pad * self.stride + (self.kernel_size - self.stride) // 2
|
||||
self.pad_right = self.pad * self.stride + (self.kernel_size - self.stride + 1) // 2
|
||||
filter = kaiser_sinc_filter1d(cutoff=0.5 / ratio,
|
||||
half_width=0.6 / ratio,
|
||||
kernel_size=self.kernel_size)
|
||||
self.register_buffer("filter", filter)
|
||||
|
||||
# x: [B, C, T]
|
||||
def forward(self, x):
|
||||
_, C, _ = x.shape
|
||||
|
||||
x = F.pad(x, (self.pad, self.pad), mode='replicate')
|
||||
x = self.ratio * F.conv_transpose1d(
|
||||
x, self.filter.expand(C, -1, -1), stride=self.stride, groups=C)
|
||||
x = x[..., self.pad_left:-self.pad_right]
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class DownSample1d(nn.Module):
|
||||
def __init__(self, ratio=2, kernel_size=None):
|
||||
super().__init__()
|
||||
self.ratio = ratio
|
||||
self.kernel_size = int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
|
||||
self.lowpass = LowPassFilter1d(cutoff=0.5 / ratio,
|
||||
half_width=0.6 / ratio,
|
||||
stride=ratio,
|
||||
kernel_size=self.kernel_size)
|
||||
|
||||
def forward(self, x):
|
||||
xx = self.lowpass(x)
|
||||
|
||||
return xx
|
||||
@@ -0,0 +1,62 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from types import SimpleNamespace
|
||||
|
||||
from .models import BigVGANVocoder
|
||||
|
||||
from comfy.utils import load_torch_file
|
||||
|
||||
# BigVGAN vocoder configuration
|
||||
_bigvgan_vocoder_config = {
|
||||
'resblock': '1',
|
||||
'num_gpus': 0,
|
||||
'batch_size': 64,
|
||||
'num_mels': 80,
|
||||
'learning_rate': 0.0001,
|
||||
'adam_b1': 0.8,
|
||||
'adam_b2': 0.99,
|
||||
'lr_decay': 0.999,
|
||||
'seed': 1234,
|
||||
'upsample_rates': [4, 4, 2, 2, 2, 2],
|
||||
'upsample_kernel_sizes': [8, 8, 4, 4, 4, 4],
|
||||
'upsample_initial_channel': 1536,
|
||||
'resblock_kernel_sizes': [3, 7, 11],
|
||||
'resblock_dilation_sizes': [
|
||||
[1, 3, 5],
|
||||
[1, 3, 5],
|
||||
[1, 3, 5]
|
||||
],
|
||||
'activation': 'snakebeta',
|
||||
'snake_logscale': True,
|
||||
'resolutions': [
|
||||
[1024, 120, 600],
|
||||
[2048, 240, 1200],
|
||||
[512, 50, 240]
|
||||
],
|
||||
'mpd_reshapes': [2, 3, 5, 7, 11],
|
||||
'use_spectral_norm': False,
|
||||
'discriminator_channel_mult': 1,
|
||||
}
|
||||
|
||||
class BigVGAN(nn.Module):
|
||||
|
||||
def __init__(self, ckpt_path):
|
||||
super().__init__()
|
||||
# Convert dictionary to namespace object for attribute access
|
||||
vocoder_cfg = SimpleNamespace(**_bigvgan_vocoder_config)
|
||||
self.vocoder = BigVGANVocoder(vocoder_cfg).eval()
|
||||
vocoder_ckpt = load_torch_file(ckpt_path)
|
||||
self.vocoder.load_state_dict(vocoder_ckpt)
|
||||
|
||||
self.weight_norm_removed = False
|
||||
self.remove_weight_norm()
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(self, x):
|
||||
assert self.weight_norm_removed, 'call remove_weight_norm() before inference'
|
||||
return self.vocoder(x)
|
||||
|
||||
def remove_weight_norm(self):
|
||||
self.vocoder.remove_weight_norm()
|
||||
self.weight_norm_removed = True
|
||||
return self
|
||||
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2020 Jungil Kong
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2020 Edward Dixon
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
@@ -0,0 +1,201 @@
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
|
||||
Copyright [yyyy] [name of copyright owner]
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
@@ -0,0 +1,29 @@
|
||||
BSD 3-Clause License
|
||||
|
||||
Copyright (c) 2019, Seungwon Park 박승원
|
||||
All rights reserved.
|
||||
|
||||
Redistribution and use in source and binary forms, with or without
|
||||
modification, are permitted provided that the following conditions are met:
|
||||
|
||||
1. Redistributions of source code must retain the above copyright notice, this
|
||||
list of conditions and the following disclaimer.
|
||||
|
||||
2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
this list of conditions and the following disclaimer in the documentation
|
||||
and/or other materials provided with the distribution.
|
||||
|
||||
3. Neither the name of the copyright holder nor the names of its
|
||||
contributors may be used to endorse or promote products derived from
|
||||
this software without specific prior written permission.
|
||||
|
||||
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
@@ -0,0 +1,16 @@
|
||||
Copyright 2020 Alexandre Défossez
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy of this software and
|
||||
associated documentation files (the "Software"), to deal in the Software without restriction,
|
||||
including without limitation the rights to use, copy, modify, merge, publish, distribute,
|
||||
sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all copies or
|
||||
substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT
|
||||
NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND
|
||||
NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM,
|
||||
DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||
@@ -0,0 +1,255 @@
|
||||
# Copyright (c) 2022 NVIDIA CORPORATION.
|
||||
# Licensed under the MIT license.
|
||||
|
||||
# Adapted from https://github.com/jik876/hifi-gan under the MIT license.
|
||||
# LICENSE is in incl_licenses directory.
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn import Conv1d, ConvTranspose1d
|
||||
from torch.nn.utils.parametrizations import weight_norm
|
||||
from torch.nn.utils.parametrize import remove_parametrizations
|
||||
|
||||
from . import activations
|
||||
from .alias_free_torch import *
|
||||
from .utils import get_padding, init_weights
|
||||
|
||||
LRELU_SLOPE = 0.1
|
||||
|
||||
|
||||
class AMPBlock1(torch.nn.Module):
|
||||
|
||||
def __init__(self, h, channels, kernel_size=3, dilation=(1, 3, 5), activation=None):
|
||||
super(AMPBlock1, self).__init__()
|
||||
self.h = h
|
||||
|
||||
self.convs1 = nn.ModuleList([
|
||||
weight_norm(
|
||||
Conv1d(channels,
|
||||
channels,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=dilation[0],
|
||||
padding=get_padding(kernel_size, dilation[0]))),
|
||||
weight_norm(
|
||||
Conv1d(channels,
|
||||
channels,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=dilation[1],
|
||||
padding=get_padding(kernel_size, dilation[1]))),
|
||||
weight_norm(
|
||||
Conv1d(channels,
|
||||
channels,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=dilation[2],
|
||||
padding=get_padding(kernel_size, dilation[2])))
|
||||
])
|
||||
self.convs1.apply(init_weights)
|
||||
|
||||
self.convs2 = nn.ModuleList([
|
||||
weight_norm(
|
||||
Conv1d(channels,
|
||||
channels,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=1,
|
||||
padding=get_padding(kernel_size, 1))),
|
||||
weight_norm(
|
||||
Conv1d(channels,
|
||||
channels,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=1,
|
||||
padding=get_padding(kernel_size, 1))),
|
||||
weight_norm(
|
||||
Conv1d(channels,
|
||||
channels,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=1,
|
||||
padding=get_padding(kernel_size, 1)))
|
||||
])
|
||||
self.convs2.apply(init_weights)
|
||||
|
||||
self.num_layers = len(self.convs1) + len(self.convs2) # total number of conv layers
|
||||
|
||||
if activation == 'snake': # periodic nonlinearity with snake function and anti-aliasing
|
||||
self.activations = nn.ModuleList([
|
||||
Activation1d(
|
||||
activation=activations.Snake(channels, alpha_logscale=h.snake_logscale))
|
||||
for _ in range(self.num_layers)
|
||||
])
|
||||
elif activation == 'snakebeta': # periodic nonlinearity with snakebeta function and anti-aliasing
|
||||
self.activations = nn.ModuleList([
|
||||
Activation1d(
|
||||
activation=activations.SnakeBeta(channels, alpha_logscale=h.snake_logscale))
|
||||
for _ in range(self.num_layers)
|
||||
])
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"activation incorrectly specified. check the config file and look for 'activation'."
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
acts1, acts2 = self.activations[::2], self.activations[1::2]
|
||||
for c1, c2, a1, a2 in zip(self.convs1, self.convs2, acts1, acts2):
|
||||
xt = a1(x)
|
||||
xt = c1(xt)
|
||||
xt = a2(xt)
|
||||
xt = c2(xt)
|
||||
x = xt + x
|
||||
|
||||
return x
|
||||
|
||||
def remove_weight_norm(self):
|
||||
for l in self.convs1:
|
||||
remove_parametrizations(l, 'weight')
|
||||
for l in self.convs2:
|
||||
remove_parametrizations(l, 'weight')
|
||||
|
||||
|
||||
class AMPBlock2(torch.nn.Module):
|
||||
|
||||
def __init__(self, h, channels, kernel_size=3, dilation=(1, 3), activation=None):
|
||||
super(AMPBlock2, self).__init__()
|
||||
self.h = h
|
||||
|
||||
self.convs = nn.ModuleList([
|
||||
weight_norm(
|
||||
Conv1d(channels,
|
||||
channels,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=dilation[0],
|
||||
padding=get_padding(kernel_size, dilation[0]))),
|
||||
weight_norm(
|
||||
Conv1d(channels,
|
||||
channels,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=dilation[1],
|
||||
padding=get_padding(kernel_size, dilation[1])))
|
||||
])
|
||||
self.convs.apply(init_weights)
|
||||
|
||||
self.num_layers = len(self.convs) # total number of conv layers
|
||||
|
||||
if activation == 'snake': # periodic nonlinearity with snake function and anti-aliasing
|
||||
self.activations = nn.ModuleList([
|
||||
Activation1d(
|
||||
activation=activations.Snake(channels, alpha_logscale=h.snake_logscale))
|
||||
for _ in range(self.num_layers)
|
||||
])
|
||||
elif activation == 'snakebeta': # periodic nonlinearity with snakebeta function and anti-aliasing
|
||||
self.activations = nn.ModuleList([
|
||||
Activation1d(
|
||||
activation=activations.SnakeBeta(channels, alpha_logscale=h.snake_logscale))
|
||||
for _ in range(self.num_layers)
|
||||
])
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"activation incorrectly specified. check the config file and look for 'activation'."
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
for c, a in zip(self.convs, self.activations):
|
||||
xt = a(x)
|
||||
xt = c(xt)
|
||||
x = xt + x
|
||||
|
||||
return x
|
||||
|
||||
def remove_weight_norm(self):
|
||||
for l in self.convs:
|
||||
remove_parametrizations(l, 'weight')
|
||||
|
||||
|
||||
class BigVGANVocoder(torch.nn.Module):
|
||||
# this is our main BigVGAN model. Applies anti-aliased periodic activation for resblocks.
|
||||
def __init__(self, h):
|
||||
super().__init__()
|
||||
self.h = h
|
||||
|
||||
self.num_kernels = len(h.resblock_kernel_sizes)
|
||||
self.num_upsamples = len(h.upsample_rates)
|
||||
|
||||
# pre conv
|
||||
self.conv_pre = weight_norm(Conv1d(h.num_mels, h.upsample_initial_channel, 7, 1, padding=3))
|
||||
|
||||
# define which AMPBlock to use. BigVGAN uses AMPBlock1 as default
|
||||
resblock = AMPBlock1 if h.resblock == '1' else AMPBlock2
|
||||
|
||||
# transposed conv-based upsamplers. does not apply anti-aliasing
|
||||
self.ups = nn.ModuleList()
|
||||
for i, (u, k) in enumerate(zip(h.upsample_rates, h.upsample_kernel_sizes)):
|
||||
self.ups.append(
|
||||
nn.ModuleList([
|
||||
weight_norm(
|
||||
ConvTranspose1d(h.upsample_initial_channel // (2**i),
|
||||
h.upsample_initial_channel // (2**(i + 1)),
|
||||
k,
|
||||
u,
|
||||
padding=(k - u) // 2))
|
||||
]))
|
||||
|
||||
# residual blocks using anti-aliased multi-periodicity composition modules (AMP)
|
||||
self.resblocks = nn.ModuleList()
|
||||
for i in range(len(self.ups)):
|
||||
ch = h.upsample_initial_channel // (2**(i + 1))
|
||||
for j, (k, d) in enumerate(zip(h.resblock_kernel_sizes, h.resblock_dilation_sizes)):
|
||||
self.resblocks.append(resblock(h, ch, k, d, activation=h.activation))
|
||||
|
||||
# post conv
|
||||
if h.activation == "snake": # periodic nonlinearity with snake function and anti-aliasing
|
||||
activation_post = activations.Snake(ch, alpha_logscale=h.snake_logscale)
|
||||
self.activation_post = Activation1d(activation=activation_post)
|
||||
elif h.activation == "snakebeta": # periodic nonlinearity with snakebeta function and anti-aliasing
|
||||
activation_post = activations.SnakeBeta(ch, alpha_logscale=h.snake_logscale)
|
||||
self.activation_post = Activation1d(activation=activation_post)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"activation incorrectly specified. check the config file and look for 'activation'."
|
||||
)
|
||||
|
||||
self.conv_post = weight_norm(Conv1d(ch, 1, 7, 1, padding=3))
|
||||
|
||||
# weight initialization
|
||||
for i in range(len(self.ups)):
|
||||
self.ups[i].apply(init_weights)
|
||||
self.conv_post.apply(init_weights)
|
||||
|
||||
def forward(self, x):
|
||||
# pre conv
|
||||
x = self.conv_pre(x)
|
||||
|
||||
for i in range(self.num_upsamples):
|
||||
# upsampling
|
||||
for i_up in range(len(self.ups[i])):
|
||||
x = self.ups[i][i_up](x)
|
||||
# AMP blocks
|
||||
xs = None
|
||||
for j in range(self.num_kernels):
|
||||
if xs is None:
|
||||
xs = self.resblocks[i * self.num_kernels + j](x)
|
||||
else:
|
||||
xs += self.resblocks[i * self.num_kernels + j](x)
|
||||
x = xs / self.num_kernels
|
||||
|
||||
# post conv
|
||||
x = self.activation_post(x)
|
||||
x = self.conv_post(x)
|
||||
x = torch.tanh(x)
|
||||
|
||||
return x
|
||||
|
||||
def remove_weight_norm(self):
|
||||
print('Removing weight norm...')
|
||||
for l in self.ups:
|
||||
for l_i in l:
|
||||
remove_parametrizations(l_i, 'weight')
|
||||
for l in self.resblocks:
|
||||
l.remove_weight_norm()
|
||||
remove_parametrizations(self.conv_pre, 'weight')
|
||||
remove_parametrizations(self.conv_post, 'weight')
|
||||
@@ -0,0 +1,20 @@
|
||||
# Adapted from https://github.com/jik876/hifi-gan under the MIT license.
|
||||
# LICENSE is in incl_licenses directory.
|
||||
|
||||
from torch.nn.utils.parametrizations import weight_norm
|
||||
|
||||
|
||||
def init_weights(m, mean=0.0, std=0.01):
|
||||
classname = m.__class__.__name__
|
||||
if classname.find("Conv") != -1:
|
||||
m.weight.data.normal_(mean, std)
|
||||
|
||||
|
||||
def apply_weight_norm(m):
|
||||
classname = m.__class__.__name__
|
||||
if classname.find("Conv") != -1:
|
||||
weight_norm(m)
|
||||
|
||||
|
||||
def get_padding(kernel_size, dilation=1):
|
||||
return int((kernel_size * dilation - dilation) / 2)
|
||||
@@ -0,0 +1,212 @@
|
||||
# Reference: # https://github.com/bytedance/Make-An-Audio-2
|
||||
from typing import Literal
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import numpy as np
|
||||
|
||||
# following is from librosa
|
||||
|
||||
def hz_to_mel(frequencies, *, htk = False):
|
||||
frequencies = np.asanyarray(frequencies)
|
||||
|
||||
if htk:
|
||||
mels: np.ndarray = 2595.0 * np.log10(1.0 + frequencies / 700.0)
|
||||
return mels
|
||||
|
||||
# Fill in the linear part
|
||||
f_min = 0.0
|
||||
f_sp = 200.0 / 3
|
||||
|
||||
mels = (frequencies - f_min) / f_sp
|
||||
|
||||
# Fill in the log-scale part
|
||||
|
||||
min_log_hz = 1000.0 # beginning of log region (Hz)
|
||||
min_log_mel = (min_log_hz - f_min) / f_sp # same (Mels)
|
||||
logstep = np.log(6.4) / 27.0 # step size for log region
|
||||
|
||||
if frequencies.ndim:
|
||||
# If we have array data, vectorize
|
||||
log_t = frequencies >= min_log_hz
|
||||
mels[log_t] = min_log_mel + np.log(frequencies[log_t] / min_log_hz) / logstep
|
||||
elif frequencies >= min_log_hz:
|
||||
# If we have scalar data, heck directly
|
||||
mels = min_log_mel + np.log(frequencies / min_log_hz) / logstep
|
||||
|
||||
return mels
|
||||
|
||||
def mel_to_hz(mels, *, htk = False):
|
||||
mels = np.asanyarray(mels)
|
||||
|
||||
if htk:
|
||||
return 700.0 * (10.0 ** (mels / 2595.0) - 1.0)
|
||||
|
||||
# Fill in the linear scale
|
||||
f_min = 0.0
|
||||
f_sp = 200.0 / 3
|
||||
freqs = f_min + f_sp * mels
|
||||
|
||||
# And now the nonlinear scale
|
||||
min_log_hz = 1000.0 # beginning of log region (Hz)
|
||||
min_log_mel = (min_log_hz - f_min) / f_sp # same (Mels)
|
||||
logstep = np.log(6.4) / 27.0 # step size for log region
|
||||
|
||||
if mels.ndim:
|
||||
# If we have vector data, vectorize
|
||||
log_t = mels >= min_log_mel
|
||||
freqs[log_t] = min_log_hz * np.exp(logstep * (mels[log_t] - min_log_mel))
|
||||
elif mels >= min_log_mel:
|
||||
# If we have scalar data, check directly
|
||||
freqs = min_log_hz * np.exp(logstep * (mels - min_log_mel))
|
||||
|
||||
return freqs
|
||||
|
||||
def mel_frequencies(n_mels = 128, *, fmin = 0.0, fmax = 11025.0, htk = False):
|
||||
min_mel = hz_to_mel(fmin, htk=htk)
|
||||
max_mel = hz_to_mel(fmax, htk=htk)
|
||||
mels = np.linspace(min_mel, max_mel, n_mels)
|
||||
hz: np.ndarray = mel_to_hz(mels, htk=htk)
|
||||
return hz
|
||||
|
||||
def librosa_mel_fn(
|
||||
*,
|
||||
sr: float,
|
||||
n_fft: int,
|
||||
n_mels: int = 128,
|
||||
fmin: float = 0.0,
|
||||
fmax = None,
|
||||
htk = False,
|
||||
norm = "slaney",
|
||||
dtype = np.float32,
|
||||
) -> np.ndarray:
|
||||
|
||||
if fmax is None:
|
||||
fmax = float(sr) / 2
|
||||
|
||||
# Initialize the weights
|
||||
n_mels = int(n_mels)
|
||||
weights = np.zeros((n_mels, int(1 + n_fft // 2)), dtype=dtype)
|
||||
|
||||
# Center freqs of each FFT bin
|
||||
fftfreqs = np.fft.rfftfreq(n=n_fft, d=1.0 / sr)
|
||||
|
||||
# 'Center freqs' of mel bands - uniformly spaced between limits
|
||||
mel_f = mel_frequencies(n_mels + 2, fmin=fmin, fmax=fmax, htk=htk)
|
||||
|
||||
fdiff = np.diff(mel_f)
|
||||
ramps = np.subtract.outer(mel_f, fftfreqs)
|
||||
|
||||
for i in range(n_mels):
|
||||
# lower and upper slopes for all bins
|
||||
lower = -ramps[i] / fdiff[i]
|
||||
upper = ramps[i + 2] / fdiff[i + 1]
|
||||
|
||||
# .. then intersect them with each other and zero
|
||||
weights[i] = np.maximum(0, np.minimum(lower, upper))
|
||||
|
||||
# Slaney-style mel is scaled to be approx constant energy per channel
|
||||
enorm = 2.0 / (mel_f[2 : n_mels + 2] - mel_f[:n_mels])
|
||||
weights *= enorm[:, np.newaxis]
|
||||
|
||||
return weights
|
||||
|
||||
|
||||
def dynamic_range_compression_torch(x, C=1, clip_val=1e-5, *, norm_fn):
|
||||
return norm_fn(torch.clamp(x, min=clip_val) * C)
|
||||
|
||||
|
||||
def spectral_normalize_torch(magnitudes, norm_fn):
|
||||
output = dynamic_range_compression_torch(magnitudes, norm_fn=norm_fn)
|
||||
return output
|
||||
|
||||
|
||||
class MelConverter(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
sampling_rate: float,
|
||||
n_fft: int,
|
||||
num_mels: int,
|
||||
hop_size: int,
|
||||
win_size: int,
|
||||
fmin: float,
|
||||
fmax: float,
|
||||
norm_fn,
|
||||
):
|
||||
super().__init__()
|
||||
self.sampling_rate = sampling_rate
|
||||
self.n_fft = n_fft
|
||||
self.num_mels = num_mels
|
||||
self.hop_size = hop_size
|
||||
self.win_size = win_size
|
||||
self.fmin = fmin
|
||||
self.fmax = fmax
|
||||
self.norm_fn = norm_fn
|
||||
|
||||
mel = librosa_mel_fn(sr=self.sampling_rate,
|
||||
n_fft=self.n_fft,
|
||||
n_mels=self.num_mels,
|
||||
fmin=self.fmin,
|
||||
fmax=self.fmax)
|
||||
mel_basis = torch.from_numpy(mel).float()
|
||||
hann_window = torch.hann_window(self.win_size)
|
||||
|
||||
self.register_buffer('mel_basis', mel_basis)
|
||||
self.register_buffer('hann_window', hann_window)
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return self.mel_basis.device
|
||||
|
||||
def forward(self, waveform: torch.Tensor, center: bool = False) -> torch.Tensor:
|
||||
waveform = waveform.clamp(min=-1., max=1.).to(self.device)
|
||||
|
||||
waveform = torch.nn.functional.pad(
|
||||
waveform.unsqueeze(1),
|
||||
[int((self.n_fft - self.hop_size) / 2),
|
||||
int((self.n_fft - self.hop_size) / 2)],
|
||||
mode='reflect')
|
||||
waveform = waveform.squeeze(1)
|
||||
|
||||
spec = torch.stft(waveform,
|
||||
self.n_fft,
|
||||
hop_length=self.hop_size,
|
||||
win_length=self.win_size,
|
||||
window=self.hann_window,
|
||||
center=center,
|
||||
pad_mode='reflect',
|
||||
normalized=False,
|
||||
onesided=True,
|
||||
return_complex=True)
|
||||
|
||||
spec = torch.view_as_real(spec)
|
||||
spec = torch.sqrt(spec.pow(2).sum(-1) + (1e-9)).float()
|
||||
spec = torch.matmul(self.mel_basis, spec)
|
||||
spec = spectral_normalize_torch(spec, self.norm_fn)
|
||||
|
||||
return spec
|
||||
|
||||
|
||||
def get_mel_converter(mode: Literal['16k', '44k']) -> MelConverter:
|
||||
if mode == '16k':
|
||||
return MelConverter(sampling_rate=16_000,
|
||||
n_fft=1024,
|
||||
num_mels=80,
|
||||
hop_size=256,
|
||||
win_size=1024,
|
||||
fmin=0,
|
||||
fmax=8_000,
|
||||
norm_fn=torch.log10)
|
||||
elif mode == '44k':
|
||||
return MelConverter(sampling_rate=44_100,
|
||||
n_fft=2048,
|
||||
num_mels=128,
|
||||
hop_size=512,
|
||||
win_size=2048,
|
||||
fmin=0,
|
||||
fmax=44100 / 2,
|
||||
norm_fn=torch.log)
|
||||
else:
|
||||
raise ValueError(f'Unknown mode: {mode}')
|
||||
@@ -0,0 +1,257 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import folder_paths
|
||||
import os
|
||||
|
||||
from .mel_converter import get_mel_converter
|
||||
from .vae.autoencoder import AutoEncoderModule
|
||||
from .vae.distributions import DiagonalGaussianDistribution
|
||||
import torchaudio
|
||||
|
||||
from comfy import model_management as mm
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
class FeaturesUtils(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
tod_vae_ckpt: str,
|
||||
bigvgan_vocoder_ckpt = None,
|
||||
mode=['16k', '44k'],
|
||||
need_vae_encoder: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.mel_converter = get_mel_converter(mode)
|
||||
self.tod = AutoEncoderModule(vae_ckpt_path=tod_vae_ckpt,
|
||||
vocoder_ckpt_path=bigvgan_vocoder_ckpt,
|
||||
mode=mode,
|
||||
need_vae_encoder=need_vae_encoder)
|
||||
|
||||
def encode_audio(self, x) -> DiagonalGaussianDistribution:
|
||||
assert self.tod is not None, 'VAE is not loaded'
|
||||
# x: (B * L)
|
||||
mel = self.mel_converter(x)
|
||||
dist = self.tod.encode(mel)
|
||||
|
||||
return dist
|
||||
|
||||
def vocode(self, mel: torch.Tensor) -> torch.Tensor:
|
||||
assert self.tod is not None, 'VAE is not loaded'
|
||||
return self.tod.vocode(mel)
|
||||
|
||||
def decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
assert self.tod is not None, 'VAE is not loaded'
|
||||
return self.tod.decode(z)
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
def wrapped_decode(self, z):
|
||||
with torch.amp.autocast('cuda', dtype=self.dtype):
|
||||
mel_decoded = self.decode(z)
|
||||
audio = self.vocode(mel_decoded)
|
||||
|
||||
return audio
|
||||
|
||||
def wrapped_encode(self, audio):
|
||||
with torch.amp.autocast('cuda', dtype=self.dtype):
|
||||
dist = self.encode_audio(audio)
|
||||
|
||||
return dist.mean
|
||||
|
||||
if not "mmaudio" in folder_paths.folder_names_and_paths:
|
||||
folder_paths.add_model_folder_path("mmaudio", os.path.join(folder_paths.models_dir, "mmaudio"))
|
||||
|
||||
class OviMMAudioVAELoader:
|
||||
"""Loads MMAudio VAE for audio encoding/decoding in Ovi"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
s.vae_files = folder_paths.get_filename_list("vae")
|
||||
s.mmaudio_files = folder_paths.get_filename_list("mmaudio")
|
||||
s.all_files = s.vae_files + s.mmaudio_files
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"vae": (s.all_files, {"tooltip": "MMAudio VAE 16k (v1-16.pth) model from models/vae or models/mmaudio"}),
|
||||
"vocoder": (s.all_files, {"tooltip": "BigVGAN vocoder (best_netG.pt) from models/vae or models/mmaudio"}),
|
||||
"precision": (["bf16", "fp16", "fp32"], {"default": "bf16"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MMAUDIOVAE",)
|
||||
RETURN_NAMES = ("mmaudio_vae",)
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "WanVideoWrapper/Ovi"
|
||||
DESCRIPTION = "Loads MMAudio VAE for Ovi audio generation"
|
||||
|
||||
def loadmodel(self, vae, vocoder, precision):
|
||||
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||
|
||||
vae_path = folder_paths.get_full_path("vae", vae) if vae in self.vae_files else folder_paths.get_full_path("mmaudio", vae)
|
||||
vocoder_path = folder_paths.get_full_path("vae", vocoder) if vocoder in self.vae_files else folder_paths.get_full_path("mmaudio", vocoder)
|
||||
|
||||
vae = FeaturesUtils(
|
||||
tod_vae_ckpt=vae_path,
|
||||
bigvgan_vocoder_ckpt=vocoder_path,
|
||||
mode='16k',
|
||||
need_vae_encoder=True
|
||||
)
|
||||
|
||||
vae.to(device=offload_device, dtype=dtype)
|
||||
vae.eval()
|
||||
|
||||
return (vae,)
|
||||
|
||||
class WanVideoDecodeOviAudio:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"mmaudio_vae": ("MMAUDIOVAE",),
|
||||
"samples": ("LATENT",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("AUDIO",)
|
||||
RETURN_NAMES = ("audio",)
|
||||
FUNCTION = "decode"
|
||||
CATEGORY = "WanVideoWrapper/Ovi"
|
||||
|
||||
def decode(self, mmaudio_vae, samples):
|
||||
mm.soft_empty_cache()
|
||||
audio_latents = samples.get("latent_ovi_audio", None)
|
||||
if audio_latents is None:
|
||||
raise ValueError("No Ovi audio latents found in input samples")
|
||||
|
||||
mmaudio_vae.to(device)
|
||||
|
||||
waveform = mmaudio_vae.wrapped_decode(audio_latents.to(device=device, dtype=mmaudio_vae.dtype))
|
||||
audio = {"waveform": waveform.unsqueeze(0).cpu().float(), "sample_rate": 16000}
|
||||
|
||||
mmaudio_vae.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
|
||||
return (audio,)
|
||||
|
||||
class WanVideoEncodeOviAudio:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"mmaudio_vae": ("MMAUDIOVAE",),
|
||||
"audio": ("AUDIO",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
RETURN_NAMES = ("samples",)
|
||||
FUNCTION = "decode"
|
||||
CATEGORY = "WanVideoWrapper/Ovi"
|
||||
|
||||
def decode(self, mmaudio_vae, audio):
|
||||
|
||||
mmaudio_vae.to(device)
|
||||
|
||||
waveform = audio.get("waveform", None)
|
||||
sample_rate = audio.get("sample_rate", None)
|
||||
if sample_rate != 16000:
|
||||
waveform = torchaudio.functional.resample(waveform, sample_rate, 16000)
|
||||
waveform = waveform.to(device=device, dtype=mmaudio_vae.dtype)[0][0].unsqueeze(0)
|
||||
|
||||
samples = mmaudio_vae.wrapped_encode(waveform)
|
||||
|
||||
mmaudio_vae.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
|
||||
return ({"latent_ovi_audio": samples},)
|
||||
|
||||
|
||||
class WanVideoAddOviAudioToLatents:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"original_samples": ("LATENT",),
|
||||
"audio_samples": ("LATENT",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
RETURN_NAMES = ("samples",)
|
||||
FUNCTION = "decode"
|
||||
CATEGORY = "WanVideoWrapper/Ovi"
|
||||
|
||||
def decode(self, original_samples, audio_samples):
|
||||
samples = original_samples.copy()
|
||||
samples.update(audio_samples)
|
||||
|
||||
return (samples,)
|
||||
|
||||
class WanVideoEmptyMMAudioLatents:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"length": ("INT", {"default": 157, "min": 1, "max": 10000, "step": 1, "tooltip": "Length of the audio latent sequence"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
RETURN_NAMES = ("samples",)
|
||||
FUNCTION = "decode"
|
||||
CATEGORY = "WanVideoWrapper/Ovi"
|
||||
|
||||
def decode(self, length):
|
||||
audio_latents = torch.zeros((length, 20), device=torch.device("cpu"), dtype=torch.float32) # 1, l c -> l, c
|
||||
|
||||
return ({"latent_ovi_audio": audio_latents},)
|
||||
|
||||
|
||||
class WanVideoOviCFG:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"original_text_embeds": ("WANVIDEOTEXTEMBEDS",),
|
||||
"ovi_negative_text_embeds": ("WANVIDEOTEXTEMBEDS",),
|
||||
"ovi_audio_cfg": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 100.0, "step": 0.01}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
|
||||
RETURN_NAMES = ("text_embeds",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper/Ovi"
|
||||
DESCRIPTION = "Adds Ovi negative text embeddings and audio CFG scale to the text embeddings dictionary"
|
||||
|
||||
def process(self, original_text_embeds, ovi_negative_text_embeds, ovi_audio_cfg):
|
||||
negative_text_embeds = ovi_negative_text_embeds.get("negative_prompt_embeds", None)
|
||||
if negative_text_embeds is None:
|
||||
negative_text_embeds = original_text_embeds["prompt_embeds"]
|
||||
|
||||
prompt_embeds_dict_copy = original_text_embeds.copy()
|
||||
prompt_embeds_dict_copy.update({
|
||||
"ovi_negative_prompt_embeds": negative_text_embeds,
|
||||
"ovi_audio_cfg": ovi_audio_cfg,
|
||||
})
|
||||
return (prompt_embeds_dict_copy,)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"OviMMAudioVAELoader": OviMMAudioVAELoader,
|
||||
"WanVideoDecodeOviAudio": WanVideoDecodeOviAudio,
|
||||
"WanVideoEncodeOviAudio": WanVideoEncodeOviAudio,
|
||||
"WanVideoOviCFG": WanVideoOviCFG,
|
||||
"WanVideoAddOviAudioToLatents": WanVideoAddOviAudioToLatents,
|
||||
"WanVideoEmptyMMAudioLatents": WanVideoEmptyMMAudioLatents,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"OviMMAudioVAELoader": "Ovi MMAudio VAE Loader",
|
||||
"WanVideoDecodeOviAudio": "WanVideo Decode Ovi Audio",
|
||||
"WanVideoEncodeOviAudio": "WanVideo Encode Ovi Audio",
|
||||
"WanVideoOviCFG": "WanVideo Ovi CFG",
|
||||
"WanVideoAddOviAudioToLatents": "WanVideo Add MMAudio To Latents",
|
||||
"WanVideoEmptyMMAudioLatents": "WanVideo Empty MMAudio Latents",
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
from typing import Literal, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from .vae import VAE, get_my_vae
|
||||
from .distributions import DiagonalGaussianDistribution
|
||||
from ..bigvgan import BigVGAN
|
||||
|
||||
from comfy.utils import load_torch_file
|
||||
|
||||
class AutoEncoderModule(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
*,
|
||||
vae_ckpt_path,
|
||||
vocoder_ckpt_path: Optional[str] = None,
|
||||
mode: Literal['16k', '44k'],
|
||||
need_vae_encoder: bool = True):
|
||||
super().__init__()
|
||||
self.vae: VAE = get_my_vae(mode).eval()
|
||||
#vae_state_dict = torch.load(vae_ckpt_path, weights_only=True, map_location='cpu')'
|
||||
vae_state_dict = load_torch_file(vae_ckpt_path)
|
||||
self.vae.load_state_dict(vae_state_dict)
|
||||
self.vae.remove_weight_norm()
|
||||
|
||||
if mode == '16k':
|
||||
assert vocoder_ckpt_path is not None
|
||||
self.vocoder = BigVGAN(vocoder_ckpt_path).eval()
|
||||
elif mode == '44k':
|
||||
raise NotImplementedError("44k mode requires BigVGANv2 which is not currently supported in this environment.")
|
||||
self.vocoder = BigVGANv2.from_pretrained('nvidia/bigvgan_v2_44khz_128band_512x',
|
||||
use_cuda_kernel=False)
|
||||
self.vocoder.remove_weight_norm()
|
||||
else:
|
||||
raise ValueError(f'Unknown mode: {mode}')
|
||||
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
if not need_vae_encoder:
|
||||
del self.vae.encoder
|
||||
|
||||
@torch.inference_mode()
|
||||
def encode(self, x: torch.Tensor) -> DiagonalGaussianDistribution:
|
||||
return self.vae.encode(x)
|
||||
|
||||
@torch.inference_mode()
|
||||
def decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
return self.vae.decode(z)
|
||||
|
||||
@torch.inference_mode()
|
||||
def vocode(self, spec: torch.Tensor) -> torch.Tensor:
|
||||
return self.vocoder(spec)
|
||||
@@ -0,0 +1,45 @@
|
||||
from typing import Optional
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
|
||||
class DiagonalGaussianDistribution:
|
||||
|
||||
def __init__(self, parameters, deterministic=False):
|
||||
self.parameters = parameters
|
||||
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
|
||||
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
|
||||
self.deterministic = deterministic
|
||||
self.std = torch.exp(0.5 * self.logvar)
|
||||
self.var = torch.exp(self.logvar)
|
||||
if self.deterministic:
|
||||
self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device)
|
||||
|
||||
def sample(self, rng: Optional[torch.Generator] = None):
|
||||
# x = self.mean + self.std * torch.randn(self.mean.shape).to(device=self.parameters.device)
|
||||
|
||||
r = torch.empty_like(self.mean).normal_(generator=rng)
|
||||
x = self.mean + self.std * r
|
||||
|
||||
return x
|
||||
|
||||
def kl(self, other=None):
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.])
|
||||
else:
|
||||
if other is None:
|
||||
|
||||
return 0.5 * torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar
|
||||
else:
|
||||
return 0.5 * (torch.pow(self.mean - other.mean, 2) / other.var +
|
||||
self.var / other.var - 1.0 - self.logvar + other.logvar)
|
||||
|
||||
def nll(self, sample, dims=[1, 2, 3]):
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.])
|
||||
logtwopi = np.log(2.0 * np.pi)
|
||||
return 0.5 * torch.sum(logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
|
||||
dim=dims)
|
||||
|
||||
def mode(self):
|
||||
return self.mean
|
||||
@@ -0,0 +1,168 @@
|
||||
# Copyright (c) 2024, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
#
|
||||
# This work is licensed under a Creative Commons
|
||||
# Attribution-NonCommercial-ShareAlike 4.0 International License.
|
||||
# You should have received a copy of the license along with this
|
||||
# work. If not, see http://creativecommons.org/licenses/by-nc-sa/4.0/
|
||||
"""Improved diffusion model architecture proposed in the paper
|
||||
"Analyzing and Improving the Training Dynamics of Diffusion Models"."""
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
#----------------------------------------------------------------------------
|
||||
# Variant of constant() that inherits dtype and device from the given
|
||||
# reference tensor by default.
|
||||
|
||||
_constant_cache = dict()
|
||||
|
||||
|
||||
def constant(value, shape=None, dtype=None, device=None, memory_format=None):
|
||||
value = np.asarray(value)
|
||||
if shape is not None:
|
||||
shape = tuple(shape)
|
||||
if dtype is None:
|
||||
dtype = torch.get_default_dtype()
|
||||
if device is None:
|
||||
device = torch.device('cpu')
|
||||
if memory_format is None:
|
||||
memory_format = torch.contiguous_format
|
||||
|
||||
key = (value.shape, value.dtype, value.tobytes(), shape, dtype, device, memory_format)
|
||||
tensor = _constant_cache.get(key, None)
|
||||
if tensor is None:
|
||||
tensor = torch.as_tensor(value.copy(), dtype=dtype, device=device)
|
||||
if shape is not None:
|
||||
tensor, _ = torch.broadcast_tensors(tensor, torch.empty(shape))
|
||||
tensor = tensor.contiguous(memory_format=memory_format)
|
||||
_constant_cache[key] = tensor
|
||||
return tensor
|
||||
|
||||
|
||||
def const_like(ref, value, shape=None, dtype=None, device=None, memory_format=None):
|
||||
if dtype is None:
|
||||
dtype = ref.dtype
|
||||
if device is None:
|
||||
device = ref.device
|
||||
return constant(value, shape=shape, dtype=dtype, device=device, memory_format=memory_format)
|
||||
|
||||
|
||||
#----------------------------------------------------------------------------
|
||||
# Normalize given tensor to unit magnitude with respect to the given
|
||||
# dimensions. Default = all dimensions except the first.
|
||||
|
||||
|
||||
def normalize(x, dim=None, eps=1e-4):
|
||||
if dim is None:
|
||||
dim = list(range(1, x.ndim))
|
||||
norm = torch.linalg.vector_norm(x, dim=dim, keepdim=True, dtype=torch.float32)
|
||||
norm = torch.add(eps, norm, alpha=np.sqrt(norm.numel() / x.numel()))
|
||||
return x / norm.to(x.dtype)
|
||||
|
||||
|
||||
class Normalize(torch.nn.Module):
|
||||
|
||||
def __init__(self, dim=None, eps=1e-4):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.eps = eps
|
||||
|
||||
def forward(self, x):
|
||||
return normalize(x, dim=self.dim, eps=self.eps)
|
||||
|
||||
|
||||
#----------------------------------------------------------------------------
|
||||
# Upsample or downsample the given tensor with the given filter,
|
||||
# or keep it as is.
|
||||
|
||||
|
||||
def resample(x, f=[1, 1], mode='keep'):
|
||||
if mode == 'keep':
|
||||
return x
|
||||
f = np.float32(f)
|
||||
assert f.ndim == 1 and len(f) % 2 == 0
|
||||
pad = (len(f) - 1) // 2
|
||||
f = f / f.sum()
|
||||
f = np.outer(f, f)[np.newaxis, np.newaxis, :, :]
|
||||
f = const_like(x, f)
|
||||
c = x.shape[1]
|
||||
if mode == 'down':
|
||||
return torch.nn.functional.conv2d(x,
|
||||
f.tile([c, 1, 1, 1]),
|
||||
groups=c,
|
||||
stride=2,
|
||||
padding=(pad, ))
|
||||
assert mode == 'up'
|
||||
return torch.nn.functional.conv_transpose2d(x, (f * 4).tile([c, 1, 1, 1]),
|
||||
groups=c,
|
||||
stride=2,
|
||||
padding=(pad, ))
|
||||
|
||||
|
||||
#----------------------------------------------------------------------------
|
||||
# Magnitude-preserving SiLU (Equation 81).
|
||||
|
||||
|
||||
def mp_silu(x):
|
||||
return torch.nn.functional.silu(x) / 0.596
|
||||
|
||||
|
||||
class MPSiLU(torch.nn.Module):
|
||||
|
||||
def forward(self, x):
|
||||
return mp_silu(x)
|
||||
|
||||
|
||||
#----------------------------------------------------------------------------
|
||||
# Magnitude-preserving sum (Equation 88).
|
||||
|
||||
|
||||
def mp_sum(a, b, t=0.5):
|
||||
return a.lerp(b, t) / np.sqrt((1 - t)**2 + t**2)
|
||||
|
||||
|
||||
#----------------------------------------------------------------------------
|
||||
# Magnitude-preserving concatenation (Equation 103).
|
||||
|
||||
|
||||
def mp_cat(a, b, dim=1, t=0.5):
|
||||
Na = a.shape[dim]
|
||||
Nb = b.shape[dim]
|
||||
C = np.sqrt((Na + Nb) / ((1 - t)**2 + t**2))
|
||||
wa = C / np.sqrt(Na) * (1 - t)
|
||||
wb = C / np.sqrt(Nb) * t
|
||||
return torch.cat([wa * a, wb * b], dim=dim)
|
||||
|
||||
|
||||
#----------------------------------------------------------------------------
|
||||
# Magnitude-preserving convolution or fully-connected layer (Equation 47)
|
||||
# with force weight normalization (Equation 66).
|
||||
|
||||
|
||||
class MPConv1D(torch.nn.Module):
|
||||
|
||||
def __init__(self, in_channels, out_channels, kernel_size):
|
||||
super().__init__()
|
||||
self.out_channels = out_channels
|
||||
self.weight = torch.nn.Parameter(torch.randn(out_channels, in_channels, kernel_size))
|
||||
|
||||
self.weight_norm_removed = False
|
||||
|
||||
def forward(self, x, gain=1):
|
||||
assert self.weight_norm_removed, 'call remove_weight_norm() before inference'
|
||||
|
||||
w = self.weight * gain
|
||||
if w.ndim == 2:
|
||||
return x @ w.t()
|
||||
assert w.ndim == 3
|
||||
return torch.nn.functional.conv1d(x, w, padding=(w.shape[-1] // 2, ))
|
||||
|
||||
def remove_weight_norm(self):
|
||||
w = self.weight.to(torch.float32)
|
||||
w = normalize(w) # traditional weight normalization
|
||||
w = w / np.sqrt(w[0].numel())
|
||||
w = w.to(self.weight.dtype)
|
||||
self.weight.data.copy_(w)
|
||||
|
||||
self.weight_norm_removed = True
|
||||
return self
|
||||
+369
@@ -0,0 +1,369 @@
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from .edm2_utils import MPConv1D
|
||||
from .vae_modules import (AttnBlock1D, Downsample1D, ResnetBlock1D,
|
||||
Upsample1D, nonlinearity)
|
||||
from .distributions import DiagonalGaussianDistribution
|
||||
|
||||
log = logging.getLogger()
|
||||
|
||||
DATA_MEAN_80D = [
|
||||
-1.6058, -1.3676, -1.2520, -1.2453, -1.2078, -1.2224, -1.2419, -1.2439, -1.2922, -1.2927,
|
||||
-1.3170, -1.3543, -1.3401, -1.3836, -1.3907, -1.3912, -1.4313, -1.4152, -1.4527, -1.4728,
|
||||
-1.4568, -1.5101, -1.5051, -1.5172, -1.5623, -1.5373, -1.5746, -1.5687, -1.6032, -1.6131,
|
||||
-1.6081, -1.6331, -1.6489, -1.6489, -1.6700, -1.6738, -1.6953, -1.6969, -1.7048, -1.7280,
|
||||
-1.7361, -1.7495, -1.7658, -1.7814, -1.7889, -1.8064, -1.8221, -1.8377, -1.8417, -1.8643,
|
||||
-1.8857, -1.8929, -1.9173, -1.9379, -1.9531, -1.9673, -1.9824, -2.0042, -2.0215, -2.0436,
|
||||
-2.0766, -2.1064, -2.1418, -2.1855, -2.2319, -2.2767, -2.3161, -2.3572, -2.3954, -2.4282,
|
||||
-2.4659, -2.5072, -2.5552, -2.6074, -2.6584, -2.7107, -2.7634, -2.8266, -2.8981, -2.9673
|
||||
]
|
||||
|
||||
DATA_STD_80D = [
|
||||
1.0291, 1.0411, 1.0043, 0.9820, 0.9677, 0.9543, 0.9450, 0.9392, 0.9343, 0.9297, 0.9276, 0.9263,
|
||||
0.9242, 0.9254, 0.9232, 0.9281, 0.9263, 0.9315, 0.9274, 0.9247, 0.9277, 0.9199, 0.9188, 0.9194,
|
||||
0.9160, 0.9161, 0.9146, 0.9161, 0.9100, 0.9095, 0.9145, 0.9076, 0.9066, 0.9095, 0.9032, 0.9043,
|
||||
0.9038, 0.9011, 0.9019, 0.9010, 0.8984, 0.8983, 0.8986, 0.8961, 0.8962, 0.8978, 0.8962, 0.8973,
|
||||
0.8993, 0.8976, 0.8995, 0.9016, 0.8982, 0.8972, 0.8974, 0.8949, 0.8940, 0.8947, 0.8936, 0.8939,
|
||||
0.8951, 0.8956, 0.9017, 0.9167, 0.9436, 0.9690, 1.0003, 1.0225, 1.0381, 1.0491, 1.0545, 1.0604,
|
||||
1.0761, 1.0929, 1.1089, 1.1196, 1.1176, 1.1156, 1.1117, 1.1070
|
||||
]
|
||||
|
||||
DATA_MEAN_128D = [
|
||||
-3.3462, -2.6723, -2.4893, -2.3143, -2.2664, -2.3317, -2.1802, -2.4006, -2.2357, -2.4597,
|
||||
-2.3717, -2.4690, -2.5142, -2.4919, -2.6610, -2.5047, -2.7483, -2.5926, -2.7462, -2.7033,
|
||||
-2.7386, -2.8112, -2.7502, -2.9594, -2.7473, -3.0035, -2.8891, -2.9922, -2.9856, -3.0157,
|
||||
-3.1191, -2.9893, -3.1718, -3.0745, -3.1879, -3.2310, -3.1424, -3.2296, -3.2791, -3.2782,
|
||||
-3.2756, -3.3134, -3.3509, -3.3750, -3.3951, -3.3698, -3.4505, -3.4509, -3.5089, -3.4647,
|
||||
-3.5536, -3.5788, -3.5867, -3.6036, -3.6400, -3.6747, -3.7072, -3.7279, -3.7283, -3.7795,
|
||||
-3.8259, -3.8447, -3.8663, -3.9182, -3.9605, -3.9861, -4.0105, -4.0373, -4.0762, -4.1121,
|
||||
-4.1488, -4.1874, -4.2461, -4.3170, -4.3639, -4.4452, -4.5282, -4.6297, -4.7019, -4.7960,
|
||||
-4.8700, -4.9507, -5.0303, -5.0866, -5.1634, -5.2342, -5.3242, -5.4053, -5.4927, -5.5712,
|
||||
-5.6464, -5.7052, -5.7619, -5.8410, -5.9188, -6.0103, -6.0955, -6.1673, -6.2362, -6.3120,
|
||||
-6.3926, -6.4797, -6.5565, -6.6511, -6.8130, -6.9961, -7.1275, -7.2457, -7.3576, -7.4663,
|
||||
-7.6136, -7.7469, -7.8815, -8.0132, -8.1515, -8.3071, -8.4722, -8.7418, -9.3975, -9.6628,
|
||||
-9.7671, -9.8863, -9.9992, -10.0860, -10.1709, -10.5418, -11.2795, -11.3861
|
||||
]
|
||||
|
||||
DATA_STD_128D = [
|
||||
2.3804, 2.4368, 2.3772, 2.3145, 2.2803, 2.2510, 2.2316, 2.2083, 2.1996, 2.1835, 2.1769, 2.1659,
|
||||
2.1631, 2.1618, 2.1540, 2.1606, 2.1571, 2.1567, 2.1612, 2.1579, 2.1679, 2.1683, 2.1634, 2.1557,
|
||||
2.1668, 2.1518, 2.1415, 2.1449, 2.1406, 2.1350, 2.1313, 2.1415, 2.1281, 2.1352, 2.1219, 2.1182,
|
||||
2.1327, 2.1195, 2.1137, 2.1080, 2.1179, 2.1036, 2.1087, 2.1036, 2.1015, 2.1068, 2.0975, 2.0991,
|
||||
2.0902, 2.1015, 2.0857, 2.0920, 2.0893, 2.0897, 2.0910, 2.0881, 2.0925, 2.0873, 2.0960, 2.0900,
|
||||
2.0957, 2.0958, 2.0978, 2.0936, 2.0886, 2.0905, 2.0845, 2.0855, 2.0796, 2.0840, 2.0813, 2.0817,
|
||||
2.0838, 2.0840, 2.0917, 2.1061, 2.1431, 2.1976, 2.2482, 2.3055, 2.3700, 2.4088, 2.4372, 2.4609,
|
||||
2.4731, 2.4847, 2.5072, 2.5451, 2.5772, 2.6147, 2.6529, 2.6596, 2.6645, 2.6726, 2.6803, 2.6812,
|
||||
2.6899, 2.6916, 2.6931, 2.6998, 2.7062, 2.7262, 2.7222, 2.7158, 2.7041, 2.7485, 2.7491, 2.7451,
|
||||
2.7485, 2.7233, 2.7297, 2.7233, 2.7145, 2.6958, 2.6788, 2.6439, 2.6007, 2.4786, 2.2469, 2.1877,
|
||||
2.1392, 2.0717, 2.0107, 1.9676, 1.9140, 1.7102, 0.9101, 0.7164
|
||||
]
|
||||
|
||||
|
||||
class VAE(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
data_dim: int,
|
||||
embed_dim: int,
|
||||
hidden_dim: int,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
if data_dim == 80:
|
||||
self.data_mean = nn.Buffer(torch.tensor(DATA_MEAN_80D, dtype=torch.float32))
|
||||
self.data_std = nn.Buffer(torch.tensor(DATA_STD_80D, dtype=torch.float32))
|
||||
elif data_dim == 128:
|
||||
self.data_mean = nn.Buffer(torch.tensor(DATA_MEAN_128D, dtype=torch.float32))
|
||||
self.data_std = nn.Buffer(torch.tensor(DATA_STD_128D, dtype=torch.float32))
|
||||
|
||||
self.data_mean = self.data_mean.view(1, -1, 1)
|
||||
self.data_std = self.data_std.view(1, -1, 1)
|
||||
|
||||
self.encoder = Encoder1D(
|
||||
dim=hidden_dim,
|
||||
ch_mult=(1, 2, 4),
|
||||
num_res_blocks=2,
|
||||
attn_layers=[3],
|
||||
down_layers=[0],
|
||||
in_dim=data_dim,
|
||||
embed_dim=embed_dim,
|
||||
)
|
||||
self.decoder = Decoder1D(
|
||||
dim=hidden_dim,
|
||||
ch_mult=(1, 2, 4),
|
||||
num_res_blocks=2,
|
||||
attn_layers=[3],
|
||||
down_layers=[0],
|
||||
in_dim=data_dim,
|
||||
out_dim=data_dim,
|
||||
embed_dim=embed_dim,
|
||||
)
|
||||
|
||||
self.embed_dim = embed_dim
|
||||
# self.quant_conv = nn.Conv1d(2 * embed_dim, 2 * embed_dim, 1)
|
||||
# self.post_quant_conv = nn.Conv1d(embed_dim, embed_dim, 1)
|
||||
|
||||
self.initialize_weights()
|
||||
|
||||
def initialize_weights(self):
|
||||
pass
|
||||
|
||||
def encode(self, x: torch.Tensor, normalize: bool = True) -> DiagonalGaussianDistribution:
|
||||
if normalize:
|
||||
x = self.normalize(x)
|
||||
moments = self.encoder(x)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
return posterior
|
||||
|
||||
def decode(self, z: torch.Tensor, unnormalize: bool = True) -> torch.Tensor:
|
||||
dec = self.decoder(z)
|
||||
if unnormalize:
|
||||
dec = self.unnormalize(dec)
|
||||
return dec
|
||||
|
||||
def normalize(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return (x - self.data_mean) / self.data_std
|
||||
|
||||
def unnormalize(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return x * self.data_std + self.data_mean
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
sample_posterior: bool = True,
|
||||
rng: Optional[torch.Generator] = None,
|
||||
normalize: bool = True,
|
||||
unnormalize: bool = True,
|
||||
) -> tuple[torch.Tensor, DiagonalGaussianDistribution]:
|
||||
|
||||
posterior = self.encode(x, normalize=normalize)
|
||||
if sample_posterior:
|
||||
z = posterior.sample(rng)
|
||||
else:
|
||||
z = posterior.mode()
|
||||
dec = self.decode(z, unnormalize=unnormalize)
|
||||
return dec, posterior
|
||||
|
||||
def load_weights(self, src_dict) -> None:
|
||||
self.load_state_dict(src_dict, strict=True)
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
return next(self.parameters()).device
|
||||
|
||||
def get_last_layer(self):
|
||||
return self.decoder.conv_out.weight
|
||||
|
||||
def remove_weight_norm(self):
|
||||
for name, m in self.named_modules():
|
||||
if isinstance(m, MPConv1D):
|
||||
m.remove_weight_norm()
|
||||
log.debug(f"Removed weight norm from {name}")
|
||||
return self
|
||||
|
||||
|
||||
class Encoder1D(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
*,
|
||||
dim: int,
|
||||
ch_mult: tuple[int] = (1, 2, 4, 8),
|
||||
num_res_blocks: int,
|
||||
attn_layers: list[int] = [],
|
||||
down_layers: list[int] = [],
|
||||
resamp_with_conv: bool = True,
|
||||
in_dim: int,
|
||||
embed_dim: int,
|
||||
double_z: bool = True,
|
||||
kernel_size: int = 3,
|
||||
clip_act: float = 256.0):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_layers = len(ch_mult)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.in_channels = in_dim
|
||||
self.clip_act = clip_act
|
||||
self.down_layers = down_layers
|
||||
self.attn_layers = attn_layers
|
||||
self.conv_in = MPConv1D(in_dim, self.dim, kernel_size=kernel_size)
|
||||
|
||||
in_ch_mult = (1, ) + tuple(ch_mult)
|
||||
self.in_ch_mult = in_ch_mult
|
||||
# downsampling
|
||||
self.down = nn.ModuleList()
|
||||
for i_level in range(self.num_layers):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_in = dim * in_ch_mult[i_level]
|
||||
block_out = dim * ch_mult[i_level]
|
||||
for i_block in range(self.num_res_blocks):
|
||||
block.append(
|
||||
ResnetBlock1D(in_dim=block_in,
|
||||
out_dim=block_out,
|
||||
kernel_size=kernel_size,
|
||||
use_norm=True))
|
||||
block_in = block_out
|
||||
if i_level in attn_layers:
|
||||
attn.append(AttnBlock1D(block_in))
|
||||
down = nn.Module()
|
||||
down.block = block
|
||||
down.attn = attn
|
||||
if i_level in down_layers:
|
||||
down.downsample = Downsample1D(block_in, resamp_with_conv)
|
||||
self.down.append(down)
|
||||
|
||||
# middle
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = ResnetBlock1D(in_dim=block_in,
|
||||
out_dim=block_in,
|
||||
kernel_size=kernel_size,
|
||||
use_norm=True)
|
||||
self.mid.attn_1 = AttnBlock1D(block_in)
|
||||
self.mid.block_2 = ResnetBlock1D(in_dim=block_in,
|
||||
out_dim=block_in,
|
||||
kernel_size=kernel_size,
|
||||
use_norm=True)
|
||||
|
||||
# end
|
||||
self.conv_out = MPConv1D(block_in,
|
||||
2 * embed_dim if double_z else embed_dim,
|
||||
kernel_size=kernel_size)
|
||||
|
||||
self.learnable_gain = nn.Parameter(torch.zeros([]))
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
# downsampling
|
||||
hs = [self.conv_in(x)]
|
||||
for i_level in range(self.num_layers):
|
||||
for i_block in range(self.num_res_blocks):
|
||||
h = self.down[i_level].block[i_block](hs[-1])
|
||||
if len(self.down[i_level].attn) > 0:
|
||||
h = self.down[i_level].attn[i_block](h)
|
||||
h = h.clamp(-self.clip_act, self.clip_act)
|
||||
hs.append(h)
|
||||
if i_level in self.down_layers:
|
||||
hs.append(self.down[i_level].downsample(hs[-1]))
|
||||
|
||||
# middle
|
||||
h = hs[-1]
|
||||
h = self.mid.block_1(h)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h)
|
||||
h = h.clamp(-self.clip_act, self.clip_act)
|
||||
|
||||
# end
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h, gain=(self.learnable_gain + 1))
|
||||
return h
|
||||
|
||||
|
||||
class Decoder1D(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
*,
|
||||
dim: int,
|
||||
out_dim: int,
|
||||
ch_mult: tuple[int] = (1, 2, 4, 8),
|
||||
num_res_blocks: int,
|
||||
attn_layers: list[int] = [],
|
||||
down_layers: list[int] = [],
|
||||
kernel_size: int = 3,
|
||||
resamp_with_conv: bool = True,
|
||||
in_dim: int,
|
||||
embed_dim: int,
|
||||
clip_act: float = 256.0):
|
||||
super().__init__()
|
||||
self.ch = dim
|
||||
self.num_layers = len(ch_mult)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.in_channels = in_dim
|
||||
self.clip_act = clip_act
|
||||
self.down_layers = [i + 1 for i in down_layers] # each downlayer add one
|
||||
|
||||
# compute in_ch_mult, block_in and curr_res at lowest res
|
||||
block_in = dim * ch_mult[self.num_layers - 1]
|
||||
|
||||
# z to block_in
|
||||
self.conv_in = MPConv1D(embed_dim, block_in, kernel_size=kernel_size)
|
||||
|
||||
# middle
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, use_norm=True)
|
||||
self.mid.attn_1 = AttnBlock1D(block_in)
|
||||
self.mid.block_2 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, use_norm=True)
|
||||
|
||||
# upsampling
|
||||
self.up = nn.ModuleList()
|
||||
for i_level in reversed(range(self.num_layers)):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_out = dim * ch_mult[i_level]
|
||||
for i_block in range(self.num_res_blocks + 1):
|
||||
block.append(ResnetBlock1D(in_dim=block_in, out_dim=block_out, use_norm=True))
|
||||
block_in = block_out
|
||||
if i_level in attn_layers:
|
||||
attn.append(AttnBlock1D(block_in))
|
||||
up = nn.Module()
|
||||
up.block = block
|
||||
up.attn = attn
|
||||
if i_level in self.down_layers:
|
||||
up.upsample = Upsample1D(block_in, resamp_with_conv)
|
||||
self.up.insert(0, up) # prepend to get consistent order
|
||||
|
||||
# end
|
||||
self.conv_out = MPConv1D(block_in, out_dim, kernel_size=kernel_size)
|
||||
self.learnable_gain = nn.Parameter(torch.zeros([]))
|
||||
|
||||
def forward(self, z):
|
||||
# z to block_in
|
||||
h = self.conv_in(z)
|
||||
|
||||
# middle
|
||||
h = self.mid.block_1(h)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h)
|
||||
h = h.clamp(-self.clip_act, self.clip_act)
|
||||
|
||||
# upsampling
|
||||
for i_level in reversed(range(self.num_layers)):
|
||||
for i_block in range(self.num_res_blocks + 1):
|
||||
h = self.up[i_level].block[i_block](h)
|
||||
if len(self.up[i_level].attn) > 0:
|
||||
h = self.up[i_level].attn[i_block](h)
|
||||
h = h.clamp(-self.clip_act, self.clip_act)
|
||||
if i_level in self.down_layers:
|
||||
h = self.up[i_level].upsample(h)
|
||||
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h, gain=(self.learnable_gain + 1))
|
||||
return h
|
||||
|
||||
|
||||
def VAE_16k(**kwargs) -> VAE:
|
||||
return VAE(data_dim=80, embed_dim=20, hidden_dim=384, **kwargs)
|
||||
|
||||
|
||||
def VAE_44k(**kwargs) -> VAE:
|
||||
return VAE(data_dim=128, embed_dim=40, hidden_dim=512, **kwargs)
|
||||
|
||||
|
||||
def get_my_vae(name: str, **kwargs) -> VAE:
|
||||
if name == '16k':
|
||||
return VAE_16k(**kwargs)
|
||||
if name == '44k':
|
||||
return VAE_44k(**kwargs)
|
||||
raise ValueError(f'Unknown model: {name}')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
network = get_my_vae('standard')
|
||||
|
||||
# print the number of parameters in terms of millions
|
||||
num_params = sum(p.numel() for p in network.parameters()) / 1e6
|
||||
print(f'Number of parameters: {num_params:.2f}M')
|
||||
@@ -0,0 +1,117 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
from .edm2_utils import (MPConv1D, mp_silu, mp_sum, normalize)
|
||||
|
||||
|
||||
def nonlinearity(x):
|
||||
# swish
|
||||
return mp_silu(x)
|
||||
|
||||
|
||||
class ResnetBlock1D(nn.Module):
|
||||
|
||||
def __init__(self, *, in_dim, out_dim=None, conv_shortcut=False, kernel_size=3, use_norm=True):
|
||||
super().__init__()
|
||||
self.in_dim = in_dim
|
||||
out_dim = in_dim if out_dim is None else out_dim
|
||||
self.out_dim = out_dim
|
||||
self.use_conv_shortcut = conv_shortcut
|
||||
self.use_norm = use_norm
|
||||
|
||||
self.conv1 = MPConv1D(in_dim, out_dim, kernel_size=kernel_size)
|
||||
self.conv2 = MPConv1D(out_dim, out_dim, kernel_size=kernel_size)
|
||||
if self.in_dim != self.out_dim:
|
||||
if self.use_conv_shortcut:
|
||||
self.conv_shortcut = MPConv1D(in_dim, out_dim, kernel_size=kernel_size)
|
||||
else:
|
||||
self.nin_shortcut = MPConv1D(in_dim, out_dim, kernel_size=1)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
|
||||
# pixel norm
|
||||
if self.use_norm:
|
||||
x = normalize(x, dim=1)
|
||||
|
||||
h = x
|
||||
h = nonlinearity(h)
|
||||
h = self.conv1(h)
|
||||
|
||||
h = nonlinearity(h)
|
||||
h = self.conv2(h)
|
||||
|
||||
if self.in_dim != self.out_dim:
|
||||
if self.use_conv_shortcut:
|
||||
x = self.conv_shortcut(x)
|
||||
else:
|
||||
x = self.nin_shortcut(x)
|
||||
|
||||
return mp_sum(x, h, t=0.3)
|
||||
|
||||
|
||||
class AttnBlock1D(nn.Module):
|
||||
|
||||
def __init__(self, in_channels, num_heads=1):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.num_heads = num_heads
|
||||
self.qkv = MPConv1D(in_channels, in_channels * 3, kernel_size=1)
|
||||
self.proj_out = MPConv1D(in_channels, in_channels, kernel_size=1)
|
||||
|
||||
def forward(self, x):
|
||||
h = x
|
||||
y = self.qkv(h)
|
||||
y = y.reshape(y.shape[0], self.num_heads, -1, 3, y.shape[-1])
|
||||
q, k, v = normalize(y, dim=2).unbind(3)
|
||||
|
||||
q = rearrange(q, 'b h c l -> b h l c')
|
||||
k = rearrange(k, 'b h c l -> b h l c')
|
||||
v = rearrange(v, 'b h c l -> b h l c')
|
||||
|
||||
h = F.scaled_dot_product_attention(q, k, v)
|
||||
h = rearrange(h, 'b h l c -> b (h c) l')
|
||||
|
||||
h = self.proj_out(h)
|
||||
|
||||
return mp_sum(x, h, t=0.3)
|
||||
|
||||
|
||||
class Upsample1D(nn.Module):
|
||||
|
||||
def __init__(self, in_channels, with_conv):
|
||||
super().__init__()
|
||||
self.with_conv = with_conv
|
||||
if self.with_conv:
|
||||
self.conv = MPConv1D(in_channels, in_channels, kernel_size=3)
|
||||
|
||||
def forward(self, x):
|
||||
x = F.interpolate(x, scale_factor=2.0, mode='nearest-exact') # support 3D tensor(B,C,T)
|
||||
if self.with_conv:
|
||||
x = self.conv(x)
|
||||
return x
|
||||
|
||||
|
||||
class Downsample1D(nn.Module):
|
||||
|
||||
def __init__(self, in_channels, with_conv):
|
||||
super().__init__()
|
||||
self.with_conv = with_conv
|
||||
if self.with_conv:
|
||||
# no asymmetric padding in torch conv, must do it ourselves
|
||||
self.conv1 = MPConv1D(in_channels, in_channels, kernel_size=1)
|
||||
self.conv2 = MPConv1D(in_channels, in_channels, kernel_size=1)
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
if self.with_conv:
|
||||
x = self.conv1(x)
|
||||
|
||||
x = F.avg_pool1d(x, kernel_size=2, stride=2)
|
||||
|
||||
if self.with_conv:
|
||||
x = self.conv2(x)
|
||||
|
||||
return x
|
||||
File diff suppressed because it is too large
Load Diff
@@ -68,6 +68,13 @@ except Exception as e:
|
||||
LYNX_NODE_CLASS_MAPPINGS = {}
|
||||
LYNX_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
try:
|
||||
from .Ovi.nodes_ovi import NODE_CLASS_MAPPINGS as OVI_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as OVI_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
log.warning(f"WanVideoWrapper WARNING: Ovi nodes not available due to error in importing them: {e}")
|
||||
OVI_NODE_CLASS_MAPPINGS = {}
|
||||
OVI_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
|
||||
@@ -88,6 +95,7 @@ NODE_CLASS_MAPPINGS.update(S2V_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(HUMO_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(SAMPLER_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(LYNX_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(OVI_NODE_CLASS_MAPPINGS)
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
@@ -109,5 +117,6 @@ NODE_DISPLAY_NAME_MAPPINGS.update(S2V_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(HUMO_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(SAMPLER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(LYNX_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(OVI_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
@@ -113,6 +113,7 @@ class EasyCacheState:
|
||||
'cache': None,
|
||||
'accumulated_error': 0.0,
|
||||
'skipped_steps': [],
|
||||
'cache_ovi': None,
|
||||
}
|
||||
return pred_id
|
||||
|
||||
|
||||
+54
-22
@@ -892,8 +892,8 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
cnt += 1
|
||||
if cnt % 100 == 0:
|
||||
pbar.update(100)
|
||||
#for name, param in transformer.named_parameters():
|
||||
# print(name, param.device, param.dtype)
|
||||
|
||||
#[print(name, param.device, param.dtype) for name, param in transformer.named_parameters()]
|
||||
|
||||
pbar.update_absolute(0)
|
||||
|
||||
@@ -1112,7 +1112,8 @@ class WanVideoModelLoader:
|
||||
if "vace_blocks.0.after_proj.weight" in sd and not "patch_embedding.weight" in sd:
|
||||
raise ValueError("You are attempting to load a VACE module as a WanVideo model, instead you should use the vace_model input and matching T2V base model")
|
||||
|
||||
# currently this can be VAE or MTV-Crafter weights
|
||||
# currently this can be VACE, MTV-Crafter, Lynx or Ovi-audio weights
|
||||
extra_audio_model = False
|
||||
if extra_model is not None:
|
||||
for _model in extra_model:
|
||||
print("Loading extra model: ", _model["path"])
|
||||
@@ -1126,32 +1127,38 @@ class WanVideoModelLoader:
|
||||
if _model["path"].endswith(".gguf"):
|
||||
raise ValueError("With GGUF extra model the main model must also be GGUF quantized model")
|
||||
extra_sd = load_torch_file(_model["path"], device=transformer_load_device, safe_load=True)
|
||||
if "audio_model.patch_embedding.0.weight" in extra_sd:
|
||||
extra_audio_model = True
|
||||
sd.update(extra_sd)
|
||||
del extra_sd
|
||||
|
||||
first_key = next(iter(sd))
|
||||
if first_key.startswith("audio_model.") and not extra_audio_model:
|
||||
sd = {key.replace("audio_model.", "", 1): value for key, value in sd.items()}
|
||||
if first_key.startswith("model.diffusion_model."):
|
||||
new_sd = {}
|
||||
for key, value in sd.items():
|
||||
new_key = key.replace("model.diffusion_model.", "", 1)
|
||||
new_sd[new_key] = value
|
||||
sd = new_sd
|
||||
sd = {key.replace("model.diffusion_model.", "", 1): value for key, value in sd.items()}
|
||||
elif first_key.startswith("model."):
|
||||
new_sd = {}
|
||||
for key, value in sd.items():
|
||||
new_key = key.replace("model.", "", 1)
|
||||
new_sd[new_key] = value
|
||||
sd = new_sd
|
||||
if not "patch_embedding.weight" in sd:
|
||||
raise ValueError("Invalid WanVideo model selected")
|
||||
dim = sd["patch_embedding.weight"].shape[0]
|
||||
sd = {key.replace("model.", "", 1): value for key, value in sd.items()}
|
||||
|
||||
if "patch_embedding.weight" in sd:
|
||||
dim = sd["patch_embedding.weight"].shape[0]
|
||||
in_channels = sd["patch_embedding.weight"].shape[1]
|
||||
elif "patch_embedding.0.weight" in sd:
|
||||
dim = sd["patch_embedding.0.weight"].shape[0]
|
||||
in_channels = sd["patch_embedding.0.weight"].shape[1]
|
||||
else:
|
||||
raise ValueError("No patch_embedding weight found, is the selected model a full WanVideo model?")
|
||||
|
||||
in_features = sd["blocks.0.self_attn.k.weight"].shape[1]
|
||||
out_features = sd["blocks.0.self_attn.k.weight"].shape[0]
|
||||
in_channels = sd["patch_embedding.weight"].shape[1]
|
||||
log.info(f"Detected model in_channels: {in_channels}")
|
||||
ffn_dim = sd["blocks.0.ffn.0.bias"].shape[0]
|
||||
ffn2_dim = sd["blocks.0.ffn.2.weight"].shape[1]
|
||||
|
||||
patch_size=(1, 2, 2)
|
||||
if "patch_embedding.0.weight" in sd:
|
||||
patch_size = [1]
|
||||
|
||||
is_humo = "audio_proj.audio_proj_glob_1.layer.weight" in sd
|
||||
is_wananimate = "pose_patch_embedding.weight" in sd
|
||||
|
||||
@@ -1273,6 +1280,7 @@ class WanVideoModelLoader:
|
||||
"dim": dim,
|
||||
"in_features": in_features,
|
||||
"out_features": out_features,
|
||||
"patch_size": patch_size,
|
||||
"ffn_dim": ffn_dim,
|
||||
"ffn2_dim": ffn2_dim,
|
||||
"eps": 1e-06,
|
||||
@@ -1309,8 +1317,32 @@ class WanVideoModelLoader:
|
||||
}
|
||||
|
||||
with init_empty_weights():
|
||||
transformer = WanModel(**TRANSFORMER_CONFIG)
|
||||
transformer.eval()
|
||||
transformer = WanModel(**TRANSFORMER_CONFIG).eval()
|
||||
|
||||
if extra_audio_model:
|
||||
log.info("Ovi extra audio model detected, initializing...")
|
||||
TRANSFORMER_CONFIG.update({
|
||||
"patch_size": [1],
|
||||
"in_dim": 20,
|
||||
"out_dim": 20,
|
||||
})
|
||||
|
||||
with init_empty_weights():
|
||||
transformer.audio_model = WanModel(**TRANSFORMER_CONFIG).eval()
|
||||
|
||||
from .wanvideo.modules.model import WanLayerNorm, WanRMSNorm
|
||||
|
||||
for block in transformer.blocks:
|
||||
block.cross_attn.k_fusion = nn.Linear(block.dim, block.dim)
|
||||
block.cross_attn.v_fusion = nn.Linear(block.dim, block.dim)
|
||||
block.cross_attn.pre_attn_norm_fusion = WanLayerNorm(block.dim, elementwise_affine=True)
|
||||
block.cross_attn.norm_k_fusion = WanRMSNorm(block.dim, eps=1e-6) if block.qk_norm else nn.Identity()
|
||||
|
||||
for block in transformer.audio_model.blocks:
|
||||
block.cross_attn.k_fusion = nn.Linear(block.dim, block.dim)
|
||||
block.cross_attn.v_fusion = nn.Linear(block.dim, block.dim)
|
||||
block.cross_attn.pre_attn_norm_fusion = WanLayerNorm(block.dim, elementwise_affine=True)
|
||||
block.cross_attn.norm_k_fusion = WanRMSNorm(block.dim, eps=1e-6) if block.qk_norm else nn.Identity()
|
||||
|
||||
#ReCamMaster
|
||||
if "blocks.0.cam_encoder.weight" in sd:
|
||||
@@ -1422,10 +1454,10 @@ class WanVideoModelLoader:
|
||||
for k, v in sd.items():
|
||||
if k.endswith(".scale_weight"):
|
||||
scale_weights[k] = v.to(device, base_dtype)
|
||||
|
||||
if "fp8_e4m3fn" in quantization:
|
||||
|
||||
if quantization == "fp8_e4m3fn":
|
||||
weight_dtype = torch.float8_e4m3fn
|
||||
elif "fp8_e5m2" in quantization:
|
||||
elif quantization == "fp8_e5m2":
|
||||
weight_dtype = torch.float8_e5m2
|
||||
else:
|
||||
weight_dtype = base_dtype
|
||||
|
||||
+142
-106
@@ -46,9 +46,7 @@ class MetaParameter(torch.nn.Parameter):
|
||||
self.quant_type = quant_type
|
||||
return self
|
||||
|
||||
def offload_transformer(transformer):
|
||||
for block in transformer.blocks:
|
||||
block.kv_cache = None
|
||||
def offload_transformer(transformer):
|
||||
transformer.teacache_state.clear_all()
|
||||
transformer.magcache_state.clear_all()
|
||||
transformer.easycache_state.clear_all()
|
||||
@@ -73,6 +71,11 @@ def offload_transformer(transformer):
|
||||
else:
|
||||
transformer.to(offload_device)
|
||||
|
||||
for block in transformer.blocks:
|
||||
block.kv_cache = None
|
||||
if transformer.audio_model is not None and hasattr(block, 'audio_block'):
|
||||
block.audio_block = None
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
@@ -257,6 +260,18 @@ class WanVideoSampler:
|
||||
if arg not in step_sig.parameters:
|
||||
scheduler_step_args.pop(arg)
|
||||
|
||||
# Ovi
|
||||
if transformer.audio_model is not None: # temporary workaround (...nothing more permanent)
|
||||
for i, block in enumerate(transformer.blocks):
|
||||
block.audio_block = transformer.audio_model.blocks[i]
|
||||
sample_scheduler_ovi = copy.deepcopy(sample_scheduler)
|
||||
rope_function = "default" # comfy rope not implemented for ovi model yet
|
||||
ovi_negative_text_embeds = text_embeds.get("ovi_negative_prompt_embeds", None)
|
||||
ovi_audio_cfg = text_embeds.get("ovi_audio_cfg", None)
|
||||
if ovi_audio_cfg is not None:
|
||||
if not isinstance(ovi_audio_cfg, list):
|
||||
ovi_audio_cfg = [ovi_audio_cfg] * (steps + 1)
|
||||
|
||||
if isinstance(cfg, list):
|
||||
if steps < len(cfg):
|
||||
log.info(f"Received {len(cfg)} cfg values, but only {steps} steps. Slicing cfg list to match steps.")
|
||||
@@ -486,6 +501,21 @@ class WanVideoSampler:
|
||||
|
||||
pos_latent = neg_latent = None
|
||||
|
||||
# Ovi
|
||||
noise_audio = latent_ovi = seq_len_ovi = None
|
||||
if transformer.audio_model is not None:
|
||||
noise_audio = samples.get("latent_ovi_audio", None) if samples is not None else None
|
||||
if noise_audio is not None:
|
||||
if not torch.any(noise_audio):
|
||||
noise_audio = torch.randn(noise_audio.shape, device=torch.device("cpu"), dtype=torch.float32, generator=seed_g)
|
||||
else:
|
||||
noise_audio = noise_audio.squeeze().movedim(0, 1).to(device, dtype)
|
||||
else:
|
||||
noise_audio = torch.randn((157, 20), device=torch.device("cpu"), dtype=torch.float32, generator=seed_g) # T C
|
||||
log.info(f"Ovi audio latent shape: {noise_audio.shape}")
|
||||
latent_ovi = noise_audio
|
||||
seq_len_ovi = noise_audio.shape[0]
|
||||
|
||||
if transformer.dim == 1536 and humo_image_cond is not None: #small humo model
|
||||
#noise = torch.cat([noise[:, :-humo_reference_count], humo_image_cond[4:, -humo_reference_count:]], dim=1)
|
||||
pos_latent = humo_image_cond[4:, -humo_reference_count:].to(device, dtype)
|
||||
@@ -734,33 +764,35 @@ class WanVideoSampler:
|
||||
saved_generator_state = samples.get("generator_state", None)
|
||||
if saved_generator_state is not None:
|
||||
seed_g.set_state(saved_generator_state)
|
||||
input_samples = samples["samples"].squeeze(0).to(noise)
|
||||
if input_samples.shape[1] != noise.shape[1]:
|
||||
input_samples = torch.cat([input_samples[:, :1].repeat(1, noise.shape[1] - input_samples.shape[1], 1, 1), input_samples], dim=1)
|
||||
input_samples = samples.get("samples", None)
|
||||
if input_samples is not None:
|
||||
input_samples = input_samples.squeeze(0).to(noise)
|
||||
if input_samples.shape[1] != noise.shape[1]:
|
||||
input_samples = torch.cat([input_samples[:, :1].repeat(1, noise.shape[1] - input_samples.shape[1], 1, 1), input_samples], dim=1)
|
||||
|
||||
if add_noise_to_samples:
|
||||
latent_timestep = timesteps[:1].to(noise)
|
||||
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
|
||||
else:
|
||||
noise = input_samples
|
||||
if add_noise_to_samples:
|
||||
latent_timestep = timesteps[:1].to(noise)
|
||||
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
|
||||
else:
|
||||
noise = input_samples
|
||||
|
||||
noise_mask = samples.get("noise_mask", None)
|
||||
if noise_mask is not None:
|
||||
log.info(f"Latent noise_mask shape: {noise_mask.shape}")
|
||||
original_image = samples.get("original_image", None)
|
||||
if original_image is None:
|
||||
original_image = input_samples
|
||||
if len(noise_mask.shape) == 4:
|
||||
noise_mask = noise_mask.squeeze(1)
|
||||
if noise_mask.shape[0] < noise.shape[1]:
|
||||
noise_mask = noise_mask.repeat(noise.shape[1] // noise_mask.shape[0], 1, 1)
|
||||
noise_mask = samples.get("noise_mask", None)
|
||||
if noise_mask is not None:
|
||||
log.info(f"Latent noise_mask shape: {noise_mask.shape}")
|
||||
original_image = samples.get("original_image", None)
|
||||
if original_image is None:
|
||||
original_image = input_samples
|
||||
if len(noise_mask.shape) == 4:
|
||||
noise_mask = noise_mask.squeeze(1)
|
||||
if noise_mask.shape[0] < noise.shape[1]:
|
||||
noise_mask = noise_mask.repeat(noise.shape[1] // noise_mask.shape[0], 1, 1)
|
||||
|
||||
noise_mask = torch.nn.functional.interpolate(
|
||||
noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
|
||||
size=(noise.shape[1], noise.shape[2], noise.shape[3]),
|
||||
mode='trilinear',
|
||||
align_corners=False
|
||||
).repeat(1, noise.shape[0], 1, 1, 1)
|
||||
noise_mask = torch.nn.functional.interpolate(
|
||||
noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
|
||||
size=(noise.shape[1], noise.shape[2], noise.shape[3]),
|
||||
mode='trilinear',
|
||||
align_corners=False
|
||||
).repeat(1, noise.shape[0], 1, 1, 1)
|
||||
|
||||
# extra latents (Pusa) and 5b
|
||||
latents_to_insert = add_index = noise_multipliers = None
|
||||
@@ -797,7 +829,8 @@ class WanVideoSampler:
|
||||
if extra_channel_latents is not None:
|
||||
extra_channel_latents = extra_channel_latents[0].to(noise)
|
||||
|
||||
latent = noise.to(device)
|
||||
latent = noise
|
||||
print("Latent shape:", latent.shape, "Latent dtype:", latent.dtype, "Latent device:", latent.device)
|
||||
|
||||
#controlnet
|
||||
controlnet_latents = controlnet = None
|
||||
@@ -869,10 +902,8 @@ class WanVideoSampler:
|
||||
# Initialize cache state
|
||||
if samples is not None:
|
||||
previous_cache_states = samples.get("cache_states", None)
|
||||
print("Using previous cache states", previous_cache_states)
|
||||
if previous_cache_states is not None:
|
||||
log.info("Using cache states from previous sampler")
|
||||
|
||||
self.cache_state = previous_cache_states["cache_state"]
|
||||
transformer.easycache_state = previous_cache_states["easycache_state"]
|
||||
transformer.magcache_state = previous_cache_states["magcache_state"]
|
||||
@@ -1069,9 +1100,9 @@ class WanVideoSampler:
|
||||
add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None, fantasy_portrait_input=None, reverse_time=False,
|
||||
mtv_motion_tokens=None, s2v_audio_input=None, s2v_ref_motion=None, s2v_motion_frames=[1, 0], s2v_pose=None,
|
||||
humo_image_cond=None, humo_image_cond_neg=None, humo_audio=None, humo_audio_neg=None, wananim_pose_latents=None,
|
||||
wananim_face_pixels=None, uni3c_data=None,):
|
||||
wananim_face_pixels=None, uni3c_data=None, latent_model_input_ovi=None):
|
||||
nonlocal transformer
|
||||
#z = z.to(dtype)
|
||||
|
||||
autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear)
|
||||
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype) if autocast_enabled else nullcontext():
|
||||
|
||||
@@ -1301,6 +1332,9 @@ class WanVideoSampler:
|
||||
"wananim_pose_strength": wananim_pose_strength,
|
||||
"wananim_face_strength": wananim_face_strength,
|
||||
"lynx_embeds": lynx_embeds, # Lynx face and reference embeddings
|
||||
"x_ovi": [latent_model_input_ovi.to(z)] if latent_model_input_ovi is not None else None, # Audio latent model input for Ovi
|
||||
"seq_len_ovi": seq_len_ovi, # Audio latent model sequence length for Ovi
|
||||
"ovi_negative_text_embeds": ovi_negative_text_embeds, # Audio latent model negative text embeds for Ovi
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
@@ -1316,17 +1350,18 @@ class WanVideoSampler:
|
||||
#conditional (positive) pass
|
||||
if pos_latent is not None: # for humo
|
||||
base_params['x'] = [torch.cat([z[:, :-humo_reference_count], pos_latent], dim=1)]
|
||||
noise_pred_cond, cache_state_cond = transformer(
|
||||
noise_pred_cond, noise_pred_ovi, cache_state_cond = transformer(
|
||||
context=positive_embeds,
|
||||
pred_id=cache_state[0] if cache_state else None,
|
||||
vace_data=vace_data, attn_cond=attn_cond,
|
||||
**base_params
|
||||
)
|
||||
noise_pred_cond = noise_pred_cond[0]
|
||||
noise_pred_ovi = noise_pred_ovi[0] if noise_pred_ovi is not None else None
|
||||
if math.isclose(cfg_scale, 1.0):
|
||||
if use_fresca:
|
||||
noise_pred_cond = fourier_filter(noise_pred_cond, fresca_scale_low, fresca_scale_high, fresca_freq_cutoff)
|
||||
return noise_pred_cond, [cache_state_cond]
|
||||
return noise_pred_cond, noise_pred_ovi, [cache_state_cond]
|
||||
|
||||
#unconditional (negative) pass
|
||||
base_params['is_uncond'] = True
|
||||
@@ -1338,12 +1373,13 @@ class WanVideoSampler:
|
||||
if neg_latent is not None:
|
||||
base_params['x'] = [torch.cat([z[:, :-humo_reference_count], neg_latent], dim=1)]
|
||||
|
||||
noise_pred_uncond, cache_state_uncond = transformer(
|
||||
noise_pred_uncond, noise_pred_ovi_uncond, cache_state_uncond = transformer(
|
||||
context=negative_embeds if humo_audio_input_neg is None else positive_embeds, #ti #t
|
||||
pred_id=cache_state[1] if cache_state else None,
|
||||
vace_data=vace_data, attn_cond=attn_cond_neg,
|
||||
**base_params)
|
||||
noise_pred_uncond = noise_pred_uncond[0]
|
||||
noise_pred_ovi_uncond = noise_pred_ovi_uncond[0] if noise_pred_ovi_uncond is not None else None
|
||||
|
||||
# HuMo
|
||||
if not math.isclose(humo_audio_cfg_scale[idx], 1.0):
|
||||
@@ -1353,7 +1389,7 @@ class WanVideoSampler:
|
||||
if t > 980 and humo_image_cond_neg_input is not None: # use image cond for first timesteps
|
||||
base_params['y'] = [humo_image_cond_neg_input]
|
||||
|
||||
noise_pred_humo_audio_uncond, cache_state_humo = transformer(
|
||||
noise_pred_humo_audio_uncond, _, cache_state_humo = transformer(
|
||||
context=negative_embeds, pred_id=cache_state[2] if cache_state else None, vace_data=None,
|
||||
**base_params)
|
||||
|
||||
@@ -1364,13 +1400,13 @@ class WanVideoSampler:
|
||||
if cache_state is not None and len(cache_state) != 4:
|
||||
cache_state.append(None)
|
||||
# audio
|
||||
noise_pred_humo_null, cache_state_humo = transformer(
|
||||
noise_pred_humo_null, _, cache_state_humo = transformer(
|
||||
context=negative_embeds, pred_id=cache_state[2] if cache_state else None, vace_data=None,
|
||||
**base_params)
|
||||
# negative
|
||||
if humo_audio_input is not None:
|
||||
base_params['humo_audio'] = humo_audio_input
|
||||
noise_pred_humo_audio, cache_state_humo2 = transformer(
|
||||
noise_pred_humo_audio, _, cache_state_humo2 = transformer(
|
||||
context=positive_embeds, pred_id=cache_state[3] if cache_state else None, vace_data=None,
|
||||
**base_params)
|
||||
noise_pred = (humo_audio_cfg_scale[idx] * (noise_pred_cond - noise_pred_humo_audio[0])
|
||||
@@ -1383,7 +1419,7 @@ class WanVideoSampler:
|
||||
if use_phantom and not math.isclose(phantom_cfg_scale[idx], 1.0):
|
||||
if cache_state is not None and len(cache_state) != 3:
|
||||
cache_state.append(None)
|
||||
noise_pred_phantom, cache_state_phantom = transformer(
|
||||
noise_pred_phantom, _, cache_state_phantom = transformer(
|
||||
context=negative_embeds, pred_id=cache_state[2] if cache_state else None, vace_data=None,
|
||||
**base_params)
|
||||
|
||||
@@ -1404,7 +1440,7 @@ class WanVideoSampler:
|
||||
base_params['multitalk_audio'] = torch.zeros_like(multitalk_audio_input)[-1:]
|
||||
audio_context = negative_embeds
|
||||
base_params['is_uncond'] = False
|
||||
noise_pred_no_audio, cache_state_audio = transformer(
|
||||
noise_pred_no_audio, _, cache_state_audio = transformer(
|
||||
context=audio_context,
|
||||
pred_id=cache_state[2] if cache_state else None,
|
||||
vace_data=vace_data,
|
||||
@@ -1419,7 +1455,7 @@ class WanVideoSampler:
|
||||
base_params['is_uncond'] = False
|
||||
if cache_state is not None and len(cache_state) != 3:
|
||||
cache_state.append(None)
|
||||
noise_pred_lynx, cache_state_lynx = transformer(
|
||||
noise_pred_lynx, _, cache_state_lynx = transformer(
|
||||
context=negative_embeds, pred_id=cache_state[2] if cache_state else None, vace_data=None,
|
||||
**base_params)
|
||||
|
||||
@@ -1433,7 +1469,7 @@ class WanVideoSampler:
|
||||
base_params['y'] = [image_cond_input] * 2 if image_cond_input is not None else None
|
||||
base_params['clip_fea'] = torch.cat([clip_fea, clip_fea], dim=0)
|
||||
cache_state_uncond = None
|
||||
[noise_pred_cond, noise_pred_uncond], cache_state_cond = transformer(
|
||||
[noise_pred_cond, noise_pred_uncond], _, cache_state_cond = transformer(
|
||||
context=positive_embeds + negative_embeds, is_uncond=False,
|
||||
pred_id=cache_state[0] if cache_state else None,
|
||||
**base_params
|
||||
@@ -1471,7 +1507,14 @@ class WanVideoSampler:
|
||||
noise_pred = noise_pred_uncond_scaled + cfg_scale * (noise_pred_cond - noise_pred_uncond_scaled)
|
||||
del noise_pred_uncond_scaled, noise_pred_cond, noise_pred_uncond
|
||||
|
||||
return noise_pred, [cache_state_cond, cache_state_uncond]
|
||||
if latent_model_input_ovi is not None:
|
||||
if ovi_audio_cfg is None:
|
||||
audio_cfg_scale = cfg_scale - 1.0 if cfg_scale > 4.0 else cfg_scale
|
||||
else:
|
||||
audio_cfg_scale = ovi_audio_cfg[idx]
|
||||
noise_pred_ovi = noise_pred_ovi_uncond + audio_cfg_scale * (noise_pred_ovi - noise_pred_ovi_uncond)
|
||||
|
||||
return noise_pred, noise_pred_ovi, [cache_state_cond, cache_state_uncond]
|
||||
|
||||
if args.preview_method in [LatentPreviewMethod.Auto, LatentPreviewMethod.Latent2RGB]: #default for latent2rgb
|
||||
from latent_preview import prepare_callback
|
||||
@@ -1546,7 +1589,7 @@ class WanVideoSampler:
|
||||
# Store initial noise for first iteration
|
||||
if freeinit_args is not None and iter_idx == 0:
|
||||
initial_noise_saved = current_latent.detach().clone()
|
||||
if samples is not None:
|
||||
if input_samples is not None:
|
||||
current_latent = input_samples.to(device)
|
||||
continue
|
||||
|
||||
@@ -1586,6 +1629,7 @@ class WanVideoSampler:
|
||||
latent[:, add_index:add_index+num_extra_frames] = entry["samples"].to(latent)
|
||||
|
||||
latent_model_input = latent.to(device)
|
||||
latent_model_input_ovi = latent_ovi.to(device) if latent_ovi is not None else None
|
||||
|
||||
current_step_percentage = idx / len(timesteps)
|
||||
|
||||
@@ -1665,7 +1709,7 @@ class WanVideoSampler:
|
||||
partial_img_emb[:, 0, :, :] = source_image_cond[:, 0, :, :].to(intermediate_device)
|
||||
|
||||
partial_zt_src = zt_src[:, c, :, :]
|
||||
vt_src_context, new_teacache = predict_with_cfg(
|
||||
vt_src_context, _, new_teacache = predict_with_cfg(
|
||||
partial_zt_src, cfg[idx],
|
||||
positive, source_embeds["negative_prompt_embeds"],
|
||||
timestep, idx, partial_img_emb, control_latents,
|
||||
@@ -1679,7 +1723,7 @@ class WanVideoSampler:
|
||||
counter[:, c, :, :] += window_mask
|
||||
vt_src /= counter
|
||||
else:
|
||||
vt_src, self.cache_state_source = predict_with_cfg(
|
||||
vt_src, _, self.cache_state_source = predict_with_cfg(
|
||||
zt_src, cfg[idx],
|
||||
source_embeds["prompt_embeds"],
|
||||
source_embeds["negative_prompt_embeds"],
|
||||
@@ -1722,7 +1766,7 @@ class WanVideoSampler:
|
||||
partial_control_latents = control_latents[:, c, :, :]
|
||||
|
||||
partial_zt_tgt = zt_tgt[:, c, :, :]
|
||||
vt_tgt_context, new_teacache = predict_with_cfg(
|
||||
vt_tgt_context, _, new_teacache = predict_with_cfg(
|
||||
partial_zt_tgt, cfg[idx],
|
||||
positive, text_embeds["negative_prompt_embeds"],
|
||||
timestep, idx, partial_img_emb, partial_control_latents,
|
||||
@@ -1736,7 +1780,7 @@ class WanVideoSampler:
|
||||
counter[:, c, :, :] += window_mask
|
||||
vt_tgt /= counter
|
||||
else:
|
||||
vt_tgt, self.cache_state = predict_with_cfg(
|
||||
vt_tgt, _,self.cache_state = predict_with_cfg(
|
||||
zt_tgt, cfg[idx],
|
||||
text_embeds["prompt_embeds"],
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
@@ -1896,7 +1940,7 @@ class WanVideoSampler:
|
||||
partial_timestep = timestep
|
||||
#print("Partial timestep:", partial_timestep)
|
||||
|
||||
noise_pred_context, new_teacache = predict_with_cfg(
|
||||
noise_pred_context, _, new_teacache = predict_with_cfg(
|
||||
partial_latent_model_input,
|
||||
cfg[idx], positive,
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
@@ -2034,24 +2078,26 @@ class WanVideoSampler:
|
||||
|
||||
if samples is not None:
|
||||
noise_mask = samples.get("noise_mask", None)
|
||||
input_samples = samples["samples"].squeeze(0).to(noise)
|
||||
# Check if we have enough frames in input_samples
|
||||
if latent_end_idx > input_samples.shape[1]:
|
||||
# We need more frames than available - pad the input_samples at the end
|
||||
pad_length = latent_end_idx - input_samples.shape[1]
|
||||
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
|
||||
input_samples = torch.cat([input_samples, last_frame], dim=1)
|
||||
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
|
||||
if noise_mask is not None:
|
||||
original_image = input_samples.to(device)
|
||||
input_samples = samples["samples"]
|
||||
if input_samples is not None:
|
||||
input_samples = input_samples.squeeze(0).to(noise)
|
||||
# Check if we have enough frames in input_samples
|
||||
if latent_end_idx > input_samples.shape[1]:
|
||||
# We need more frames than available - pad the input_samples at the end
|
||||
pad_length = latent_end_idx - input_samples.shape[1]
|
||||
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
|
||||
input_samples = torch.cat([input_samples, last_frame], dim=1)
|
||||
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
|
||||
if noise_mask is not None:
|
||||
original_image = input_samples.to(device)
|
||||
|
||||
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
|
||||
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
|
||||
|
||||
if add_noise_to_samples:
|
||||
latent_timestep = timesteps[0]
|
||||
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
|
||||
else:
|
||||
noise = input_samples
|
||||
if add_noise_to_samples:
|
||||
latent_timestep = timesteps[0]
|
||||
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
|
||||
else:
|
||||
noise = input_samples
|
||||
|
||||
# diff diff prep
|
||||
if noise_mask is not None:
|
||||
@@ -2154,22 +2200,6 @@ class WanVideoSampler:
|
||||
else:
|
||||
positive = text_embeds["prompt_embeds"]
|
||||
|
||||
window_vace_data = None
|
||||
# if vace_data is not None:
|
||||
# window_vace_data = []
|
||||
# for vace_entry in vace_data:
|
||||
# partial_context = vace_entry["context"][0][:, latent_start_idx:latent_end_idx]
|
||||
# if has_ref:
|
||||
# partial_context[:, 0] = vace_entry["context"][0][:, 0]
|
||||
|
||||
# window_vace_data.append({
|
||||
# "context": [partial_context],
|
||||
# "scale": vace_entry["scale"],
|
||||
# "start": vace_entry["start"],
|
||||
# "end": vace_entry["end"],
|
||||
# "seq_len": vace_entry["seq_len"]
|
||||
# })
|
||||
|
||||
# uni3c slices
|
||||
if uni3c_embeds is not None:
|
||||
vae.to(device)
|
||||
@@ -2226,9 +2256,9 @@ class WanVideoSampler:
|
||||
if humo_image_cond is None or not is_first_clip:
|
||||
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
noise_pred, _, self.cache_state = predict_with_cfg(
|
||||
latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"],
|
||||
timestep, i, y, clip_embeds, control_latents, window_vace_data, partial_unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
timestep, i, y, clip_embeds, control_latents, None, partial_unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input,
|
||||
humo_image_cond=partial_humo_cond_input, humo_image_cond_neg=partial_humo_cond_neg_input, humo_audio=partial_humo_audio, humo_audio_neg=partial_humo_audio_neg,
|
||||
uni3c_data = uni3c_data)
|
||||
@@ -2457,7 +2487,7 @@ class WanVideoSampler:
|
||||
for i, t in enumerate(tqdm(timesteps, desc=f"Sampling audio indices {left_idx}-{right_idx}", position=0)):
|
||||
latent_model_input = latent.to(device)
|
||||
timestep = torch.tensor([t]).to(device)
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
noise_pred, _, self.cache_state = predict_with_cfg(
|
||||
latent_model_input,
|
||||
cfg[idx],
|
||||
text_embeds["prompt_embeds"],
|
||||
@@ -2621,24 +2651,26 @@ class WanVideoSampler:
|
||||
face_images_in = face_images[:, :, start:end].to(device, torch.float32) if face_images is not None else None
|
||||
|
||||
if samples is not None:
|
||||
input_samples = samples["samples"].squeeze(0).to(noise)
|
||||
# Check if we have enough frames in input_samples
|
||||
if end_latent > input_samples.shape[1]:
|
||||
# We need more frames than available - pad the input_samples at the end
|
||||
pad_length = end_latent - input_samples.shape[1]
|
||||
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
|
||||
input_samples = torch.cat([input_samples, last_frame], dim=1)
|
||||
input_samples = input_samples[:, start_latent:end_latent]
|
||||
if noise_mask is not None:
|
||||
original_image = input_samples.to(device)
|
||||
input_samples = samples["samples"]
|
||||
if input_samples is not None:
|
||||
input_samples = input_samples.squeeze(0).to(noise)
|
||||
# Check if we have enough frames in input_samples
|
||||
if latent_end_idx > input_samples.shape[1]:
|
||||
# We need more frames than available - pad the input_samples at the end
|
||||
pad_length = latent_end_idx - input_samples.shape[1]
|
||||
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
|
||||
input_samples = torch.cat([input_samples, last_frame], dim=1)
|
||||
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
|
||||
if noise_mask is not None:
|
||||
original_image = input_samples.to(device)
|
||||
|
||||
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
|
||||
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
|
||||
|
||||
if add_noise_to_samples:
|
||||
latent_timestep = timesteps[0]
|
||||
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
|
||||
else:
|
||||
noise = input_samples
|
||||
if add_noise_to_samples:
|
||||
latent_timestep = timesteps[0]
|
||||
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
|
||||
else:
|
||||
noise = input_samples
|
||||
|
||||
# diff diff prep
|
||||
noise_mask = samples.get("noise_mask", None)
|
||||
@@ -2701,7 +2733,7 @@ class WanVideoSampler:
|
||||
timestep = timesteps[i]
|
||||
latent_model_input = latent.to(device)
|
||||
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
noise_pred, _, self.cache_state = predict_with_cfg(
|
||||
latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"],
|
||||
timestep, i, cache_state=self.cache_state, image_cond=image_cond_in, clip_fea=clip_fea, wananim_face_pixels=face_images_in,
|
||||
wananim_pose_latents=pose_input_slice, uni3c_data=uni3c_data_input,
|
||||
@@ -2794,16 +2826,16 @@ class WanVideoSampler:
|
||||
|
||||
#region normal inference
|
||||
else:
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
latent_model_input,
|
||||
noise_pred, noise_pred_ovi, self.cache_state = predict_with_cfg(
|
||||
latent_model_input,
|
||||
cfg[idx], text_embeds["prompt_embeds"], text_embeds["negative_prompt_embeds"],
|
||||
timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, multitalk_audio_embeds=multitalk_audio_embeds, mtv_motion_tokens=mtv_motion_tokens, s2v_audio_input=s2v_audio_input,
|
||||
humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg,
|
||||
wananim_face_pixels=wananim_face_pixels, wananim_pose_latents=wananim_pose_latents, uni3c_data = uni3c_data,
|
||||
wananim_face_pixels=wananim_face_pixels, wananim_pose_latents=wananim_pose_latents, uni3c_data = uni3c_data, latent_model_input_ovi=latent_model_input_ovi
|
||||
)
|
||||
if bidirectional_sampling:
|
||||
noise_pred_flipped, self.cache_state = predict_with_cfg(
|
||||
noise_pred_flipped, _,self.cache_state = predict_with_cfg(
|
||||
latent_model_input_flipped,
|
||||
cfg[idx], text_embeds["prompt_embeds"], text_embeds["negative_prompt_embeds"],
|
||||
timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
@@ -2858,6 +2890,9 @@ class WanVideoSampler:
|
||||
**scheduler_step_args)[0].squeeze(0)
|
||||
latent_backwards = torch.flip(latent_backwards, dims=[1])
|
||||
latent = latent * 0.5 + latent_backwards * 0.5
|
||||
|
||||
if latent_ovi is not None:
|
||||
latent_ovi = sample_scheduler_ovi.step(noise_pred_ovi.unsqueeze(0), t, latent_ovi.to(device).unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
|
||||
|
||||
#InfiniteTalk first frame handling
|
||||
if (extra_latents is not None
|
||||
@@ -2941,7 +2976,8 @@ class WanVideoSampler:
|
||||
"drop_last": drop_last,
|
||||
"generator_state": seed_g.get_state(),
|
||||
"original_image": original_image.cpu() if original_image is not None else None,
|
||||
"cache_states": cache_states
|
||||
"cache_states": cache_states,
|
||||
"latent_ovi_audio": latent_ovi.unsqueeze(0).transpose(1, 2).cpu() if latent_ovi is not None else None,
|
||||
},{
|
||||
"samples": callback_latent.unsqueeze(0).cpu() if callback is not None else None,
|
||||
})
|
||||
|
||||
@@ -160,7 +160,8 @@ def patch_weight_to_device(self, key, device_to=None, inplace_update=False, back
|
||||
else:
|
||||
set_func(out_weight, inplace_update=inplace_update, seed=string_to_seed(key))
|
||||
|
||||
def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, dtype=None, base_dtype=None, state_dict=None, low_mem_load=False, control_lora=False, scale_weights={}):
|
||||
def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, dtype=None,
|
||||
base_dtype=None, state_dict=None, low_mem_load=False, control_lora=False, scale_weights={}):
|
||||
model.patch_weight_to_device = types.MethodType(patch_weight_to_device, model)
|
||||
to_load = []
|
||||
for n, m in model.model.named_modules():
|
||||
|
||||
+191
-52
@@ -238,7 +238,7 @@ def sinusoidal_embedding_1d(dim, position):
|
||||
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
|
||||
return x
|
||||
|
||||
def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0):
|
||||
def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0, freqs_scaling=1.0):
|
||||
assert dim % 2 == 0
|
||||
exponents = torch.arange(0, dim, 2, dtype=torch.float64).div(dim)
|
||||
inv_theta_pow = 1.0 / torch.pow(theta, exponents)
|
||||
@@ -246,6 +246,8 @@ def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0):
|
||||
if k > 0:
|
||||
print(f"RifleX: Using {k}th freq")
|
||||
inv_theta_pow[k-1] = 0.9 * 2 * torch.pi / L_test
|
||||
|
||||
inv_theta_pow *= freqs_scaling
|
||||
|
||||
freqs = torch.outer(torch.arange(max_seq_len), inv_theta_pow)
|
||||
freqs = torch.polar(torch.ones_like(freqs), freqs)
|
||||
@@ -254,6 +256,13 @@ def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0):
|
||||
@torch.autocast(device_type=mm.get_autocast_device(mm.get_torch_device()), enabled=False)
|
||||
@torch.compiler.disable()
|
||||
def rope_apply(x, grid_sizes, freqs, reverse_time=False):
|
||||
x_ndim = grid_sizes.shape[-1]
|
||||
if x_ndim == 3:
|
||||
return rope_apply_3d(x, grid_sizes, freqs, reverse_time=reverse_time)
|
||||
else:
|
||||
return rope_apply_1d(x, grid_sizes, freqs)
|
||||
|
||||
def rope_apply_3d(x, grid_sizes, freqs, reverse_time=False):
|
||||
n, c = x.size(2), x.size(3) // 2
|
||||
|
||||
# split freqs
|
||||
@@ -295,6 +304,30 @@ def rope_apply(x, grid_sizes, freqs, reverse_time=False):
|
||||
return torch.stack(output).to(x.dtype)
|
||||
|
||||
|
||||
def rope_apply_1d(x, grid_sizes, freqs):
|
||||
n, c = x.size(2), x.size(3) // 2 ## b l h d
|
||||
c_rope = freqs.shape[1] # number of complex dims to rotate
|
||||
assert c_rope <= c, "RoPE dimensions cannot exceed half of hidden size"
|
||||
|
||||
# loop over samples
|
||||
output = []
|
||||
for i, (l, ) in enumerate(grid_sizes.tolist()):
|
||||
seq_len = l
|
||||
# precompute multipliers
|
||||
x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(
|
||||
seq_len, n, -1, 2)) # [l n d//2]
|
||||
x_i_rope = x_i[:, :, :c_rope] * freqs[:seq_len, None, :] # [L, N, c_rope]
|
||||
x_i_passthrough = x_i[:, :, c_rope:] # untouched dims
|
||||
x_i = torch.cat([x_i_rope, x_i_passthrough], dim=2)
|
||||
|
||||
# apply rotary embedding
|
||||
x_i = torch.view_as_real(x_i).flatten(2)
|
||||
x_i = torch.cat([x_i, x[i, seq_len:]])
|
||||
|
||||
# append to collection
|
||||
output.append(x_i)
|
||||
return torch.stack(output).to(x.dtype)
|
||||
|
||||
class WanRMSNorm(nn.Module):
|
||||
|
||||
def __init__(self, dim, eps=1e-5):
|
||||
@@ -630,6 +663,7 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
super().__init__(in_features, out_features, num_heads, qk_norm, eps, kv_dim=kv_dim, rms_norm_function=rms_norm_function)
|
||||
self.attention_mode = attention_mode
|
||||
self.ip_adapter = None
|
||||
self.k_fusion = None
|
||||
|
||||
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0,
|
||||
num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy",
|
||||
@@ -688,8 +722,19 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
adapter_x = adapter_x.flatten(2)
|
||||
x[:, :orig_seq_len] = x[:, :orig_seq_len] + adapter_x * ip_scale
|
||||
|
||||
return self.o(x)
|
||||
if self.k_fusion is not None:
|
||||
# compute target attention
|
||||
target_seq = self.pre_attn_norm_fusion(kwargs["target_seq"])
|
||||
k_target = self.norm_k_fusion(self.k_fusion(target_seq)).view(b, -1, n, d)
|
||||
v_target = self.v_fusion(target_seq).view(b, -1, n, d)
|
||||
|
||||
q = rope_apply(q, grid_sizes, kwargs["src_freqs"])
|
||||
k_target = rope_apply(k_target, kwargs["target_grid_sizes"], kwargs["target_freqs"])
|
||||
target_x = attention(q, k_target, v_target, k_lens=kwargs["target_seq_lens"]).flatten(2)
|
||||
|
||||
x = x.add(target_x)
|
||||
|
||||
return self.o(x)
|
||||
|
||||
class WanI2VCrossAttention(WanSelfAttention):
|
||||
|
||||
@@ -918,18 +963,18 @@ class WanAttentionBlock(nn.Module):
|
||||
self.cross_attn.ip_adapter = WanLynxIPCrossAttention(cross_attention_dim=2048, dim=self.dim, n_registers=0, bias=False)
|
||||
|
||||
#@torch.compiler.disable()
|
||||
def get_mod(self, e):
|
||||
def get_mod(self, e, modulation):
|
||||
if e.dim() == 3:
|
||||
return (self.modulation + e).chunk(6, dim=1) # 1, 6, dim
|
||||
return (modulation + e).chunk(6, dim=1) # 1, 6, dim
|
||||
elif e.dim() == 4:
|
||||
e_mod = self.modulation.unsqueeze(2) + e
|
||||
e_mod = modulation.unsqueeze(2) + e
|
||||
return [ei.squeeze(1) for ei in e_mod.unbind(dim=1)]
|
||||
|
||||
def modulate(self, x, shift_msa, scale_msa, seg_idx=None):
|
||||
|
||||
def modulate(self, norm_x, shift_msa, scale_msa, seg_idx=None):
|
||||
"""
|
||||
Modulate x with shift and scale. If seg_idx is provided, apply segmented modulation.
|
||||
"""
|
||||
norm_x = self.norm1(x)
|
||||
if seg_idx is not None:
|
||||
parts = []
|
||||
for i in range(2):
|
||||
@@ -983,6 +1028,7 @@ class WanAttentionBlock(nn.Module):
|
||||
mtv_motion_tokens=None, mtv_motion_rotary_emb=None, mtv_strength=1.0, mtv_freqs=None, #mtv crafter
|
||||
humo_audio_input=None, humo_audio_scale=1.0, #humo audio
|
||||
lynx_x_ip=None, lynx_ref_feature=None, lynx_ip_scale=1.0, lynx_ref_scale=1.0, #lynx
|
||||
x_ovi=None, e_ovi=None, freqs_ovi=None, context_ovi=None, seq_lens_ovi=None, grid_sizes_ovi=None #ovi
|
||||
):
|
||||
r"""
|
||||
Args:
|
||||
@@ -1000,18 +1046,22 @@ class WanAttentionBlock(nn.Module):
|
||||
self.seg_idx = [0, self.seg_idx, x.size(1)]
|
||||
e = e[0]
|
||||
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device))
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device), self.modulation)
|
||||
del e
|
||||
input_x = self.modulate(x, shift_msa, scale_msa, seg_idx=self.seg_idx)
|
||||
input_x = self.modulate(self.norm1(x), shift_msa, scale_msa, seg_idx=self.seg_idx)
|
||||
del shift_msa, scale_msa
|
||||
|
||||
if x_ip is not None:
|
||||
shift_msa_ip, scale_msa_ip, gate_msa_ip, shift_mlp_ip, scale_mlp_ip, gate_mlp_ip = self.get_mod(e_ip.to(x.device))
|
||||
input_x_ip = self.modulate(x_ip, shift_msa_ip, scale_msa_ip)
|
||||
shift_msa_ip, scale_msa_ip, gate_msa_ip, shift_mlp_ip, scale_mlp_ip, gate_mlp_ip = self.get_mod(e_ip.to(x.device), self.modulation)
|
||||
input_x_ip = self.modulate(self.norm1(x_ip), shift_msa_ip, scale_msa_ip)
|
||||
self.cond_size = input_x_ip.shape[1]
|
||||
input_x = torch.concat([input_x, input_x_ip], dim=1)
|
||||
self.kv_cache = None
|
||||
|
||||
if x_ovi is not None:
|
||||
shift_msa_ovi, scale_msa_ovi, gate_msa_ovi, shift_mlp_ovi, scale_mlp_ovi, gate_mlp_ovi = self.get_mod(e_ovi.to(x.device), self.audio_block.modulation)
|
||||
input_x_ovi = self.modulate(self.audio_block.norm1(x_ovi), shift_msa_ovi, scale_msa_ovi)
|
||||
|
||||
if camera_embed is not None:
|
||||
# encode ReCamMaster camera
|
||||
camera_embed = self.cam_encoder(camera_embed.to(x))
|
||||
@@ -1054,8 +1104,16 @@ class WanAttentionBlock(nn.Module):
|
||||
elif self.rope_func == "comfy_chunked":
|
||||
q, k = apply_rope_comfy_chunked(q, k, freqs)
|
||||
else:
|
||||
q=rope_apply(q, grid_sizes, freqs, reverse_time=reverse_time)
|
||||
k=rope_apply(k, grid_sizes, freqs, reverse_time=reverse_time)
|
||||
q = rope_apply(q, grid_sizes, freqs, reverse_time=reverse_time)
|
||||
k = rope_apply(k, grid_sizes, freqs, reverse_time=reverse_time)
|
||||
|
||||
if x_ovi is not None:
|
||||
q_ovi, k_ovi, v_ovi = self.audio_block.self_attn.qkv_fn(input_x_ovi)
|
||||
q_ovi = rope_apply(q_ovi, grid_sizes_ovi, freqs_ovi)
|
||||
k_ovi = rope_apply(k_ovi, grid_sizes_ovi, freqs_ovi)
|
||||
y_ovi = self.audio_block.self_attn.forward(q_ovi, k_ovi, v_ovi, seq_lens_ovi)
|
||||
x_ovi = x_ovi.addcmul(y_ovi, gate_msa_ovi)
|
||||
|
||||
|
||||
# FETA
|
||||
if enhance_enabled:
|
||||
@@ -1076,7 +1134,7 @@ class WanAttentionBlock(nn.Module):
|
||||
current_step=current_step,
|
||||
video_attention_split_steps=video_attention_split_steps
|
||||
)
|
||||
elif ref_target_masks is not None:
|
||||
elif ref_target_masks is not None: #multi/infinite talk
|
||||
y, x_ref_attn_map = self.self_attn.forward_multitalk(q, k, v, seq_lens, grid_sizes, ref_target_masks)
|
||||
elif self.attention_mode == "radial_sage_attention":
|
||||
if self.dense_block or self.dense_timesteps is not None and current_step < self.dense_timesteps:
|
||||
@@ -1091,7 +1149,7 @@ class WanAttentionBlock(nn.Module):
|
||||
y = self.self_attn.forward(q, k, v, seq_lens, attention_mode_override="sageattn_3")
|
||||
else:
|
||||
y = self.self_attn.forward(q, k, v, seq_lens, attention_mode_override="sageattn")
|
||||
elif x_ip is not None and self.kv_cache is None:
|
||||
elif x_ip is not None and self.kv_cache is None: #stand-in
|
||||
# First pass: cache IP keys/values and compute attention
|
||||
self.kv_cache = {"k_ip": k_ip.detach(), "v_ip": v_ip.detach()}
|
||||
y = self.self_attn.forward_ip(q, k, v, q_ip, k_ip, v_ip, seq_lens)
|
||||
@@ -1112,17 +1170,19 @@ class WanAttentionBlock(nn.Module):
|
||||
if enhance_enabled:
|
||||
y.mul_(feta_scores)
|
||||
|
||||
#ReCamMaster
|
||||
# ReCamMaster
|
||||
if camera_embed is not None:
|
||||
y = self.projector(y)
|
||||
|
||||
# Stand-in
|
||||
if x_ip is not None:
|
||||
y, y_ip = (
|
||||
y[:, : -self.cond_size],
|
||||
y[:, -self.cond_size :],
|
||||
)
|
||||
|
||||
if self.zero_timestep:
|
||||
# S2V
|
||||
if self.zero_timestep:
|
||||
z = []
|
||||
for i in range(2):
|
||||
z.append(y[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_msa[:, i:i + 1])
|
||||
@@ -1134,7 +1194,30 @@ class WanAttentionBlock(nn.Module):
|
||||
|
||||
# cross-attention & ffn function
|
||||
if context is not None:
|
||||
if split_attn:
|
||||
if x_ovi is not None:
|
||||
#audio
|
||||
og_ovi_x = x_ovi
|
||||
x_ovi = x_ovi + self.audio_block.cross_attn(self.audio_block.norm3(x_ovi), context_ovi, grid_sizes_ovi,
|
||||
src_freqs=freqs_ovi,
|
||||
target_seq=x,
|
||||
target_seq_lens=seq_lens,
|
||||
target_grid_sizes=grid_sizes,
|
||||
target_freqs=freqs)
|
||||
y = self.audio_block.ffn(torch.addcmul(shift_mlp_ovi, self.audio_block.norm2(x_ovi), 1 + scale_mlp_ovi))
|
||||
x_ovi = x_ovi.addcmul(y, gate_mlp_ovi)
|
||||
|
||||
assert not torch.equal(og_ovi_x, x_ovi), "Audio should be changed after cross-attention!"
|
||||
|
||||
# video
|
||||
x = x + self.cross_attn(self.norm3(x), context, grid_sizes,
|
||||
src_freqs=freqs,
|
||||
target_seq=og_ovi_x,
|
||||
target_seq_lens=seq_lens_ovi,
|
||||
target_grid_sizes=grid_sizes_ovi,
|
||||
target_freqs=freqs_ovi)
|
||||
y = self.ffn(torch.addcmul(shift_mlp, self.norm2(x), 1 + scale_mlp))
|
||||
x = x.addcmul(y, gate_mlp)
|
||||
elif split_attn:
|
||||
if nag_context is not None:
|
||||
raise NotImplementedError("nag_context is not supported in split_cross_attn_ffn")
|
||||
x = self.split_cross_attn_ffn(x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed, grid_sizes)
|
||||
@@ -1154,12 +1237,12 @@ class WanAttentionBlock(nn.Module):
|
||||
x = x.addcmul(y, gate_mlp)
|
||||
del gate_mlp
|
||||
|
||||
if x_ip is not None:
|
||||
if x_ip is not None: #stand-in
|
||||
x_ip = x_ip.addcmul(y_ip, gate_msa_ip)
|
||||
y_ip = self.ffn(torch.addcmul(shift_mlp_ip, self.norm2(x_ip), 1 + scale_mlp_ip))
|
||||
x_ip = x_ip.addcmul(y_ip, gate_mlp_ip)
|
||||
|
||||
return x, x_ip, lynx_ref_feature
|
||||
return x, x_ip, lynx_ref_feature, x_ovi
|
||||
|
||||
|
||||
def cross_attn_ffn(self, x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed,
|
||||
@@ -1172,7 +1255,7 @@ class WanAttentionBlock(nn.Module):
|
||||
audio_proj=audio_proj, audio_scale=audio_scale,
|
||||
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond,
|
||||
rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs,
|
||||
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=self.original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale)
|
||||
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=self.original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, )
|
||||
# MultiTalk
|
||||
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
|
||||
x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding,
|
||||
@@ -1318,14 +1401,14 @@ class BaseWanAttentionBlock(WanAttentionBlock):
|
||||
self.block_id = block_id
|
||||
|
||||
def forward(self, x, vace_hints=None, vace_context_scale=[1.0], **kwargs):
|
||||
x, x_ip, lynx_ref_feature = super().forward(x, **kwargs)
|
||||
x, x_ip, lynx_ref_feature, x_ovi = super().forward(x, **kwargs)
|
||||
if vace_hints is None:
|
||||
return x, x_ip, lynx_ref_feature
|
||||
return x, x_ip, lynx_ref_feature, x_ovi
|
||||
|
||||
if self.block_id is not None:
|
||||
for i in range(len(vace_hints)):
|
||||
x.add_(vace_hints[i][self.block_id].to(x.device), alpha=vace_context_scale[i])
|
||||
return x, x_ip, lynx_ref_feature
|
||||
return x, x_ip, lynx_ref_feature, x_ovi
|
||||
|
||||
class Head(nn.Module):
|
||||
|
||||
@@ -1528,6 +1611,8 @@ class WanModel(torch.nn.Module):
|
||||
# lynx
|
||||
lynx_ip_layers=None,
|
||||
lynx_ref_layers=None,
|
||||
# ovi
|
||||
is_ovi_audio_model=False,
|
||||
):
|
||||
r"""
|
||||
Initialize the diffusion model backbone.
|
||||
@@ -1646,9 +1731,20 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
self.base_dtype = dtype
|
||||
|
||||
self.is_ovi_audio_model = patch_size == [1]
|
||||
|
||||
self.audio_model = None
|
||||
|
||||
# embeddings
|
||||
self.patch_embedding = nn.Conv3d(
|
||||
in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||
if not self.is_ovi_audio_model:
|
||||
self.patch_embedding = nn.Conv3d(in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||
else:
|
||||
from ...Ovi.audio_model_layers import ChannelLastConv1d, ConvMLP
|
||||
self.patch_embedding = nn.Sequential(
|
||||
ChannelLastConv1d(in_dim, dim, kernel_size=7, padding=3),
|
||||
nn.SiLU(),
|
||||
ConvMLP(dim, dim * 4, kernel_size=7, padding=3),
|
||||
)
|
||||
|
||||
self.original_patch_embedding = self.patch_embedding
|
||||
self.expanded_patch_embedding = self.patch_embedding
|
||||
@@ -2056,6 +2152,7 @@ class WanModel(torch.nn.Module):
|
||||
wananim_pose_latents=None, wananim_face_pixel_values=None,
|
||||
wananim_pose_strength=1.0, wananim_face_strength=1.0,
|
||||
lynx_embeds=None,
|
||||
x_ovi=None, seq_len_ovi=None, ovi_negative_text_embeds=None,
|
||||
):
|
||||
r"""
|
||||
Forward pass through the diffusion model
|
||||
@@ -2136,7 +2233,7 @@ class WanModel(torch.nn.Module):
|
||||
merged_audio_emb = audio_emb[:, s2v_motion_frames[1]:, :]
|
||||
|
||||
# params
|
||||
device = self.patch_embedding.weight.device
|
||||
device = self.main_device
|
||||
|
||||
if freqs is not None and freqs.device != device:
|
||||
freqs = freqs.to(device)
|
||||
@@ -2161,17 +2258,21 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
# patch embed
|
||||
if control_lora_enabled:
|
||||
self.expanded_patch_embedding.to(device)
|
||||
x = [
|
||||
self.expanded_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype)
|
||||
for u in x
|
||||
]
|
||||
self.expanded_patch_embedding.to(self.main_device)
|
||||
x = [self.expanded_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x]
|
||||
else:
|
||||
self.original_patch_embedding.to(self.main_device)
|
||||
x = [
|
||||
self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype)
|
||||
for u in x
|
||||
]
|
||||
x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x]
|
||||
|
||||
# ovi audio model
|
||||
if self.audio_model is not None:
|
||||
x_ovi = [self.audio_model.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x_ovi[0].dtype) for u in x_ovi]
|
||||
grid_sizes_ovi = torch.stack([torch.tensor(u.shape[1:2], dtype=torch.long) for u in x_ovi])
|
||||
seq_lens_ovi = torch.tensor([u.size(1) for u in x_ovi], dtype=torch.int32)
|
||||
x_ovi = torch.cat([torch.cat([u, u.new_zeros(1, seq_len_ovi - u.size(1), u.size(2))], dim=1) for u in x_ovi])
|
||||
d = self.dim // self.num_heads
|
||||
freqs_ovi = rope_params(1024, d - 4 * (d // 6), freqs_scaling=0.19676).to(self.main_device)
|
||||
x_ovi = x_ovi.to(self.main_device, self.base_dtype)
|
||||
|
||||
# WanAnimate
|
||||
motion_vec = None
|
||||
@@ -2190,10 +2291,11 @@ class WanModel(torch.nn.Module):
|
||||
fun_camera = self.control_adapter(fun_camera)
|
||||
x = [u + v for u, v in zip(x, fun_camera)]
|
||||
|
||||
# grid sizes and seq len
|
||||
grid_sizes = torch.stack([torch.tensor(u.shape[2:], device=device, dtype=torch.long) for u in x])
|
||||
original_grid_sizes = grid_sizes.clone()
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
|
||||
self.original_seq_len = x[0].shape[1]
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32)
|
||||
assert seq_lens.max() <= seq_len
|
||||
|
||||
@@ -2201,8 +2303,6 @@ class WanModel(torch.nn.Module):
|
||||
if self.trainable_cond_mask is not None:
|
||||
cond_mask_weight = self.trainable_cond_mask.weight.to(x[0]).unsqueeze(1).unsqueeze(1)
|
||||
|
||||
self.original_seq_len = x[0].shape[1]
|
||||
|
||||
if add_cond is not None:
|
||||
add_cond = self.add_conv_in(add_cond.to(self.add_conv_in.weight.dtype)).to(x[0].dtype)
|
||||
add_cond = add_cond.flatten(2).transpose(1, 2)
|
||||
@@ -2243,10 +2343,7 @@ class WanModel(torch.nn.Module):
|
||||
x = [torch.cat([u, end_ref_latent.unsqueeze(0)], dim=1) for end_ref_latent, u in zip(end_ref_latent, x)]
|
||||
|
||||
|
||||
x = torch.cat([
|
||||
torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))],
|
||||
dim=1) for u in x
|
||||
])
|
||||
x = torch.cat([torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))], dim=1) for u in x])
|
||||
|
||||
if self.trainable_cond_mask is not None:
|
||||
x = x + cond_mask_weight[0]
|
||||
@@ -2334,6 +2431,21 @@ class WanModel(torch.nn.Module):
|
||||
e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(time_embed_dtype)) # b, dim
|
||||
e0 = self.time_projection(e).unflatten(1, (6, self.dim)) # b, 6, dim
|
||||
|
||||
if self.audio_model is not None:
|
||||
#if t.dim() == 1:
|
||||
# t_ovi = t.unsqueeze(1).expand(t.size(0), seq_len_ovi)
|
||||
if t.dim() == 2:
|
||||
last_timestep = t[:, -1:]
|
||||
padding = last_timestep.expand(t.size(0), seq_len_ovi - t.size(1))
|
||||
t_ovi = torch.cat([t, padding], dim=1)
|
||||
|
||||
e_ovi = self.audio_model.time_embedding(sinusoidal_embedding_1d(self.audio_model.freq_dim, t_ovi.flatten()).to(time_embed_dtype)).unsqueeze(0) # b, dim
|
||||
e0_ovi = self.audio_model.time_projection(e_ovi).unflatten(2, (6, self.dim)).movedim(1, 2) # B, seq_len, 6, dim
|
||||
else:
|
||||
e_ovi = self.audio_model.time_embedding(sinusoidal_embedding_1d(self.audio_model.freq_dim, t.flatten()).to(time_embed_dtype)) # b, dim
|
||||
e0_ovi = self.audio_model.time_projection(e_ovi).unflatten(1, (6, self.dim)) # b, 6, dim
|
||||
|
||||
|
||||
#S2V zero timestep
|
||||
if self.zero_timestep:
|
||||
e = e[:-1]
|
||||
@@ -2386,13 +2498,17 @@ class WanModel(torch.nn.Module):
|
||||
raise NotImplementedError("nag_context is not supported with EchoShot")
|
||||
inner_c = [[u.shape[0] for u in context]]
|
||||
|
||||
if self.audio_model is not None:
|
||||
if is_uncond and ovi_negative_text_embeds is not None:
|
||||
context_ovi = ovi_negative_text_embeds
|
||||
else:
|
||||
context_ovi = context
|
||||
context_ovi = self.audio_model.text_embedding(
|
||||
torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context_ovi]).to(text_embed_dtype))
|
||||
|
||||
context = self.text_embedding(
|
||||
torch.stack([
|
||||
torch.cat(
|
||||
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
|
||||
for u in context
|
||||
]).to(text_embed_dtype))
|
||||
|
||||
torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context]).to(text_embed_dtype))
|
||||
|
||||
# NAG
|
||||
if nag_context is not None:
|
||||
nag_context = self.text_embedding(
|
||||
@@ -2547,6 +2663,7 @@ class WanModel(torch.nn.Module):
|
||||
previous_raw_input = state.get('previous_raw_input')
|
||||
previous_raw_output = state.get('previous_raw_output')
|
||||
cache = state.get('cache')
|
||||
cache_ovi = state.get('cache_ovi') if self.audio_model is not None else None
|
||||
accumulated_error = state.get('accumulated_error')
|
||||
k = state.get('k', 1)
|
||||
|
||||
@@ -2565,6 +2682,8 @@ class WanModel(torch.nn.Module):
|
||||
if accumulated_error < self.easycache_thresh:
|
||||
should_calc = False
|
||||
x = raw_input + cache.to(x.device)
|
||||
if cache_ovi is not None:
|
||||
x_ovi = x_ovi + cache_ovi.to(x_ovi.device)
|
||||
state['skipped_steps'].append(current_step)
|
||||
else:
|
||||
should_calc = True
|
||||
@@ -2579,6 +2698,8 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
if self.enable_easycache:
|
||||
original_x = x.clone().to(self.cache_device)
|
||||
if x_ovi is not None:
|
||||
original_x_ovi = x_ovi.clone().to(self.cache_device)
|
||||
if should_calc:
|
||||
if self.enable_teacache or self.enable_magcache:
|
||||
original_x = x.clone().to(self.cache_device)
|
||||
@@ -2623,6 +2744,13 @@ class WanModel(torch.nn.Module):
|
||||
lynx_ip_scale=lynx_ip_scale,
|
||||
lynx_ref_scale=lynx_ref_scale,
|
||||
)
|
||||
if self.audio_model is not None:
|
||||
kwargs['e_ovi'] = e0_ovi.to(self.base_dtype)
|
||||
kwargs['context_ovi'] = context_ovi
|
||||
kwargs['grid_sizes_ovi'] = grid_sizes_ovi
|
||||
kwargs['seq_lens_ovi'] = seq_lens_ovi
|
||||
kwargs['freqs_ovi'] = freqs_ovi
|
||||
|
||||
|
||||
if vace_data is not None:
|
||||
vace_hint_list = []
|
||||
@@ -2706,7 +2834,7 @@ class WanModel(torch.nn.Module):
|
||||
if b in self.slg_blocks and is_uncond:
|
||||
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
|
||||
continue
|
||||
x, x_ip, lynx_ref_feature = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, **kwargs) #run block
|
||||
x, x_ip, lynx_ref_feature, x_ovi = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, **kwargs) #run block
|
||||
if self.audio_injector is not None and s2v_audio_input is not None:
|
||||
x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v
|
||||
if block.has_face_fuser_block and motion_vec is not None:
|
||||
@@ -2758,8 +2886,11 @@ class WanModel(torch.nn.Module):
|
||||
previous_raw_output=x_out,
|
||||
cache=x.to(original_x.device) - original_x,
|
||||
k = output_change / input_change,
|
||||
accumulated_error = 0.0
|
||||
accumulated_error = 0.0,
|
||||
cache_ovi = x_ovi.clone().to(original_x.device) - original_x_ovi if x_ovi is not None else None
|
||||
)
|
||||
|
||||
|
||||
|
||||
if self.enable_easycache and (self.easycache_start_step <= current_step <= self.easycache_end_step) and pred_id is not None:
|
||||
self.easycache_state.update(
|
||||
@@ -2785,9 +2916,17 @@ class WanModel(torch.nn.Module):
|
||||
x = x[:, :self.original_seq_len]
|
||||
|
||||
x = self.head(x, e.to(x.device))
|
||||
|
||||
if x_ovi is not None:
|
||||
x_ovi = self.audio_model.head(x_ovi, e_ovi.to(x_ovi.device))
|
||||
grid_sizes_ovi = [gs[0] for gs in grid_sizes_ovi]
|
||||
assert len(x) == len(grid_sizes_ovi)
|
||||
x_ovi = [u[:gs] for u, gs in zip(x_ovi, grid_sizes_ovi)]
|
||||
x_ovi = [u.float() for u in x_ovi]
|
||||
|
||||
x = self.unpatchify(x, original_grid_sizes) # type: ignore[arg-type]
|
||||
x = [u.float() for u in x]
|
||||
return (x, pred_id) if pred_id is not None else (x, None)
|
||||
return (x, x_ovi, pred_id) if pred_id is not None else (x, x_ovi, None)
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
r"""
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import math
|
||||
|
||||
import librosa
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
@@ -63,31 +61,6 @@ class AudioEncoder():
|
||||
|
||||
self.video_rate = 30
|
||||
|
||||
def extract_audio_feat(self,
|
||||
audio_path,
|
||||
return_all_layers=False,
|
||||
dtype=torch.float32):
|
||||
audio_input, sample_rate = librosa.load(audio_path, sr=16000)
|
||||
|
||||
input_values = self.processor(
|
||||
audio_input, sampling_rate=sample_rate,
|
||||
return_tensors="pt").input_values
|
||||
|
||||
# INFERENCE
|
||||
|
||||
# retrieve logits & take argmax
|
||||
res = self.model(
|
||||
input_values.to(self.model.device), output_hidden_states=True)
|
||||
if return_all_layers:
|
||||
feat = torch.cat(res.hidden_states)
|
||||
else:
|
||||
feat = res.hidden_states[-1]
|
||||
feat = linear_interpolation(
|
||||
feat, input_fps=50, output_fps=self.video_rate)
|
||||
|
||||
z = feat.to(dtype) # Encoding for the motion
|
||||
return z
|
||||
|
||||
def get_audio_embed_bucket(self,
|
||||
audio_embed,
|
||||
stride=2,
|
||||
|
||||
@@ -21,17 +21,17 @@ def whitespace_clean(text):
|
||||
return text
|
||||
|
||||
|
||||
def canonicalize(text, keep_punctuation_exact_string=None):
|
||||
text = text.replace('_', ' ')
|
||||
if keep_punctuation_exact_string:
|
||||
text = keep_punctuation_exact_string.join(
|
||||
part.translate(str.maketrans('', '', string.punctuation))
|
||||
for part in text.split(keep_punctuation_exact_string))
|
||||
else:
|
||||
text = text.translate(str.maketrans('', '', string.punctuation))
|
||||
text = text.lower()
|
||||
text = re.sub(r'\s+', ' ', text)
|
||||
return text.strip()
|
||||
# def canonicalize(text, keep_punctuation_exact_string=None):
|
||||
# text = text.replace('_', ' ')
|
||||
# if keep_punctuation_exact_string:
|
||||
# text = keep_punctuation_exact_string.join(
|
||||
# part.translate(str.maketrans('', '', string.punctuation))
|
||||
# for part in text.split(keep_punctuation_exact_string))
|
||||
# else:
|
||||
# text = text.translate(str.maketrans('', '', string.punctuation))
|
||||
# text = text.lower()
|
||||
# text = re.sub(r'\s+', ' ', text)
|
||||
# return text.strip()
|
||||
|
||||
|
||||
class HuggingfaceTokenizer:
|
||||
|
||||
Reference in New Issue
Block a user