(
self,
height=224,
width=224,
patch_size=16,
in_channels=3,
embed_dim=768,
layer_norm=False,
flatten=True,
bias=True,
interpolation_scale=1,
pos_embed_type="sincos",
pos_embed_max_size=None, # For SD3 cropping
)
| 184 | """2D Image to Patch Embedding with support for SD3 cropping.""" |
| 185 | |
| 186 | def __init__( |
| 187 | self, |
| 188 | height=224, |
| 189 | width=224, |
| 190 | patch_size=16, |
| 191 | in_channels=3, |
| 192 | embed_dim=768, |
| 193 | layer_norm=False, |
| 194 | flatten=True, |
| 195 | bias=True, |
| 196 | interpolation_scale=1, |
| 197 | pos_embed_type="sincos", |
| 198 | pos_embed_max_size=None, # For SD3 cropping |
| 199 | ): |
| 200 | super().__init__() |
| 201 | |
| 202 | num_patches = (height // patch_size) * (width // patch_size) |
| 203 | self.flatten = flatten |
| 204 | self.layer_norm = layer_norm |
| 205 | self.pos_embed_max_size = pos_embed_max_size |
| 206 | |
| 207 | self.proj = nn.Conv2d( |
| 208 | in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias |
| 209 | ) |
| 210 | if layer_norm: |
| 211 | self.norm = nn.LayerNorm(embed_dim, elementwise_affine=False, eps=1e-6) |
| 212 | else: |
| 213 | self.norm = None |
| 214 | |
| 215 | self.patch_size = patch_size |
| 216 | self.height, self.width = height // patch_size, width // patch_size |
| 217 | self.base_size = height // patch_size |
| 218 | self.interpolation_scale = interpolation_scale |
| 219 | |
| 220 | # Calculate positional embeddings based on max size or default |
| 221 | if pos_embed_max_size: |
| 222 | grid_size = pos_embed_max_size |
| 223 | else: |
| 224 | grid_size = int(num_patches**0.5) |
| 225 | |
| 226 | if pos_embed_type is None: |
| 227 | self.pos_embed = None |
| 228 | elif pos_embed_type == "sincos": |
| 229 | pos_embed = get_2d_sincos_pos_embed( |
| 230 | embed_dim, grid_size, base_size=self.base_size, interpolation_scale=self.interpolation_scale |
| 231 | ) |
| 232 | persistent = True if pos_embed_max_size else False |
| 233 | self.register_buffer("pos_embed", torch.from_numpy(pos_embed).float().unsqueeze(0), persistent=persistent) |
| 234 | else: |
| 235 | raise ValueError(f"Unsupported pos_embed_type: {pos_embed_type}") |
| 236 | |
| 237 | def cropped_pos_embed(self, height, width): |
| 238 | """Crops positional embeddings for SD3 compatibility.""" |
nothing calls this directly
no test coverage detected