(self, latent)
| 552 | return spatial_pos_embed |
| 553 | |
| 554 | def forward(self, latent): |
| 555 | if self.pos_embed_max_size is not None: |
| 556 | height, width = latent.shape[-2:] |
| 557 | else: |
| 558 | height, width = latent.shape[-2] // self.patch_size, latent.shape[-1] // self.patch_size |
| 559 | latent = self.proj(latent) |
| 560 | if self.flatten: |
| 561 | latent = latent.flatten(2).transpose(1, 2) # BCHW -> BNC |
| 562 | if self.layer_norm: |
| 563 | latent = self.norm(latent) |
| 564 | if self.pos_embed is None: |
| 565 | return latent.to(latent.dtype) |
| 566 | # Interpolate or crop positional embeddings as needed |
| 567 | if self.pos_embed_max_size: |
| 568 | pos_embed = self.cropped_pos_embed(height, width) |
| 569 | else: |
| 570 | if self.height != height or self.width != width: |
| 571 | pos_embed = get_2d_sincos_pos_embed( |
| 572 | embed_dim=self.pos_embed.shape[-1], |
| 573 | grid_size=(height, width), |
| 574 | base_size=self.base_size, |
| 575 | interpolation_scale=self.interpolation_scale, |
| 576 | device=latent.device, |
| 577 | output_type="pt", |
| 578 | ) |
| 579 | pos_embed = pos_embed.float().unsqueeze(0) |
| 580 | else: |
| 581 | pos_embed = self.pos_embed |
| 582 | |
| 583 | return (latent + pos_embed).to(latent.dtype) |
| 584 | |
| 585 | |
| 586 | class LuminaPatchEmbed(nn.Module): |
nothing calls this directly
no test coverage detected