(self, x: torch.Tensor,
hidden_z: Optional[torch.Tensor] = None,
heads_z: Optional[torch.Tensor] = None,
mha_z: Optional[torch.Tensor] = None,
intermediate_z: Optional[torch.Tensor] = None,
ffn_z: Optional[torch.Tensor] = None,
embed_dim_z: Optional[torch.Tensor] = None)
| 491 | self.transformer.set_grad_checkpointing(enable) |
| 492 | |
| 493 | def forward(self, x: torch.Tensor, |
| 494 | hidden_z: Optional[torch.Tensor] = None, |
| 495 | heads_z: Optional[torch.Tensor] = None, |
| 496 | mha_z: Optional[torch.Tensor] = None, |
| 497 | intermediate_z: Optional[torch.Tensor] = None, |
| 498 | ffn_z: Optional[torch.Tensor] = None, |
| 499 | embed_dim_z: Optional[torch.Tensor] = None): |
| 500 | |
| 501 | self.hidden_z = hidden_z |
| 502 | self.embed_dim_z = embed_dim_z |
| 503 | |
| 504 | x = x.to(self.conv1.weight.device) |
| 505 | x = self.conv1(x) # shape = [*, width, grid, grid] |
| 506 | # shape = [*, width, grid ** 2] |
| 507 | x = x.reshape(x.shape[0], x.shape[1], -1) |
| 508 | x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width] |
| 509 | # the first token is the class token. |
| 510 | x = torch.cat( |
| 511 | [self.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device), |
| 512 | x], dim=1) # shape = [*, 1 + grid ** 2, width] |
| 513 | x = x + self.positional_embedding.to(x.dtype) # 128, 50, 768 |
| 514 | |
| 515 | if hidden_z is not None: |
| 516 | x = torch.mul(x, hidden_z) |
| 517 | x = self.ln_pre(x, hidden_z=hidden_z) |
| 518 | |
| 519 | x = x.permute(1, 0, 2) # NLD -> LND 50, 128, 768 |
| 520 | x = self.transformer(x, |
| 521 | hidden_z=hidden_z, |
| 522 | heads_z=heads_z, |
| 523 | mha_z=mha_z, |
| 524 | intermediate_z=intermediate_z, |
| 525 | ffn_z=ffn_z) |
| 526 | |
| 527 | x = x.permute(1, 0, 2) # LND -> NLD |
| 528 | |
| 529 | # select class token |
| 530 | x = self.ln_post(x[:, 0, :], hidden_z=hidden_z) |
| 531 | |
| 532 | if self.proj is not None: |
| 533 | x = self.get_proj_feature(x) |
| 534 | |
| 535 | return x |
| 536 | |
| 537 | def get_proj_feature(self, x): |
| 538 | if self.proj is not None: |
nothing calls this directly
no test coverage detected