MCPcopy Create free account
hub / github.com/GuyTevet/motion-diffusion-model / forward

Method forward

model/mdm.py:343–357  ·  view source on GitHub ↗
(self, x)

Source from the content-addressed store, hash-verified

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
360class OutputProcess(nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected