Multi label + general jank
This is probably not what the original paper had in mind for "interpolating between classes" but it mostly works.
This commit is contained in:
@@ -258,8 +258,8 @@ class DiT(nn.Module):
|
||||
# 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:]
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user