Skip to content

flux2vae

Classes

fastvideo.models.vaes.flux2vae.AutoencoderKLFlux2

AutoencoderKLFlux2(config: Flux2VAEConfig)

Bases: Module, ParallelTiledVAE

A VAE model with KL loss for encoding images into latents and decoding latent representations into images.

This model inherits from [ParallelTiledVAE] for tiling support and uses standard diffusers Encoder/Decoder components for Flux2 image generation.

Source code in fastvideo/models/vaes/flux2vae.py
def __init__(
    self,
    config: Flux2VAEConfig,
):
    nn.Module.__init__(self)
    ParallelTiledVAE.__init__(self, config=config)

    self.config = config
    arch_config = config.arch_config

    in_channels: int = arch_config.in_channels
    out_channels: int = arch_config.out_channels
    down_block_types: Tuple[str, ...] = arch_config.down_block_types
    up_block_types: Tuple[str, ...] = arch_config.up_block_types
    block_out_channels: Tuple[int, ...] = arch_config.block_out_channels
    layers_per_block: int = arch_config.layers_per_block
    act_fn: str = arch_config.act_fn
    latent_channels: int = arch_config.latent_channels
    norm_num_groups: int = arch_config.norm_num_groups
    sample_size: int = arch_config.sample_size
    force_upcast: bool = arch_config.force_upcast
    use_quant_conv: bool = arch_config.use_quant_conv
    use_post_quant_conv: bool = arch_config.use_post_quant_conv
    mid_block_add_attention: bool = arch_config.mid_block_add_attention
    batch_norm_eps: float = arch_config.batch_norm_eps
    batch_norm_momentum: float = arch_config.batch_norm_momentum
    patch_size: Tuple[int, int] = arch_config.patch_size

    # pass init params to Encoder
    self.encoder = Encoder(
        in_channels=in_channels,
        out_channels=latent_channels,
        down_block_types=down_block_types,
        block_out_channels=block_out_channels,
        layers_per_block=layers_per_block,
        act_fn=act_fn,
        norm_num_groups=norm_num_groups,
        double_z=True,
        mid_block_add_attention=mid_block_add_attention,
    )

    # pass init params to Decoder
    self.decoder = Decoder(
        in_channels=latent_channels,
        out_channels=out_channels,
        up_block_types=up_block_types,
        block_out_channels=block_out_channels,
        layers_per_block=layers_per_block,
        norm_num_groups=norm_num_groups,
        act_fn=act_fn,
        mid_block_add_attention=mid_block_add_attention,
    )

    self.quant_conv = (
        nn.Conv2d(2 * latent_channels, 2 * latent_channels, 1)
        if use_quant_conv
        else None
    )
    self.post_quant_conv = (
        nn.Conv2d(latent_channels, latent_channels, 1)
        if use_post_quant_conv
        else None
    )

    self.bn = nn.BatchNorm2d(
        math.prod(patch_size) * latent_channels,
        eps=batch_norm_eps,
        momentum=batch_norm_momentum,
        affine=False,
        track_running_stats=True,
    )

    self.use_slicing = False
    self.use_tiling = False

    # only relevant if vae tiling is enabled
    self.tile_sample_min_size = sample_size
    sample_size_val = (
        sample_size[0]
        if isinstance(sample_size, (list, tuple))
        else sample_size
    )
    self.tile_latent_min_size = int(
        sample_size_val / (2 ** (len(block_out_channels) - 1))
    )
    self.tile_overlap_factor = 0.25

Attributes

fastvideo.models.vaes.flux2vae.AutoencoderKLFlux2.attn_processors property
attn_processors: Dict[str, AttentionProcessor]

Returns:

Type Description
Dict[str, AttentionProcessor]

dict of attention processors: A dictionary containing all attention processors used in the model with

Dict[str, AttentionProcessor]

indexed by its weight name.

Methods:

fastvideo.models.vaes.flux2vae.AutoencoderKLFlux2.decode
decode(z: FloatTensor, return_dict: bool = True, generator=None) -> Union[DecoderOutput, FloatTensor]

