| 130 | |
| 131 | |
| 132 | def decode(self, hidden, N, H, W): |
| 133 | BN, hw, _ = hidden.shape |
| 134 | B = BN // N |
| 135 | |
| 136 | final_output = [] |
| 137 | |
| 138 | hidden = hidden.reshape(B*N, hw, -1) |
| 139 | |
| 140 | register_token = self.register_token.repeat(B, N, 1, 1).reshape(B*N, *self.register_token.shape[-2:]) |
| 141 | |
| 142 | # Concatenate special tokens with patch tokens |
| 143 | hidden = torch.cat([register_token, hidden], dim=1) |
| 144 | hw = hidden.shape[1] |
| 145 | |
| 146 | if self.pos_type.startswith('rope'): |
| 147 | pos = self.position_getter(B * N, H//self.patch_size, W//self.patch_size, hidden.device) |
| 148 | |
| 149 | if self.patch_start_idx > 0: |
| 150 | # do not use position embedding for special tokens (camera and register tokens) |
| 151 | # so set pos to 0 for the special tokens |
| 152 | pos = pos + 1 |
| 153 | pos_special = torch.zeros(B * N, self.patch_start_idx, 2).to(hidden.device).to(pos.dtype) |
| 154 | pos = torch.cat([pos_special, pos], dim=1) |
| 155 | |
| 156 | for i in range(len(self.decoder)): |
| 157 | blk = self.decoder[i] |
| 158 | |
| 159 | if i % 2 == 0: |
| 160 | pos = pos.reshape(B*N, hw, -1) |
| 161 | hidden = hidden.reshape(B*N, hw, -1) |
| 162 | else: |
| 163 | pos = pos.reshape(B, N*hw, -1) |
| 164 | hidden = hidden.reshape(B, N*hw, -1) |
| 165 | |
| 166 | hidden = blk(hidden, xpos=pos) |
| 167 | |
| 168 | if i+1 in [len(self.decoder)-1, len(self.decoder)]: |
| 169 | final_output.append(hidden.reshape(B*N, hw, -1)) |
| 170 | |
| 171 | return torch.cat([final_output[0], final_output[1]], dim=-1), pos.reshape(B*N, hw, -1) |
| 172 | |
| 173 | def forward(self, imgs): |
| 174 | imgs = (imgs - self.image_mean) / self.image_std |