Skip to content

vllm_omni.diffusion.models.seedvr2.vae

SeedVR2 causal VAE with temporal tiling and spatial height sharding.

CausalConv3d

Bases: Conv3d

Replicate the first frame to pad only the past temporal context.

temporal_padding instance-attribute

temporal_padding = 2 * padding[0]

forward

forward(
    x: Tensor, context: TemporalContext | None = None
) -> Tensor

Decoder3d

Bases: Module

conv_in instance-attribute

conv_in = CausalConv3d(latent_channels, channels[-1])

conv_norm_out instance-attribute

conv_norm_out = nn.GroupNorm(
    groups, channels[0], eps=1e-06
)

conv_out instance-attribute

conv_out = CausalConv3d(channels[0], 3)

mid_block instance-attribute

mid_block = MidBlock3d(channels[-1], groups)

up_blocks instance-attribute

up_blocks = nn.ModuleList(
    [
        DecoderBlock3d(
            widths[max(0, i - 1)], width, groups, i
        )
        for i, width in enumerate(widths)
    ]
)

forward

forward(
    x: Tensor, context: TemporalContext | None = None
) -> Tensor

DecoderBlock3d

Bases: Module

resnets instance-attribute

resnets = nn.ModuleList(
    [ResnetBlock3d(in_channels, out_channels, groups)]
    + [
        ResnetBlock3d(out_channels, out_channels, groups)
        for _ in range(2)
    ]
)

upsamplers instance-attribute

upsamplers = nn.ModuleList(
    [Upsample3d(out_channels, temporal=index < 2)]
    if index < 3
    else []
)

forward

forward(
    x: Tensor, context: TemporalContext | None = None
) -> Tensor

Downsample3d

Bases: Module

conv instance-attribute

conv = CausalConv3d(
    channels,
    channels,
    (3 if temporal else 1, 3, 3),
    (2 if temporal else 1, 2, 2),
    (1 if temporal else 0, 0, 0),
)

forward

forward(
    x: Tensor, context: TemporalContext | None = None
) -> Tensor

Encoder3d

Bases: Module

conv_in instance-attribute

conv_in = CausalConv3d(3, channels[0])

conv_norm_out instance-attribute

conv_norm_out = nn.GroupNorm(
    groups, channels[-1], eps=1e-06
)

conv_out instance-attribute

conv_out = CausalConv3d(channels[-1], 2 * latent_channels)

down_blocks instance-attribute

down_blocks = nn.ModuleList(
    [
        EncoderBlock3d(
            channels[max(0, i - 1)], width, groups, i
        )
        for i, width in enumerate(channels)
    ]
)

mid_block instance-attribute

mid_block = MidBlock3d(channels[-1], groups)

forward

forward(
    x: Tensor, context: TemporalContext | None = None
) -> Tensor

EncoderBlock3d

Bases: Module

downsamplers instance-attribute

downsamplers = nn.ModuleList(
    [Downsample3d(out_channels, temporal=index >= 1)]
    if index < 3
    else []
)

resnets instance-attribute

resnets = nn.ModuleList(
    [
        ResnetBlock3d(in_channels, out_channels, groups),
        ResnetBlock3d(out_channels, out_channels, groups),
    ]
)

forward

forward(
    x: Tensor, context: TemporalContext | None = None
) -> Tensor

MidBlock3d

Bases: Module

attentions instance-attribute

attentions = nn.ModuleList(
    [
        Attention(
            channels,
            heads=1,
            dim_head=channels,
            eps=1e-06,
            norm_num_groups=groups,
            residual_connection=True,
            bias=True,
            upcast_softmax=True,
            _from_deprecated_attn_block=True,
        )
    ]
)

resnets instance-attribute

resnets = nn.ModuleList(
    [
        ResnetBlock3d(channels, channels, groups)
        for _ in range(2)
    ]
)

forward

forward(
    x: Tensor, context: TemporalContext | None = None
) -> Tensor

ResnetBlock3d

Bases: Module

conv1 instance-attribute

conv1 = CausalConv3d(in_channels, out_channels)

conv2 instance-attribute

conv2 = CausalConv3d(out_channels, out_channels)

conv_shortcut instance-attribute

conv_shortcut = (
    CausalConv3d(
        in_channels,
        out_channels,
        (1, 1, 1),
        padding=(0, 0, 0),
    )
    if in_channels != out_channels
    else None
)

norm1 instance-attribute

norm1 = nn.GroupNorm(groups, in_channels, eps=1e-06)

norm2 instance-attribute

norm2 = nn.GroupNorm(groups, out_channels, eps=1e-06)

forward

forward(
    x: Tensor, context: TemporalContext | None = None
) -> Tensor

SeedVR2VAE

Bases: Module, DistributedVaeMixin

Whole-clip VAE; sampling uses the caller's request-local generator.

decoder instance-attribute

decoder = Decoder3d(channels, groups, latent_channels)

encoder instance-attribute

encoder = Encoder3d(channels, groups, latent_channels)

spatial_scale_factor class-attribute instance-attribute

spatial_scale_factor = 8

temporal_scale_factor class-attribute instance-attribute

temporal_scale_factor = 4

use_tiling instance-attribute

use_tiling = False

decode

decode(
    z: Tensor, *, chunk_size: int | None = None
) -> Tensor

encode

encode(
    x: Tensor, *, chunk_size: int | None = None
) -> DiagonalGaussianDistribution

set_parallel_size

set_parallel_size(
    parallel_size: int, mode: str = "tile"
) -> None

SpatialContext dataclass

group instance-attribute

group: ProcessGroup

rank instance-attribute

rank: int

units instance-attribute

units: tuple[int, ...]

convolution

convolution(conv: Conv3d, x: Tensor) -> Tensor

gather

gather(x: Tensor) -> Tensor

normalize

normalize(norm: GroupNorm, x: Tensor) -> Tensor

split

split(x: Tensor) -> Tensor

TemporalContext dataclass

Past convolution inputs owned by one encode or decode call.

cache_temporal class-attribute instance-attribute

cache_temporal: bool = False

first class-attribute instance-attribute

first: bool = True

history class-attribute instance-attribute

history: dict[Module, Tensor] = field(default_factory=dict)

spatial class-attribute instance-attribute

spatial: SpatialContext | None = None

tile_spatial class-attribute instance-attribute

tile_spatial: bool = False

Upsample3d

Bases: Module

conv instance-attribute

conv = CausalConv3d(channels, channels)

temporal_ratio instance-attribute

temporal_ratio = 2 if temporal else 1

upscale_conv instance-attribute

upscale_conv = nn.Conv3d(
    channels,
    channels * 4 * self.temporal_ratio,
    kernel_size=1,
)

forward

forward(
    x: Tensor, context: TemporalContext | None = None
) -> Tensor

tiled_convolution

tiled_convolution(
    conv: Conv3d, x: Tensor, rows: int = 128
) -> Tensor

Tile output rows; only each tile's spatial padding/workspace is materialized.