(self, x)
| 966 | return {} |
| 967 | |
| 968 | def forward(self, x): |
| 969 | x = x[0] |
| 970 | x = self.patch_embed(x) |
| 971 | |
| 972 | T = self.cfg.DATA.NUM_FRAMES // self.patch_stride[0] |
| 973 | H = self.cfg.DATA.TRAIN_CROP_SIZE // self.patch_stride[1] |
| 974 | W = self.cfg.DATA.TRAIN_CROP_SIZE // self.patch_stride[2] |
| 975 | B, N, C = x.shape |
| 976 | |
| 977 | if self.cls_embed_on: |
| 978 | cls_tokens = self.cls_token.expand( |
| 979 | B, -1, -1 |
| 980 | ) # stole cls_tokens impl from Phil Wang, thanks |
| 981 | x = torch.cat((cls_tokens, x), dim=1) |
| 982 | |
| 983 | if self.sep_pos_embed: |
| 984 | pos_embed = self.pos_embed_spatial.repeat( |
| 985 | 1, self.patch_dims[0], 1 |
| 986 | ) + torch.repeat_interleave( |
| 987 | self.pos_embed_temporal, |
| 988 | self.patch_dims[1] * self.patch_dims[2], |
| 989 | dim=1, |
| 990 | ) |
| 991 | pos_embed_cls = torch.cat([self.pos_embed_class, pos_embed], 1) |
| 992 | x = x + pos_embed_cls |
| 993 | else: |
| 994 | x = x + self.pos_embed |
| 995 | |
| 996 | if self.drop_rate: |
| 997 | x = self.pos_drop(x) |
| 998 | |
| 999 | if self.norm_stem: |
| 1000 | x = self.norm_stem(x) |
| 1001 | |
| 1002 | thw = [T, H, W] |
| 1003 | for blk in self.blocks: |
| 1004 | x, thw = blk(x, thw) |
| 1005 | |
| 1006 | x = self.norm(x) |
| 1007 | if self.cls_embed_on: |
| 1008 | x = x[:, 0] |
| 1009 | else: |
| 1010 | x = x.mean(1) |
| 1011 | |
| 1012 | x = self.head(x) |
| 1013 | return x |
nothing calls this directly
no outgoing calls
no test coverage detected