MCPcopy Create free account
hub / github.com/YesianRohn/TextSSR / __init__

Method __init__

diffusers/src/diffusers/models/embeddings.py:186–235  ·  view source on GitHub ↗
(
        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
    )

Source from the content-addressed store, hash-verified

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."""

Callers

nothing calls this directly

Calls 2

get_2d_sincos_pos_embedFunction · 0.85
__init__Method · 0.45

Tested by

no test coverage detected