From 7ccec8ec2ceb5753f13de857d1a6229012e32c24 Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Thu, 30 Apr 2026 04:27:13 -0600 Subject: [PATCH] Add checkpointing and a proper decode for flux 2 VAEs so they can be used with DFE --- .../diffusion_models/flux2/flux2_model.py | 15 ++++ .../diffusion_models/flux2/src/autoencoder.py | 79 +++++++++++++++---- 2 files changed, 80 insertions(+), 14 deletions(-) diff --git a/extensions_built_in/diffusion_models/flux2/flux2_model.py b/extensions_built_in/diffusion_models/flux2/flux2_model.py index 2a09f818..42d827da 100644 --- a/extensions_built_in/diffusion_models/flux2/flux2_model.py +++ b/extensions_built_in/diffusion_models/flux2/flux2_model.py @@ -523,3 +523,18 @@ class Flux2Model(BaseModel): latents = self.vae.encode(images) return latents + + def decode_latents(self, latents, device=None, dtype=None): + if device is None: + device = self.vae_device_torch + if dtype is None: + dtype = self.vae_torch_dtype + + # Move to vae to device if on cpu + if self.vae.device == torch.device("cpu"): + self.vae.to(device) + latents = latents.to(device, dtype=dtype) + + images = self.vae.decode(latents) + + return images diff --git a/extensions_built_in/diffusion_models/flux2/src/autoencoder.py b/extensions_built_in/diffusion_models/flux2/src/autoencoder.py index d5baef27..385c5fee 100644 --- a/extensions_built_in/diffusion_models/flux2/src/autoencoder.py +++ b/extensions_built_in/diffusion_models/flux2/src/autoencoder.py @@ -4,6 +4,7 @@ import torch from einops import rearrange from torch import Tensor, nn import math +import torch.utils.checkpoint as ckpt @dataclass @@ -177,24 +178,41 @@ class Encoder(nn.Module): self.conv_out = nn.Conv2d( block_in, 2 * z_channels, kernel_size=3, stride=1, padding=1 ) + self.gradient_checkpointing = False + + def enable_gradient_checkpointing(self): + self.gradient_checkpointing = True def forward(self, x: Tensor) -> Tensor: # downsampling hs = [self.conv_in(x)] for i_level in range(self.num_resolutions): for i_block in range(self.num_res_blocks): - h = self.down[i_level].block[i_block](hs[-1]) - if len(self.down[i_level].attn) > 0: - h = self.down[i_level].attn[i_block](h) + if torch.is_grad_enabled() and self.gradient_checkpointing: + h = ckpt.checkpoint(self.down[i_level].block[i_block], hs[-1]) + if len(self.down[i_level].attn) > 0: + h = ckpt.checkpoint(self.down[i_level].attn[i_block], h) + else: + h = self.down[i_level].block[i_block](hs[-1]) + if len(self.down[i_level].attn) > 0: + h = self.down[i_level].attn[i_block](h) hs.append(h) if i_level != self.num_resolutions - 1: - hs.append(self.down[i_level].downsample(hs[-1])) + if torch.is_grad_enabled() and self.gradient_checkpointing: + hs.append(ckpt.checkpoint(self.down[i_level].downsample, hs[-1])) + else: + hs.append(self.down[i_level].downsample(hs[-1])) # middle h = hs[-1] - h = self.mid.block_1(h) - h = self.mid.attn_1(h) - h = self.mid.block_2(h) + if torch.is_grad_enabled() and self.gradient_checkpointing: + h = ckpt.checkpoint(self.mid.block_1, h) + h = ckpt.checkpoint(self.mid.attn_1, h) + h = ckpt.checkpoint(self.mid.block_2, h) + else: + h = self.mid.block_1(h) + h = self.mid.attn_1(h) + h = self.mid.block_2(h) # end h = self.norm_out(h) h = swish(h) @@ -261,6 +279,10 @@ class Decoder(nn.Module): num_groups=32, num_channels=block_in, eps=1e-6, affine=True ) self.conv_out = nn.Conv2d(block_in, out_ch, kernel_size=3, stride=1, padding=1) + self.gradient_checkpointing = False + + def enable_gradient_checkpointing(self): + self.gradient_checkpointing = True def forward(self, z: Tensor) -> Tensor: z = self.post_quant_conv(z) @@ -272,20 +294,33 @@ class Decoder(nn.Module): h = self.conv_in(z) # middle - h = self.mid.block_1(h) - h = self.mid.attn_1(h) - h = self.mid.block_2(h) + if torch.is_grad_enabled() and self.gradient_checkpointing: + h = ckpt.checkpoint(self.mid.block_1, h) + h = ckpt.checkpoint(self.mid.attn_1, h) + h = ckpt.checkpoint(self.mid.block_2, h) + else: + h = self.mid.block_1(h) + h = self.mid.attn_1(h) + h = self.mid.block_2(h) # cast to proper dtype h = h.to(upscale_dtype) # upsampling for i_level in reversed(range(self.num_resolutions)): for i_block in range(self.num_res_blocks + 1): - h = self.up[i_level].block[i_block](h) - if len(self.up[i_level].attn) > 0: - h = self.up[i_level].attn[i_block](h) + if torch.is_grad_enabled() and self.gradient_checkpointing: + h = ckpt.checkpoint(self.up[i_level].block[i_block], h) + if len(self.up[i_level].attn) > 0: + h = ckpt.checkpoint(self.up[i_level].attn[i_block], h) + else: + h = self.up[i_level].block[i_block](h) + if len(self.up[i_level].attn) > 0: + h = self.up[i_level].attn[i_block](h) if i_level != 0: - h = self.up[i_level].upsample(h) + if torch.is_grad_enabled() and self.gradient_checkpointing: + h = ckpt.checkpoint(self.up[i_level].upsample, h) + else: + h = self.up[i_level].upsample(h) # end h = self.norm_out(h) @@ -326,6 +361,17 @@ class AutoEncoder(nn.Module): affine=False, track_running_stats=True, ) + self._gradient_checkpointing = False + + @property + def gradient_checkpointing(self): + return self._gradient_checkpointing + + @gradient_checkpointing.setter + def gradient_checkpointing(self, value: bool): + self._gradient_checkpointing = value + self.encoder.gradient_checkpointing = value + self.decoder.gradient_checkpointing = value @property def device(self): @@ -335,6 +381,11 @@ class AutoEncoder(nn.Module): def dtype(self): return next(self.parameters()).dtype + def enable_gradient_checkpointing(self): + self.gradient_checkpointing = True + self.encoder.enable_gradient_checkpointing() + self.decoder.enable_gradient_checkpointing() + def normalize(self, z): self.bn.eval() return self.bn(z)