(self, x)
| 341 | self.velEmbedding = nn.Linear(self.input_feats, self.latent_dim) |
| 342 | |
| 343 | def forward(self, x): |
| 344 | bs, njoints, nfeats, nframes = x.shape |
| 345 | x = x.permute((3, 0, 1, 2)).reshape(nframes, bs, njoints*nfeats) |
| 346 | |
| 347 | if self.data_rep in ['rot6d', 'xyz', 'hml_vec']: |
| 348 | x = self.poseEmbedding(x) # [seqlen, bs, d] |
| 349 | return x |
| 350 | elif self.data_rep == 'rot_vel': |
| 351 | first_pose = x[[0]] # [1, bs, 150] |
| 352 | first_pose = self.poseEmbedding(first_pose) # [1, bs, d] |
| 353 | vel = x[1:] # [seqlen-1, bs, 150] |
| 354 | vel = self.velEmbedding(vel) # [seqlen-1, bs, d] |
| 355 | return torch.cat((first_pose, vel), axis=0) # [seqlen, bs, d] |
| 356 | else: |
| 357 | raise ValueError |
| 358 | |
| 359 | |
| 360 | class OutputProcess(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected