PixArt allow on-the-fly PE interpolation
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user