Fix DiT forward pass cond
This commit is contained in:
+24
-29
@@ -230,51 +230,46 @@ class DiT(nn.Module):
|
|||||||
imgs = x.reshape(shape=(x.shape[0], c, h * p, h * p))
|
imgs = x.reshape(shape=(x.shape[0], c, h * p, h * p))
|
||||||
return imgs
|
return imgs
|
||||||
|
|
||||||
def forward(self, x, timesteps, context, y, **kwargs):
|
def forward_raw(self, x, t, y):
|
||||||
"""
|
"""
|
||||||
Forward pass of DiT.
|
Forward pass of DiT.
|
||||||
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
||||||
t: (N,) tensor of diffusion timesteps
|
t: (N,) tensor of diffusion timesteps
|
||||||
y: (N,) tensor of class labels
|
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
|
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
|
t = self.t_embedder(t) # (N, D)
|
||||||
y = self.y_embedder(y, self.training) # (N, D)
|
y = self.y_embedder(y, self.training) # (N, D)
|
||||||
c = t + y # (N, D)
|
c = t + y # (N, D)
|
||||||
for block in self.blocks:
|
for block in self.blocks:
|
||||||
x = block(x, c) # (N, T, D)
|
x = block(x, c) # (N, T, D)
|
||||||
x = self.final_layer(x, c) # (N, T, patch_size ** 2 * out_channels)
|
x = self.final_layer(x, c) # (N, T, patch_size ** 2 * out_channels)
|
||||||
x = self.unpatchify(x) # (N, out_channels, H, W)
|
x = self.unpatchify(x) # (N, out_channels, H, W)
|
||||||
|
return x
|
||||||
|
|
||||||
# adapt to comfy outputs
|
def forward(self, x, timesteps, context, y=None, **kwargs):
|
||||||
x = x.to(torch.float)
|
"""
|
||||||
eps, rest = x[:, :self.in_channels], x[:, self.in_channels:]
|
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
|
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 #
|
# Sine/Cosine Positional Embedding Functions #
|
||||||
#################################################################################
|
#################################################################################
|
||||||
|
|||||||
+2
-2
@@ -78,7 +78,7 @@ class DiTCondLabelSelect:
|
|||||||
global label_data
|
global label_data
|
||||||
class_labels = [int(k) for k,v in label_data.items() if v == label_name]
|
class_labels = [int(k) for k,v in label_data.items() if v == label_name]
|
||||||
y = torch.tensor([[class_labels[0]]]).to(torch.int)
|
y = torch.tensor([[class_labels[0]]]).to(torch.int)
|
||||||
return ([[y, {"pooled_output": []}]], )
|
return ([[y, {}]], )
|
||||||
|
|
||||||
class DiTCondLabelEmpty:
|
class DiTCondLabelEmpty:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -100,7 +100,7 @@ class DiTCondLabelEmpty:
|
|||||||
# [ID of last class + 1] == [num_classes]
|
# [ID of last class + 1] == [num_classes]
|
||||||
y_null = model.model.dit_config["num_classes"]
|
y_null = model.model.dit_config["num_classes"]
|
||||||
y = torch.tensor([[y_null]]).to(torch.int)
|
y = torch.tensor([[y_null]]).to(torch.int)
|
||||||
return ([[y, {"pooled_output": []}]], )
|
return ([[y, {}]], )
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"DitCheckpointLoader" : DitCheckpointLoader,
|
"DitCheckpointLoader" : DitCheckpointLoader,
|
||||||
|
|||||||
Reference in New Issue
Block a user