Crops positional embeddings for SD3 compatibility.
(self, height, width)
| 529 | raise ValueError(f"Unsupported pos_embed_type: {pos_embed_type}") |
| 530 | |
| 531 | def cropped_pos_embed(self, height, width): |
| 532 | """Crops positional embeddings for SD3 compatibility.""" |
| 533 | if self.pos_embed_max_size is None: |
| 534 | raise ValueError("`pos_embed_max_size` must be set for cropping.") |
| 535 | |
| 536 | height = height // self.patch_size |
| 537 | width = width // self.patch_size |
| 538 | if height > self.pos_embed_max_size: |
| 539 | raise ValueError( |
| 540 | f"Height ({height}) cannot be greater than `pos_embed_max_size`: {self.pos_embed_max_size}." |
| 541 | ) |
| 542 | if width > self.pos_embed_max_size: |
| 543 | raise ValueError( |
| 544 | f"Width ({width}) cannot be greater than `pos_embed_max_size`: {self.pos_embed_max_size}." |
| 545 | ) |
| 546 | |
| 547 | top = (self.pos_embed_max_size - height) // 2 |
| 548 | left = (self.pos_embed_max_size - width) // 2 |
| 549 | spatial_pos_embed = self.pos_embed.reshape(1, self.pos_embed_max_size, self.pos_embed_max_size, -1) |
| 550 | spatial_pos_embed = spatial_pos_embed[:, top : top + height, left : left + width, :] |
| 551 | spatial_pos_embed = spatial_pos_embed.reshape(1, -1, spatial_pos_embed.shape[-1]) |
| 552 | return spatial_pos_embed |
| 553 | |
| 554 | def forward(self, latent): |
| 555 | if self.pos_embed_max_size is not None: |