Minor optimization to stop a double .to() being needed

This commit is contained in:
gchapman
2024-06-18 21:45:10 +01:00
parent 03a81eddd5
commit 1d6c291e75
+3 -3
View File
@@ -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: