@@ -128,7 +128,9 @@ class PixArt(nn.Module):
|
|||||||
t = self.t_embedder(t) # (N, D)
|
t = self.t_embedder(t) # (N, D)
|
||||||
t0 = self.t_block(t)
|
t0 = self.t_block(t)
|
||||||
y = self.y_embedder(y, self.training) # (N, 1, L, D)
|
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)
|
mask = mask.squeeze(1).squeeze(1)
|
||||||
y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1])
|
y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1])
|
||||||
y_lens = mask.sum(dim=1).tolist()
|
y_lens = mask.sum(dim=1).tolist()
|
||||||
|
|||||||
@@ -167,7 +167,9 @@ class PixArtMS(PixArt):
|
|||||||
t = t + torch.cat([csize, ar], dim=1)
|
t = t + torch.cat([csize, ar], dim=1)
|
||||||
t0 = self.t_block(t)
|
t0 = self.t_block(t)
|
||||||
y = self.y_embedder(y, self.training) # (N, D)
|
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)
|
mask = mask.squeeze(1).squeeze(1)
|
||||||
y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1])
|
y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1])
|
||||||
y_lens = mask.sum(dim=1).tolist()
|
y_lens = mask.sum(dim=1).tolist()
|
||||||
|
|||||||
+1
-1
@@ -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)
|
previewer = latent_preview.get_previewer(model.load_device, model.model.latent_format)
|
||||||
|
|
||||||
## Noise schedule.
|
## 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)
|
noise_schedule = NoiseScheduleVP(schedule=noise_schedule_vp, betas=betas)
|
||||||
|
|
||||||
## Convert your discrete-time `model` to the continuous-time
|
## Convert your discrete-time `model` to the continuous-time
|
||||||
|
|||||||
Reference in New Issue
Block a user