PixArt small fix

Issue #3 + some code changes from upstream.
This commit is contained in:
City
2023-12-06 21:31:07 +01:00
parent 389c16f2f5
commit b31d8c517b
3 changed files with 7 additions and 3 deletions
+3 -1
View File
@@ -128,7 +128,9 @@ class PixArt(nn.Module):
t = self.t_embedder(t) # (N, D)
t0 = self.t_block(t)
y = self.y_embedder(y, self.training) # (N, 1, L, D)
if self.training:
if mask is not None:
if mask.shape[0] != y.shape[0]:
mask = mask.repeat(y.shape[0] // mask.shape[0], 1)
mask = mask.squeeze(1).squeeze(1)
y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1])
y_lens = mask.sum(dim=1).tolist()
+3 -1
View File
@@ -167,7 +167,9 @@ class PixArtMS(PixArt):
t = t + torch.cat([csize, ar], dim=1)
t0 = self.t_block(t)
y = self.y_embedder(y, self.training) # (N, D)
if self.training:
if mask is not None:
if mask.shape[0] != y.shape[0]:
mask = mask.repeat(y.shape[0] // mask.shape[0], 1)
mask = mask.squeeze(1).squeeze(1)
y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1])
y_lens = mask.sum(dim=1).tolist()
+1 -1
View File
@@ -38,7 +38,7 @@ def sample_pixart(model, seed, steps, cfg, noise_schedule, noise_schedule_vp, po
previewer = latent_preview.get_previewer(model.load_device, model.model.latent_format)
## Noise schedule.
betas = torch.tensor(gd.get_named_beta_schedule(noise_schedule, steps))
betas = torch.tensor(gd.get_named_beta_schedule(noise_schedule, 1000))
noise_schedule = NoiseScheduleVP(schedule=noise_schedule_vp, betas=betas)
## Convert your discrete-time `model` to the continuous-time