(self, latent)
| 258 | return spatial_pos_embed |
| 259 | |
| 260 | def forward(self, latent): |
| 261 | if self.pos_embed_max_size is not None: |
| 262 | height, width = latent.shape[-2:] |
| 263 | else: |
| 264 | height, width = latent.shape[-2] // self.patch_size, latent.shape[-1] // self.patch_size |
| 265 | |
| 266 | latent = self.proj(latent) |
| 267 | if self.flatten: |
| 268 | latent = latent.flatten(2).transpose(1, 2) # BCHW -> BNC |
| 269 | if self.layer_norm: |
| 270 | latent = self.norm(latent) |
| 271 | if self.pos_embed is None: |
| 272 | return latent.to(latent.dtype) |
| 273 | # Interpolate or crop positional embeddings as needed |
| 274 | if self.pos_embed_max_size: |
| 275 | pos_embed = self.cropped_pos_embed(height, width) |
| 276 | else: |
| 277 | if self.height != height or self.width != width: |
| 278 | pos_embed = get_2d_sincos_pos_embed( |
| 279 | embed_dim=self.pos_embed.shape[-1], |
| 280 | grid_size=(height, width), |
| 281 | base_size=self.base_size, |
| 282 | interpolation_scale=self.interpolation_scale, |
| 283 | ) |
| 284 | pos_embed = torch.from_numpy(pos_embed).float().unsqueeze(0).to(latent.device) |
| 285 | else: |
| 286 | pos_embed = self.pos_embed |
| 287 | |
| 288 | return (latent + pos_embed).to(latent.dtype) |
| 289 | |
| 290 | |
| 291 | class LuminaPatchEmbed(nn.Module): |
nothing calls this directly
no test coverage detected