PixArt allow on-the-fly PE interpolation

This commit is contained in:
City
2024-06-19 02:26:01 +02:00
parent 5101719ca8
commit 2ad0b09cb9
2 changed files with 11 additions and 2 deletions
+2
View File
@@ -72,6 +72,7 @@ class PixArt(nn.Module):
drop_path: float = 0., drop_path: float = 0.,
caption_channels=4096, caption_channels=4096,
pe_interpolation=1.0, pe_interpolation=1.0,
pe_precision=None,
config=None, config=None,
model_max_length=120, model_max_length=120,
qk_norm=False, qk_norm=False,
@@ -85,6 +86,7 @@ class PixArt(nn.Module):
self.patch_size = patch_size self.patch_size = patch_size
self.num_heads = num_heads self.num_heads = num_heads
self.pe_interpolation = pe_interpolation self.pe_interpolation = pe_interpolation
self.pe_precision = pe_precision
self.depth = depth self.depth = depth
self.x_embedder = PatchEmbed(input_size, patch_size, in_channels, hidden_size, bias=True) self.x_embedder = PatchEmbed(input_size, patch_size, in_channels, hidden_size, bias=True)
+9 -2
View File
@@ -98,7 +98,8 @@ class PixArtMS(PixArt):
pred_sigma=True, pred_sigma=True,
drop_path: float = 0., drop_path: float = 0.,
caption_channels=4096, caption_channels=4096,
pe_interpolation=1., pe_interpolation=None,
pe_precision=None,
config=None, config=None,
model_max_length=120, model_max_length=120,
micro_condition=True, micro_condition=True,
@@ -168,10 +169,16 @@ class PixArtMS(PixArt):
x = x.to(self.dtype) x = x.to(self.dtype)
timestep = t.to(self.dtype) timestep = t.to(self.dtype)
y = y.to(self.dtype) y = y.to(self.dtype)
pe_interpolation = self.pe_interpolation
if pe_interpolation is None or self.pe_precision is not None:
# calculate pe_interpolation on-the-fly
pe_interpolation = round((x.shape[-1]+x.shape[-2])/2.0 / (512/8.0), self.pe_precision or 0)
self.h, self.w = x.shape[-2]//self.patch_size, x.shape[-1]//self.patch_size self.h, self.w = x.shape[-2]//self.patch_size, x.shape[-1]//self.patch_size
pos_embed = torch.from_numpy( pos_embed = torch.from_numpy(
get_2d_sincos_pos_embed( get_2d_sincos_pos_embed(
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=pe_interpolation,
base_size=self.base_size base_size=self.base_size
) )
).unsqueeze(0).to(x.device).to(self.dtype) ).unsqueeze(0).to(x.device).to(self.dtype)