(self, x)
| 129 | nn.init.constant_(m.weight, 1.0) |
| 130 | |
| 131 | def patchify(self, x): |
| 132 | bsz, c, h, w = x.shape |
| 133 | p = self.patch_size |
| 134 | h_, w_ = h // p, w // p |
| 135 | |
| 136 | x = x.reshape(bsz, c, h_, p, w_, p) |
| 137 | x = torch.einsum('nchpwq->nhwcpq', x) |
| 138 | x = x.reshape(bsz, h_ * w_, c * p ** 2) |
| 139 | return x # [n, l, d] |
| 140 | |
| 141 | def unpatchify(self, x): |
| 142 | bsz = x.shape[0] |