@@ -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()
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user