Decode a batch of images.

Parameters:

Name Type Description Default
z `torch.Tensor`

Input batch of latent vectors.

required
return_dict `bool`, *optional*, defaults to `True`

Whether to return a [~models.vae.DecoderOutput] instead of a plain tuple.

True

Returns:

Type Description
Union[DecoderOutput, FloatTensor]

[~models.vae.DecoderOutput] or tuple: If return_dict is True, a [~models.vae.DecoderOutput] is returned, otherwise a plain tuple is returned.

Source code in fastvideo/models/vaes/flux2vae.py
def decode(
    self, z: torch.FloatTensor, return_dict: bool = True, generator=None
) -> Union[DecoderOutput, torch.FloatTensor]:
    """
    Decode a batch of images.

    Args:
        z (`torch.Tensor`): Input batch of latent vectors.
        return_dict (`bool`, *optional*, defaults to `True`):
            Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.

    Returns:
        [`~models.vae.DecoderOutput`] or `tuple`:
            If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
            returned.

    """
    if self.use_slicing and z.shape[0] > 1:
        decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)]
        decoded = torch.cat(decoded_slices)
    else:
        decoded = self._decode(z).sample

    if not return_dict:
        return (decoded,)

    return DecoderOutput(sample=decoded)
fastvideo.models.vaes.flux2vae.AutoencoderKLFlux2.encode
encode(x: Tensor, return_dict: bool = True) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]

Encode a batch of images into latents.

Parameters:

Name Type Description Default
x `torch.Tensor`

Input batch of images.

required
return_dict `bool`, *optional*, defaults to `True`

Whether to return a [~models.autoencoder_kl.AutoencoderKLOutput] instead of a plain tuple.

True

Returns:

Type Description
Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]

The latent representations of the encoded images. If return_dict is True, a

Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]

[~models.autoencoder_kl.AutoencoderKLOutput] is returned, otherwise a plain tuple is returned.

Source code in fastvideo/models/vaes/flux2vae.py
def encode(
    self, x: torch.Tensor, return_dict: bool = True
) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]:
    """
    Encode a batch of images into latents.

    Args:
        x (`torch.Tensor`): Input batch of images.
        return_dict (`bool`, *optional*, defaults to `True`):
            Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.

    Returns:
            The latent representations of the encoded images. If `return_dict` is True, a
            [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned.
    """

    if x.ndim == 5:
        assert x.shape[2] == 1
        x = x.squeeze(2)

    if self.use_slicing and x.shape[0] > 1:
        encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)]
        h = torch.cat(encoded_slices)
    else:
        h = self._encode(x)

    posterior = DiagonalGaussianDistribution(h)
    if not return_dict:
        return (posterior,)

    return AutoencoderKLOutput(latent_dist=posterior)
fastvideo.models.vaes.flux2vae.AutoencoderKLFlux2.forward
forward(sample: Tensor, sample_posterior: bool = False, return_dict: bool = True, generator: Optional[Generator] = None) -> Union[DecoderOutput, Tensor]

Parameters:

Name Type Description Default
sample `torch.Tensor`

Input sample.

required
sample_posterior `bool`, *optional*, defaults to `False`

Whether to sample from the posterior.

False
return_dict `bool`, *optional*, defaults to `True`

Whether or not to return a [DecoderOutput] instead of a plain tuple.

True
Source code in fastvideo/models/vaes/flux2vae.py
def forward(
    self,
    sample: torch.Tensor,
    sample_posterior: bool = False,
    return_dict: bool = True,
    generator: Optional[torch.Generator] = None,
) -> Union[DecoderOutput, torch.Tensor]:
    r"""
    Args:
        sample (`torch.Tensor`): Input sample.
        sample_posterior (`bool`, *optional*, defaults to `False`):
            Whether to sample from the posterior.
        return_dict (`bool`, *optional*, defaults to `True`):
            Whether or not to return a [`DecoderOutput`] instead of a plain tuple.
    """
    x = sample
    posterior = self.encode(x).latent_dist
    if sample_posterior:
        z = posterior.sample(generator=generator)
    else:
        z = posterior.mode()
    dec = self.decode(z).sample

    if not return_dict:
        return (dec,)

    return DecoderOutput(sample=dec)
