Fix DiT forward pass cond

This commit is contained in:
City
2023-10-29 12:51:20 +01:00
parent 1d5f745bbf
commit aee031b692
2 changed files with 32 additions and 37 deletions
+30 -35
View File
@@ -230,51 +230,46 @@ class DiT(nn.Module):
imgs = x.reshape(shape=(x.shape[0], c, h * p, h * p))
return imgs
def forward(self, x, timesteps, context, y, **kwargs):
def forward_raw(self, x, t, y):
"""
Forward pass of DiT.
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
t: (N,) tensor of diffusion timesteps
y: (N,) tensor of class labels
"""
# adapt to comfy inputs
t = timesteps
y = context.to(torch.int)
y = y[:, 0]
x = self.x_embedder(x) + self.pos_embed # (N, T, D), where T = H * W / patch_size ** 2
t = self.t_embedder(t) # (N, D) + change to model dtype
y = self.y_embedder(y, self.training) # (N, D)
c = t + y # (N, D)
x = self.x_embedder(x) + self.pos_embed # (N, T, D), where T = H * W / patch_size ** 2
t = self.t_embedder(t) # (N, D)
y = self.y_embedder(y, self.training) # (N, D)
c = t + y # (N, D)
for block in self.blocks:
x = block(x, c) # (N, T, D)
x = self.final_layer(x, c) # (N, T, patch_size ** 2 * out_channels)
x = self.unpatchify(x) # (N, out_channels, H, W)
x = block(x, c) # (N, T, D)
x = self.final_layer(x, c) # (N, T, patch_size ** 2 * out_channels)
x = self.unpatchify(x) # (N, out_channels, H, W)
return x
# adapt to comfy outputs
x = x.to(torch.float)
eps, rest = x[:, :self.in_channels], x[:, self.in_channels:]
def forward(self, x, timesteps, context, y=None, **kwargs):
"""
Forward pass that adapts comfy input to original forward function
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
timesteps: (N,) tensor of diffusion timesteps
context: (N, [LabelID]) conditioning
y: extra conditioning.
"""
## Remove outer array from cond
context = context[:, 0]
## run original forward pass
out = self.forward_raw(
x = x.to(self.dtype),
t = timesteps.to(self.dtype),
y = context.to(torch.int),
)
## only return EPS
out = out.to(torch.float)
eps, rest = out[:, :self.in_channels], out[:, self.in_channels:]
return eps
def forward_with_cfg(self, x, t, y, cfg_scale):
"""
Forward pass of DiT, but also batches the unconditional forward pass for classifier-free guidance.
"""
# https://github.com/openai/glide-text2im/blob/main/notebooks/text2im.ipynb
half = x[: len(x) // 2]
combined = torch.cat([half, half], dim=0)
model_out = self.forward(combined, t, y)
# For exact reproducibility reasons, we apply classifier-free guidance on only
# three channels by default. The standard approach to cfg applies it to all channels.
# This can be done by uncommenting the following line and commenting-out the line following that.
eps, rest = model_out[:, :self.in_channels], model_out[:, self.in_channels:]
# eps, rest = model_out[:, :3], model_out[:, 3:]
cond_eps, uncond_eps = torch.split(eps, len(eps) // 2, dim=0)
half_eps = uncond_eps + cfg_scale * (cond_eps - uncond_eps)
eps = torch.cat([half_eps, half_eps], dim=0)
return torch.cat([eps, rest], dim=1)
#################################################################################
# Sine/Cosine Positional Embedding Functions #
#################################################################################
+2 -2
View File
@@ -78,7 +78,7 @@ class DiTCondLabelSelect:
global label_data
class_labels = [int(k) for k,v in label_data.items() if v == label_name]
y = torch.tensor([[class_labels[0]]]).to(torch.int)
return ([[y, {"pooled_output": []}]], )
return ([[y, {}]], )
class DiTCondLabelEmpty:
@classmethod
@@ -100,7 +100,7 @@ class DiTCondLabelEmpty:
# [ID of last class + 1] == [num_classes]
y_null = model.model.dit_config["num_classes"]
y = torch.tensor([[y_null]]).to(torch.int)
return ([[y, {"pooled_output": []}]], )
return ([[y, {}]], )
NODE_CLASS_MAPPINGS = {
"DitCheckpointLoader" : DitCheckpointLoader,