Add gradient checkpointing to ideogram4 vae
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user