| 54 | |
| 55 | |
| 56 | def forward_vit(pretrained, x): |
| 57 | b, c, h, w = x.shape |
| 58 | |
| 59 | glob = pretrained.model.forward_flex(x) |
| 60 | |
| 61 | layer_1 = pretrained.activations["1"] |
| 62 | layer_2 = pretrained.activations["2"] |
| 63 | layer_3 = pretrained.activations["3"] |
| 64 | layer_4 = pretrained.activations["4"] |
| 65 | |
| 66 | layer_1 = pretrained.act_postprocess1[0:2](layer_1) |
| 67 | layer_2 = pretrained.act_postprocess2[0:2](layer_2) |
| 68 | layer_3 = pretrained.act_postprocess3[0:2](layer_3) |
| 69 | layer_4 = pretrained.act_postprocess4[0:2](layer_4) |
| 70 | |
| 71 | unflatten = nn.Sequential( |
| 72 | nn.Unflatten( |
| 73 | 2, |
| 74 | torch.Size( |
| 75 | [ |
| 76 | h // pretrained.model.patch_size[1], |
| 77 | w // pretrained.model.patch_size[0], |
| 78 | ] |
| 79 | ), |
| 80 | ) |
| 81 | ) |
| 82 | |
| 83 | if layer_1.ndim == 3: |
| 84 | layer_1 = unflatten(layer_1) |
| 85 | if layer_2.ndim == 3: |
| 86 | layer_2 = unflatten(layer_2) |
| 87 | if layer_3.ndim == 3: |
| 88 | layer_3 = unflatten(layer_3) |
| 89 | if layer_4.ndim == 3: |
| 90 | layer_4 = unflatten(layer_4) |
| 91 | |
| 92 | layer_1 = pretrained.act_postprocess1[3 : len(pretrained.act_postprocess1)](layer_1) |
| 93 | layer_2 = pretrained.act_postprocess2[3 : len(pretrained.act_postprocess2)](layer_2) |
| 94 | layer_3 = pretrained.act_postprocess3[3 : len(pretrained.act_postprocess3)](layer_3) |
| 95 | layer_4 = pretrained.act_postprocess4[3 : len(pretrained.act_postprocess4)](layer_4) |
| 96 | |
| 97 | return layer_1, layer_2, layer_3, layer_4 |
| 98 | |
| 99 | |
| 100 | def _resize_pos_embed(self, posemb, gs_h, gs_w): |