FantasyPortrait reqs and cleanup

unsure if I want to add onnx in the requirements at this stage...
This commit is contained in:
kijai
2025-08-14 18:53:08 +03:00
parent 43447dbaf9
commit 0c4426fc2b
4 changed files with 10 additions and 80 deletions
+5 -2
View File
@@ -3,9 +3,12 @@ import time
import cv2
import numpy as np
import onnx
import onnxruntime
try:
import onnx
import onnxruntime
except ImportError:
print("ONNX and ONNX Runtime are required for FantasyPortrait. Please install them using 'pip install onnx onnxruntime-gpu'")
def create_onnx_session(onnx_path, gpu_id=None) -> onnxruntime.InferenceSession:
start = time.perf_counter()
-77
View File
@@ -3,8 +3,6 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from ..wanvideo.modules.attention import attention
def FeedForward(dim, mult=4):
inner_dim = int(dim * mult)
return nn.Sequential(
@@ -51,81 +49,6 @@ class MultiProjModel(nn.Module):
return result_final
class SingleStreamBlockProcessor(nn.Module):
def __init__(self, context_dim, hidden_dim):
super().__init__()
self.context_dim = context_dim
self.hidden_dim = hidden_dim
self.ip_adapter_single_stream_k_proj = nn.Linear(
context_dim, hidden_dim, bias=False
)
self.ip_adapter_single_stream_v_proj = nn.Linear(
context_dim, hidden_dim, bias=False
)
nn.init.zeros_(self.ip_adapter_single_stream_k_proj.weight)
nn.init.zeros_(self.ip_adapter_single_stream_v_proj.weight)
def __call__(
self,
attn: nn.Module,
x: torch.Tensor,
context: torch.Tensor,
context_lens: torch.Tensor,
adapter_proj: torch.Tensor,
adapter_context_lens: torch.Tensor,
latents_num_frames: int = 21,
ip_scale: float = 1.0,
adapter_attn_mask: torch.Tensor = None,
) -> torch.Tensor:
context_img = context[:, :257]
context = context[:, 257:]
b, n, d = x.size(0), attn.num_heads, attn.head_dim
# compute query, key, value
q = attn.norm_q(attn.q(x)).view(b, -1, n, d)
k = attn.norm_k(attn.k(context)).view(b, -1, n, d)
v = attn.v(context).view(b, -1, n, d)
k_img = attn.norm_k_img(attn.k_img(context_img)).view(b, -1, n, d)
v_img = attn.v_img(context_img).view(b, -1, n, d)
img_x = attention(q, k_img, v_img)
# compute attention
x = attention(q, k, v)
x = x.flatten(2)
img_x = img_x.flatten(2)
if len(adapter_proj.shape) == 4:
adapter_q = q.view(b * latents_num_frames, -1, n, d)
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(
b * latents_num_frames, -1, n, d
)
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(
b * latents_num_frames, -1, n, d
)
adapter_x = attention(
adapter_q, ip_key, ip_value, attn_mask=adapter_attn_mask
)
adapter_x = adapter_x.view(b, q.size(1), n, d)
adapter_x = adapter_x.flatten(2)
elif len(adapter_proj.shape) == 3:
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(
b, -1, n, d
)
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(
b, -1, n, d
)
adapter_x = attention(q, ip_key, ip_value, attn_mask=adapter_attn_mask)
adapter_x = adapter_x.flatten(2)
x = x + img_x + adapter_x * ip_scale
x = attn.o(x)
return x
class PerceiverAttention(nn.Module):
def __init__(self, *, dim, dim_head=64, heads=8):
super().__init__()