support loading Fun camera model
Input not functional yet
This commit is contained in:
@@ -623,6 +623,8 @@ class WanVideoModelLoader:
|
||||
"inject_sample_info": True if "fps_embedding.weight" in sd else False,
|
||||
"add_ref_conv": True if "ref_conv.weight" in sd else False,
|
||||
"in_dim_ref_conv": sd["ref_conv.weight"].shape[1] if "ref_conv.weight" in sd else None,
|
||||
"add_control_adapter": True if "control_adapter.conv.weight" in sd else False,
|
||||
"in_dim_control_adapter": sd["control_adapter.conv.weight"].shape[1] if "control_adapter.conv.weight" in sd else None,
|
||||
}
|
||||
|
||||
with init_empty_weights():
|
||||
@@ -737,7 +739,7 @@ class WanVideoModelLoader:
|
||||
if quantization == "fp8_e4m3fn_fast_no_ffn":
|
||||
params_to_keep.update({"ffn"})
|
||||
print(params_to_keep)
|
||||
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep, sd=sd)
|
||||
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep)
|
||||
|
||||
del sd
|
||||
|
||||
|
||||
@@ -750,6 +750,8 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
inject_sample_info=False,
|
||||
add_ref_conv=False,
|
||||
in_dim_ref_conv=16,
|
||||
add_control_adapter=False,
|
||||
in_dim_control_adapter=24,
|
||||
):
|
||||
r"""
|
||||
Initialize the diffusion model backbone.
|
||||
@@ -911,6 +913,12 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
else:
|
||||
self.ref_conv = None
|
||||
|
||||
if add_control_adapter:
|
||||
from .wan_camera_adapter import SimpleAdapter
|
||||
self.control_adapter = SimpleAdapter(in_dim_control_adapter, dim, kernel_size=patch_size[1:], stride=patch_size[1:])
|
||||
else:
|
||||
self.control_adapter = None
|
||||
|
||||
def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None):
|
||||
log.info(f"Swapping {blocks_to_swap + 1} transformer blocks")
|
||||
self.blocks_to_swap = blocks_to_swap
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
#https://github.com/aigc-apps/VideoX-Fun/blob/wan_fun_v1.1/videox_fun/models/wan_camera_adapter.py
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
class SimpleAdapter(nn.Module):
|
||||
def __init__(self, in_dim, out_dim, kernel_size, stride, num_residual_blocks=1):
|
||||
super(SimpleAdapter, self).__init__()
|
||||
|
||||
# Pixel Unshuffle: reduce spatial dimensions by a factor of 8
|
||||
self.pixel_unshuffle = nn.PixelUnshuffle(downscale_factor=8)
|
||||
|
||||
# Convolution: reduce spatial dimensions by a factor
|
||||
# of 2 (without overlap)
|
||||
self.conv = nn.Conv2d(in_dim * 64, out_dim, kernel_size=kernel_size, stride=stride, padding=0)
|
||||
|
||||
# Residual blocks for feature extraction
|
||||
self.residual_blocks = nn.Sequential(
|
||||
*[ResidualBlock(out_dim) for _ in range(num_residual_blocks)]
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
# Reshape to merge the frame dimension into batch
|
||||
bs, c, f, h, w = x.size()
|
||||
x = x.permute(0, 2, 1, 3, 4).contiguous().view(bs * f, c, h, w)
|
||||
|
||||
# Pixel Unshuffle operation
|
||||
x_unshuffled = self.pixel_unshuffle(x)
|
||||
|
||||
# Convolution operation
|
||||
x_conv = self.conv(x_unshuffled)
|
||||
|
||||
# Feature extraction with residual blocks
|
||||
out = self.residual_blocks(x_conv)
|
||||
|
||||
# Reshape to restore original bf dimension
|
||||
out = out.view(bs, f, out.size(1), out.size(2), out.size(3))
|
||||
|
||||
# Permute dimensions to reorder (if needed), e.g., swap channels and feature frames
|
||||
out = out.permute(0, 2, 1, 3, 4)
|
||||
|
||||
return out
|
||||
|
||||
class ResidualBlock(nn.Module):
|
||||
def __init__(self, dim):
|
||||
super(ResidualBlock, self).__init__()
|
||||
self.conv1 = nn.Conv2d(dim, dim, kernel_size=3, padding=1)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
self.conv2 = nn.Conv2d(dim, dim, kernel_size=3, padding=1)
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
out = self.relu(self.conv1(x))
|
||||
out = self.conv2(out)
|
||||
out += residual
|
||||
return out
|
||||
|
||||
# Example usage
|
||||
# in_dim = 3
|
||||
# out_dim = 64
|
||||
# adapter = SimpleAdapterWithReshape(in_dim, out_dim)
|
||||
# x = torch.randn(1, in_dim, 4, 64, 64) # e.g., batch size = 1, channels = 3, frames/features = 4
|
||||
# output = adapter(x)
|
||||
# print(output.shape) # Should reflect transformed dimensions
|
||||
Reference in New Issue
Block a user