Skip to content

wanvideo

Classes

fastvideo.models.dits.wanvideo.WanI2VCrossAttention

WanI2VCrossAttention(dim: int, num_heads: int, window_size=(-1, -1), qk_norm=True, eps=1e-06, supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None, quant_config: QuantizationConfig | None = None, prefix: str = '')

Bases: WanSelfAttention

Source code in fastvideo/models/dits/wanvideo.py
def __init__(
    self,
    dim: int,
    num_heads: int,
    window_size=(-1, -1),
    qk_norm=True,
    eps=1e-6,
    supported_attention_backends: tuple[AttentionBackendEnum, ...]
    | None = None,
    quant_config: QuantizationConfig | None = None,
    prefix: str = "",
) -> None:
    super().__init__(dim, num_heads, window_size, qk_norm, eps,
                     supported_attention_backends, quant_config=quant_config, prefix=prefix)

    self.add_k_proj = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.add_k_proj")
    self.add_v_proj = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.add_v_proj")
    self.norm_added_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
    self.norm_added_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()

Methods:

fastvideo.models.dits.wanvideo.WanI2VCrossAttention.forward
forward(x, context, context_lens)

Parameters:

Name Type Description Default
x Tensor

Shape [B, L1, C]

required
context Tensor

Shape [B, L2, C]

required
context_lens Tensor

Shape [B]

required
Source code in fastvideo/models/dits/wanvideo.py
def forward(self, x, context, context_lens):
    r"""
    Args:
        x(Tensor): Shape [B, L1, C]
        context(Tensor): Shape [B, L2, C]
        context_lens(Tensor): Shape [B]
    """
    context_img = context[:, :257]
    context = context[:, 257:]
    b, n, d = x.size(0), self.num_heads, self.head_dim

    # compute query, key, value
    q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
    k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
    v = self.to_v(context)[0].view(b, -1, n, d)
    k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
        b, -1, n, d)
    v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
    img_x = self.attn(q, k_img, v_img)
    # compute attention
    x = self.attn(q, k, v) if k.size(1) > 0 else torch.zeros_like(q)

    # output
    x = x.flatten(2)
    img_x = img_x.flatten(2)
    x = x + img_x
    x, _ = self.to_out(x)
    return x

fastvideo.models.dits.wanvideo.WanSelfAttention

WanSelfAttention(dim: int, num_heads: int, window_size=(-1, -1), qk_norm=True, eps=1e-06, parallel_attention=False, quant_config: QuantizationConfig | None = None, prefix: str = '')

Bases: Module

Source code in fastvideo/models/dits/wanvideo.py
def __init__(self,
             dim: int,
             num_heads: int,
             window_size=(-1, -1),
             qk_norm=True,
             eps=1e-6,
             parallel_attention=False,
             quant_config: QuantizationConfig | None = None,
             prefix: str = "") -> None:
    assert dim % num_heads == 0
    super().__init__()
    self.dim = dim
    self.num_heads = num_heads
    self.head_dim = dim // num_heads
    self.window_size = window_size
    self.qk_norm = qk_norm
    self.eps = eps
    self.parallel_attention = parallel_attention

    # layers
    self.to_q = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_q")
    self.to_k = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_k")
    self.to_v = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_v")
    self.to_out = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_out")
    self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
    self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()

    # Scaled dot product attention
    self.attn = LocalAttention(
        num_heads=num_heads,
        head_size=self.head_dim,
        dropout_rate=0,
        softmax_scale=None,
        causal=False,
        supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
                                      AttentionBackendEnum.TORCH_SDPA))

Methods:

fastvideo.models.dits.wanvideo.WanSelfAttention.forward
forward(x: Tensor, context: Tensor, context_lens: int)

Parameters:

Name Type Description Default
x Tensor

Shape [B, L, num_heads, C / num_heads]

required
seq_lens Tensor

Shape [B]

required
grid_sizes Tensor

Shape [B, 3], the second dimension contains (F, H, W)

required
freqs Tensor

Rope freqs, shape [1024, C / num_heads / 2]

required
Source code in fastvideo/models/dits/wanvideo.py
def forward(self, x: torch.Tensor, context: torch.Tensor,
            context_lens: int):
    r"""
    Args:
        x(Tensor): Shape [B, L, num_heads, C / num_heads]
        seq_lens(Tensor): Shape [B]
        grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
        freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
    """
    pass

fastvideo.models.dits.wanvideo.WanT2VCrossAttention

WanT2VCrossAttention(dim: int, num_heads: int, window_size=(-1, -1), qk_norm=True, eps=1e-06, parallel_attention=False, quant_config: QuantizationConfig | None = None, prefix: str = '')

Bases: WanSelfAttention

Source code in fastvideo/models/dits/wanvideo.py
def __init__(self,
             dim: int,
             num_heads: int,
             window_size=(-1, -1),
             qk_norm=True,
             eps=1e-6,
             parallel_attention=False,
             quant_config: QuantizationConfig | None = None,
             prefix: str = "") -> None:
    assert dim % num_heads == 0
    super().__init__()
    self.dim = dim
    self.num_heads = num_heads
    self.head_dim = dim // num_heads
    self.window_size = window_size
    self.qk_norm = qk_norm
    self.eps = eps
    self.parallel_attention = parallel_attention

    # layers
    self.to_q = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_q")
    self.to_k = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_k")
    self.to_v = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_v")
    self.to_out = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_out")
    self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
    self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()

    # Scaled dot product attention
    self.attn = LocalAttention(
        num_heads=num_heads,
        head_size=self.head_dim,
        dropout_rate=0,
        softmax_scale=None,
        causal=False,
        supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
                                      AttentionBackendEnum.TORCH_SDPA))

Methods:

fastvideo.models.dits.wanvideo.WanT2VCrossAttention.forward
forward(x, context, context_lens, crossattn_cache=None)

Parameters:

Name Type Description Default
x Tensor

Shape [B, L1, C]

required
context Tensor

Shape [B, L2, C]

required
context_lens Tensor

Shape [B]

required
Source code in fastvideo/models/dits/wanvideo.py
def forward(self, x, context, context_lens, crossattn_cache=None):
    r"""
    Args:
        x(Tensor): Shape [B, L1, C]
        context(Tensor): Shape [B, L2, C]
        context_lens(Tensor): Shape [B]
    """
    b, n, d = x.size(0), self.num_heads, self.head_dim

    # compute query, key, value
    q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)

    if crossattn_cache is not None:
        if not crossattn_cache["is_init"]:
            crossattn_cache["is_init"] = True
            k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
            v = self.to_v(context)[0].view(b, -1, n, d)
            crossattn_cache["k"] = k
            crossattn_cache["v"] = v
        else:
            k = crossattn_cache["k"]
            v = crossattn_cache["v"]
    else:
        k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
        v = self.to_v(context)[0].view(b, -1, n, d)

    # compute attention
    x = self.attn(q, k, v) if k.size(1) > 0 else torch.zeros_like(q)

    # output
    x = x.flatten(2)
    x, _ = self.to_out(x)
    return x

Functions: