# Codeblock 7b def forward(self, x): print(f'original\t\t: {x.size()}') x = self.patcher(x) #(1) print(f'after patcher\t\t: {x.size()}') x = torch.cat([self.class_token, self.dist_token, x], dim=1) #(2) print(f'after concat\t\t: {x.size()}') x = x + self.pos_embedding #(3) print(f'after pos embed\t\t: {x.size()}') for i, encoder in enumerate(self.encoders): x = encoder(x) #(4) print(f"after encoder #{i}\t: {x.size()}") x = self.norm_out(x) #(5) print(f'after norm\t\t: {x.size()}') class_out = x[:, 0] #(6) print(f'class_out\t\t: {class_out.size()}') dist_out = x[:, 1] #(7) print(f'dist_out\t\t: {dist_out.size()}') class_out = self.class_head(class_out) #(8) print(f'after class_head\t: {class_out.size()}') dist_out = self.dist_head(dist_out) #(9) print(f'after dist_head\t\t: {class_out.size()}') return class_out, dist_out