fastvideo.models.vaes.flux2vae.AutoencoderKLFlux2.set_attn_processor
set_attn_processor(processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]])

Sets the attention processor to use to compute attention.

Parameters:

Name Type Description Default
processor `dict` of `AttentionProcessor` or only `AttentionProcessor`

The instantiated processor class or a dictionary of processor classes that will be set as the processor for all Attention layers.

If processor is a dict, the key needs to define the path to the corresponding cross attention processor. This is strongly recommended when setting trainable attention processors.

required
Source code in fastvideo/models/vaes/flux2vae.py
def set_attn_processor(
    self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]
):
    r"""
    Sets the attention processor to use to compute attention.

    Parameters:
        processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
            The instantiated processor class or a dictionary of processor classes that will be set as the processor
            for **all** `Attention` layers.

            If `processor` is a dict, the key needs to define the path to the corresponding cross attention
            processor. This is strongly recommended when setting trainable attention processors.

    """
    count = len(self.attn_processors.keys())

    if isinstance(processor, dict) and len(processor) != count:
        raise ValueError(
            f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
            f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
        )

    def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
        if hasattr(module, "set_processor"):
            if not isinstance(processor, dict):
                module.set_processor(processor)
            else:
                module.set_processor(processor.pop(f"{name}.processor"))

        for sub_name, child in module.named_children():
            fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)

    for name, module in self.named_children():
        fn_recursive_attn_processor(name, module, processor)
fastvideo.models.vaes.flux2vae.AutoencoderKLFlux2.set_default_attn_processor
set_default_attn_processor()

Disables custom attention processors and sets the default attention implementation.

Source code in fastvideo/models/vaes/flux2vae.py
def set_default_attn_processor(self):
    """
    Disables custom attention processors and sets the default attention implementation.
    """
    if all(
        proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS
        for proc in self.attn_processors.values()
    ):
        processor = AttnAddedKVProcessor()
    elif all(
        proc.__class__ in CROSS_ATTENTION_PROCESSORS
        for proc in self.attn_processors.values()
    ):
        processor = AttnProcessor()
    else:
        raise ValueError(
            f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}"
        )

    self.set_attn_processor(processor)
fastvideo.models.vaes.flux2vae.AutoencoderKLFlux2.tiled_decode
tiled_decode(z: Tensor, return_dict: bool = True) -> Union[DecoderOutput, Tensor]

Decode a batch of images using a tiled decoder.

Parameters:

Name Type Description Default
z `torch.Tensor`

Input batch of latent vectors.

required
return_dict `bool`, *optional*, defaults to `True`

Whether or not to return a [~models.vae.DecoderOutput] instead of a plain tuple.

True

Returns:

Type Description
Union[DecoderOutput, Tensor]

[~models.vae.DecoderOutput] or tuple: If return_dict is True, a [~models.vae.DecoderOutput] is returned, otherwise a plain tuple is returned.

