| 17 | |
| 18 | class PatchEmbedDust3R(PatchEmbed): |
| 19 | def forward(self, x, **kw): |
| 20 | B, C, H, W = x.shape |
| 21 | assert ( |
| 22 | H % self.patch_size[0] == 0 |
| 23 | ), f"Input image height ({H}) is not a multiple of patch size ({self.patch_size[0]})." |
| 24 | assert ( |
| 25 | W % self.patch_size[1] == 0 |
| 26 | ), f"Input image width ({W}) is not a multiple of patch size ({self.patch_size[1]})." |
| 27 | x = self.proj(x) |
| 28 | pos = self.position_getter(B, x.size(2), x.size(3), x.device) |
| 29 | if self.flatten: |
| 30 | x = x.flatten(2).transpose(1, 2) # BCHW -> BNC |
| 31 | x = self.norm(x) |
| 32 | return x, pos |
| 33 | |
| 34 | |
| 35 | class ManyAR_PatchEmbed(PatchEmbed): |