| 288 | """ |
| 289 | |
| 290 | def __init__( |
| 291 | self, |
| 292 | height: int = 224, |
| 293 | width: int = 224, |
| 294 | patch_size: int = 16, |
| 295 | in_channels: int = 3, |
| 296 | embed_dim: int = 768, |
| 297 | layer_norm: bool = False, |
| 298 | flatten: bool = True, |
| 299 | bias: bool = True, |
| 300 | interpolation_scale: int = 1, |
| 301 | pos_embed_type: str = "sincos", |
| 302 | pos_embed_max_size: Optional[int] = None, # For SD3 cropping |
| 303 | dtype=None): |
| 304 | from diffusers.models.embeddings import \ |
| 305 | get_2d_sincos_pos_embed as get_2d_sincos_pos_embed_torch |
| 306 | |
| 307 | from .conv import Conv2d |
| 308 | from .normalization import LayerNorm |
| 309 | |
| 310 | super().__init__() |
| 311 | |
| 312 | num_patches = (height // patch_size) * (width // patch_size) |
| 313 | self.flatten = flatten |
| 314 | self.layer_norm = layer_norm |
| 315 | self.pos_embed_max_size = pos_embed_max_size |
| 316 | |
| 317 | self.proj = Conv2d(in_channels, |
| 318 | embed_dim, |
| 319 | kernel_size=(patch_size, patch_size), |
| 320 | stride=(patch_size, patch_size), |
| 321 | bias=bias, |
| 322 | dtype=dtype) |
| 323 | if layer_norm: |
| 324 | self.norm = LayerNorm(embed_dim, |
| 325 | elementwise_affine=False, |
| 326 | eps=1e-6, |
| 327 | dtype=dtype) |
| 328 | else: |
| 329 | self.norm = None |
| 330 | |
| 331 | self.patch_size = patch_size |
| 332 | self.height, self.width = height // patch_size, width // patch_size |
| 333 | self.base_size = height // patch_size |
| 334 | self.interpolation_scale = interpolation_scale |
| 335 | |
| 336 | # Calculate positional embeddings based on max size or default |
| 337 | if pos_embed_max_size: |
| 338 | grid_size = pos_embed_max_size |
| 339 | else: |
| 340 | grid_size = int(num_patches**0.5) |
| 341 | |
| 342 | if pos_embed_type is None: |
| 343 | self.pos_embed = None |
| 344 | elif pos_embed_type == "sincos": |
| 345 | pos_embed = get_2d_sincos_pos_embed_torch( |
| 346 | embed_dim, |
| 347 | grid_size, |