Source code in fastvideo/models/vaes/flux2vae.py
def tiled_decode(
    self, z: torch.Tensor, return_dict: bool = True
) -> Union[DecoderOutput, torch.Tensor]:
    r"""
    Decode a batch of images using a tiled decoder.

    Args:
        z (`torch.Tensor`): Input batch of latent vectors.
        return_dict (`bool`, *optional*, defaults to `True`):
            Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.

    Returns:
        [`~models.vae.DecoderOutput`] or `tuple`:
            If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
            returned.
    """
    overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor))
    blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor)
    row_limit = self.tile_sample_min_size - blend_extent

    # Split z into overlapping 64x64 tiles and decode them separately.
    # The tiles have an overlap to avoid seams between tiles.
    rows = []
    for i in range(0, z.shape[2], overlap_size):
        row = []
        for j in range(0, z.shape[3], overlap_size):
            tile = z[
                :,
                :,
                i : i + self.tile_latent_min_size,
                j : j + self.tile_latent_min_size,
            ]
            if self.post_quant_conv is not None:
                tile = self.post_quant_conv(tile)
            decoded = self.decoder(tile)
            row.append(decoded)
        rows.append(row)
    result_rows = []
    for i, row in enumerate(rows):
        result_row = []
        for j, tile in enumerate(row):
            # blend the above tile and the left tile
            # to the current tile and add the current tile to the result row
            if i > 0:
                tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
            if j > 0:
                tile = self.blend_h(row[j - 1], tile, blend_extent)
            result_row.append(tile[:, :, :row_limit, :row_limit])
        result_rows.append(torch.cat(result_row, dim=3))

    dec = torch.cat(result_rows, dim=2)
    if not return_dict:
        return (dec,)

    return DecoderOutput(sample=dec)
fastvideo.models.vaes.flux2vae.AutoencoderKLFlux2.tiled_encode
tiled_encode(x: Tensor, return_dict: bool = True) -> AutoencoderKLOutput

Encode a batch of images using a tiled encoder.

When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the output, but they should be much less noticeable.

Parameters:

Name Type Description Default
x `torch.Tensor`

Input batch of images.

required
return_dict `bool`, *optional*, defaults to `True`

Whether or not to return a [~models.autoencoder_kl.AutoencoderKLOutput] instead of a plain tuple.

True

Returns:

Type Description
AutoencoderKLOutput

[~models.autoencoder_kl.AutoencoderKLOutput] or tuple: If return_dict is True, a [~models.autoencoder_kl.AutoencoderKLOutput] is returned, otherwise a plain tuple is returned.

Source code in fastvideo/models/vaes/flux2vae.py
def tiled_encode(
    self, x: torch.Tensor, return_dict: bool = True
) -> AutoencoderKLOutput:
    r"""Encode a batch of images using a tiled encoder.

    When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several
    steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is
    different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the
    tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the
    output, but they should be much less noticeable.

    Args:
        x (`torch.Tensor`): Input batch of images.
        return_dict (`bool`, *optional*, defaults to `True`):
            Whether or not to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.

    Returns:
        [`~models.autoencoder_kl.AutoencoderKLOutput`] or `tuple`:
            If return_dict is True, a [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain
            `tuple` is returned.
    """
    deprecation_message = (
        "The tiled_encode implementation supporting the `return_dict` parameter is deprecated. In the future, the "
        "implementation of this method will be replaced with that of `_tiled_encode` and you will no longer be able "
        "to pass `return_dict`. You will also have to create a `DiagonalGaussianDistribution()` from the returned value."
    )

    overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor))
    blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
    row_limit = self.tile_latent_min_size - blend_extent

    # Split the image into 512x512 tiles and encode them separately.
    rows = []
    for i in range(0, x.shape[2], overlap_size):
        row = []
        for j in range(0, x.shape[3], overlap_size):
            tile = x[
                :,
                :,
                i : i + self.tile_sample_min_size,
                j : j + self.tile_sample_min_size,
            ]
            tile = self.encoder(tile)
            if self.quant_conv is not None:
                tile = self.quant_conv(tile)
            row.append(tile)
        rows.append(row)
    result_rows = []
    for i, row in enumerate(rows):
        result_row = []
        for j, tile in enumerate(row):
            # blend the above tile and the left tile
            # to the current tile and add the current tile to the result row
            if i > 0:
                tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
            if j > 0:
                tile = self.blend_h(row[j - 1], tile, blend_extent)
            result_row.append(tile[:, :, :row_limit, :row_limit])
        result_rows.append(torch.cat(result_row, dim=3))

    moments = torch.cat(result_rows, dim=2)
    posterior = DiagonalGaussianDistribution(moments)

    if not return_dict:
        return (posterior,)

    return AutoencoderKLOutput(latent_dist=posterior)