Add gradient checkpointing to ideogram4 vae

This commit is contained in:
Jaret Burkett
2026-06-15 10:00:40 -06:00
parent c730d64478
commit fcccc0fbd2

View File

@@ -7,6 +7,7 @@ import re
from dataclasses import dataclass, field
import torch
import torch.utils.checkpoint as ckpt
from einops import rearrange
from torch import Tensor, nn
@@ -182,24 +183,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]) # type: ignore[index, operator]
if len(self.down[i_level].attn) > 0: # type: ignore[arg-type]
h = self.down[i_level].attn[i_block](h) # type: ignore[index, operator]
if torch.is_grad_enabled() and self.gradient_checkpointing:
h = ckpt.checkpoint(self.down[i_level].block[i_block], hs[-1]) # type: ignore[index, operator]
if len(self.down[i_level].attn) > 0: # type: ignore[arg-type]
h = ckpt.checkpoint(self.down[i_level].attn[i_block], h) # type: ignore[index, operator]
else:
h = self.down[i_level].block[i_block](hs[-1]) # type: ignore[index, operator]
if len(self.down[i_level].attn) > 0: # type: ignore[arg-type]
h = self.down[i_level].attn[i_block](h) # type: ignore[index, operator]
hs.append(h)
if i_level != self.num_resolutions - 1:
hs.append(self.down[i_level].downsample(hs[-1])) # type: ignore[operator]
if torch.is_grad_enabled() and self.gradient_checkpointing:
hs.append(ckpt.checkpoint(self.down[i_level].downsample, hs[-1])) # type: ignore[operator]
else:
hs.append(self.down[i_level].downsample(hs[-1])) # type: ignore[operator]
# middle
h = hs[-1]
h = self.mid.block_1(h) # type: ignore[operator]
h = self.mid.attn_1(h) # type: ignore[operator]
h = self.mid.block_2(h) # type: ignore[operator]
if torch.is_grad_enabled() and self.gradient_checkpointing:
h = ckpt.checkpoint(self.mid.block_1, h) # type: ignore[operator]
h = ckpt.checkpoint(self.mid.attn_1, h) # type: ignore[operator]
h = ckpt.checkpoint(self.mid.block_2, h) # type: ignore[operator]
else:
h = self.mid.block_1(h) # type: ignore[operator]
h = self.mid.attn_1(h) # type: ignore[operator]
h = self.mid.block_2(h) # type: ignore[operator]
# end
h = self.norm_out(h)
h = swish(h)
@@ -266,6 +284,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)
@@ -277,20 +299,33 @@ class Decoder(nn.Module):
h = self.conv_in(z)
# middle
h = self.mid.block_1(h) # type: ignore[operator]
h = self.mid.attn_1(h) # type: ignore[operator]
h = self.mid.block_2(h) # type: ignore[operator]
if torch.is_grad_enabled() and self.gradient_checkpointing:
h = ckpt.checkpoint(self.mid.block_1, h) # type: ignore[operator]
h = ckpt.checkpoint(self.mid.attn_1, h) # type: ignore[operator]
h = ckpt.checkpoint(self.mid.block_2, h) # type: ignore[operator]
else:
h = self.mid.block_1(h) # type: ignore[operator]
h = self.mid.attn_1(h) # type: ignore[operator]
h = self.mid.block_2(h) # type: ignore[operator]
# 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) # type: ignore[index, operator]
if len(self.up[i_level].attn) > 0: # type: ignore[arg-type]
h = self.up[i_level].attn[i_block](h) # type: ignore[index, operator]
if torch.is_grad_enabled() and self.gradient_checkpointing:
h = ckpt.checkpoint(self.up[i_level].block[i_block], h) # type: ignore[index, operator]
if len(self.up[i_level].attn) > 0: # type: ignore[arg-type]
h = ckpt.checkpoint(self.up[i_level].attn[i_block], h) # type: ignore[index, operator]
else:
h = self.up[i_level].block[i_block](h) # type: ignore[index, operator]
if len(self.up[i_level].attn) > 0: # type: ignore[arg-type]
h = self.up[i_level].attn[i_block](h) # type: ignore[index, operator]
if i_level != 0:
h = self.up[i_level].upsample(h) # type: ignore[operator]
if torch.is_grad_enabled() and self.gradient_checkpointing:
h = ckpt.checkpoint(self.up[i_level].upsample, h) # type: ignore[operator]
else:
h = self.up[i_level].upsample(h) # type: ignore[operator]
# end
h = self.norm_out(h)
@@ -331,6 +366,22 @@ class AutoEncoder(nn.Module):
affine=False,
track_running_stats=True,
)
self._gradient_checkpointing = False
@property
def gradient_checkpointing(self) -> bool:
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
def enable_gradient_checkpointing(self):
self.gradient_checkpointing = True
self.encoder.enable_gradient_checkpointing()
self.decoder.enable_gradient_checkpointing()
@property
def device(self) -> torch.device: