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:
City
2023-09-06 19:15:29 +02:00
parent 9840963f5a
commit 06256b6269
3 changed files with 116 additions and 21 deletions
+2 -2
View File
@@ -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)