From b31d8c517b732bc630a31e9770dbcaf9c1925fbb Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Wed, 6 Dec 2023 21:31:07 +0100 Subject: [PATCH] PixArt small fix Issue #3 + some code changes from upstream. --- PixArt/models/PixArt.py | 4 +++- PixArt/models/PixArtMS.py | 4 +++- PixArt/sampler.py | 2 +- 3 files changed, 7 insertions(+), 3 deletions(-) diff --git a/PixArt/models/PixArt.py b/PixArt/models/PixArt.py index 7f106fe..4428ab7 100644 --- a/PixArt/models/PixArt.py +++ b/PixArt/models/PixArt.py @@ -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() diff --git a/PixArt/models/PixArtMS.py b/PixArt/models/PixArtMS.py index 6c1dce3..a9ef67b 100644 --- a/PixArt/models/PixArtMS.py +++ b/PixArt/models/PixArtMS.py @@ -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() diff --git a/PixArt/sampler.py b/PixArt/sampler.py index ae9e656..fa2f39c 100644 --- a/PixArt/sampler.py +++ b/PixArt/sampler.py @@ -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