(self, latent)
| 538 | return spatial_pos_embed |
| 539 | |
| 540 | def forward(self, latent): |
| 541 | if self.pos_embed_max_size is not None: |
| 542 | height, width = latent.shape[-2:] |
| 543 | else: |
| 544 | height, width = latent.shape[-2] // self.patch_size, latent.shape[-1] // self.patch_size |
| 545 | latent = self.proj(latent) |
| 546 | if self.flatten: |
| 547 | latent = latent.flatten(2).transpose(1, 2) # BCHW -> BNC |
| 548 | if self.layer_norm: |
| 549 | latent = self.norm(latent) |
| 550 | if self.pos_embed is None: |
| 551 | return latent.to(latent.dtype) |
| 552 | # Interpolate or crop positional embeddings as needed |
| 553 | if self.pos_embed_max_size: |
| 554 | pos_embed = self.cropped_pos_embed(height, width) |
| 555 | else: |
| 556 | if self.height != height or self.width != width: |
| 557 | pos_embed = get_2d_sincos_pos_embed( |
| 558 | embed_dim=self.pos_embed.shape[-1], |
| 559 | grid_size=(height, width), |
| 560 | base_size=self.base_size, |
| 561 | interpolation_scale=self.interpolation_scale, |
| 562 | device=latent.device, |
| 563 | output_type="pt", |
| 564 | ) |
| 565 | pos_embed = pos_embed.float().unsqueeze(0) |
| 566 | else: |
| 567 | pos_embed = self.pos_embed |
| 568 | |
| 569 | return (latent + pos_embed).to(latent.dtype) |
| 570 | |
| 571 | |
| 572 | class LuminaPatchEmbed(nn.Module): |
nothing calls this directly
no test coverage detected