From 1d6c291e752904e08ef96185adeb7034f3391358 Mon Sep 17 00:00:00 2001 From: gchapman Date: Tue, 18 Jun 2024 21:45:10 +0100 Subject: [PATCH] Minor optimization to stop a double .to() being needed --- PixArt/models/PixArtMS.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/PixArt/models/PixArtMS.py b/PixArt/models/PixArtMS.py index 79b7614..957811e 100644 --- a/PixArt/models/PixArtMS.py +++ b/PixArt/models/PixArtMS.py @@ -174,7 +174,7 @@ class PixArtMS(PixArt): self.pos_embed.shape[-1], (self.h, self.w), pe_interpolation=self.pe_interpolation, 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 t = self.t_embedder(timestep) # (N, D) @@ -224,7 +224,7 @@ class PixArtMS(PixArt): device=x.device ).repeat(bs, 1) 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: data_info["aspect_ratio"] = torch.tensor( [[x.shape[2]/x.shape[3]]], @@ -232,7 +232,7 @@ class PixArtMS(PixArt): device=x.device ).repeat(bs, 1) 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 if len(context.shape) == 3: