Minor optimization to stop a double .to() being needed
This commit is contained in:
@@ -174,7 +174,7 @@ class PixArtMS(PixArt):
|
|||||||
self.pos_embed.shape[-1], (self.h, self.w), pe_interpolation=self.pe_interpolation,
|
self.pos_embed.shape[-1], (self.h, self.w), pe_interpolation=self.pe_interpolation,
|
||||||
base_size=self.base_size
|
base_size=self.base_size
|
||||||
)
|
)
|
||||||
).unsqueeze(0).to(x.device).to(self.dtype)
|
).unsqueeze(0).to(device=x.device, dtype=self.dtype)
|
||||||
|
|
||||||
x = self.x_embedder(x) + pos_embed # (N, T, D), where T = H * W / patch_size ** 2
|
x = self.x_embedder(x) + pos_embed # (N, T, D), where T = H * W / patch_size ** 2
|
||||||
t = self.t_embedder(timestep) # (N, D)
|
t = self.t_embedder(timestep) # (N, D)
|
||||||
@@ -224,7 +224,7 @@ class PixArtMS(PixArt):
|
|||||||
device=x.device
|
device=x.device
|
||||||
).repeat(bs, 1)
|
).repeat(bs, 1)
|
||||||
else:
|
else:
|
||||||
data_info["img_hw"] = img_hw.to(x.dtype).to(x.device)
|
data_info["img_hw"] = img_hw.to(dtype=x.dtype, device=x.device)
|
||||||
if aspect_ratio is None or True:
|
if aspect_ratio is None or True:
|
||||||
data_info["aspect_ratio"] = torch.tensor(
|
data_info["aspect_ratio"] = torch.tensor(
|
||||||
[[x.shape[2]/x.shape[3]]],
|
[[x.shape[2]/x.shape[3]]],
|
||||||
@@ -232,7 +232,7 @@ class PixArtMS(PixArt):
|
|||||||
device=x.device
|
device=x.device
|
||||||
).repeat(bs, 1)
|
).repeat(bs, 1)
|
||||||
else:
|
else:
|
||||||
data_info["aspect_ratio"] = aspect_ratio.to(x.dtype).to(x.device)
|
data_info["aspect_ratio"] = aspect_ratio.to(dtype=x.dtype, device=x.device)
|
||||||
|
|
||||||
## Still accepts the input w/o that dim but returns garbage
|
## Still accepts the input w/o that dim but returns garbage
|
||||||
if len(context.shape) == 3:
|
if len(context.shape) == 3:
|
||||||
|
|||||||
Reference in New Issue
Block a user