MCPcopy Create free account
hub / github.com/MotrixLab/FineMoGen / PoseEncoder

Class PoseEncoder

mogen/models/transformers/finemogen.py:61–115  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

59
60
61class PoseEncoder(nn.Module):
62
63 def __init__(self,
64 dataset_name="human_ml3d",
65 latent_dim=64,
66 input_dim=263):
67 super().__init__()
68 self.dataset_name = dataset_name
69 if dataset_name == "human_ml3d":
70 func = get_t2m_slice
71 self.head_slice = get_part_slice([12, 15], func)
72 self.stem_slice = get_part_slice([3, 6, 9], func)
73 self.larm_slice = get_part_slice([14, 17, 19, 21], func)
74 self.rarm_slice = get_part_slice([13, 16, 18, 20], func)
75 self.lleg_slice = get_part_slice([2, 5, 8, 11], func)
76 self.rleg_slice = get_part_slice([1, 4, 7, 10], func)
77 self.root_slice = get_part_slice([0], func)
78 self.body_slice = get_part_slice([_ for _ in range(22)], func)
79 elif dataset_name == "kit_ml":
80 func = get_kit_slice
81 self.head_slice = get_part_slice([4], func)
82 self.stem_slice = get_part_slice([1, 2, 3], func)
83 self.larm_slice = get_part_slice([8, 9, 10], func)
84 self.rarm_slice = get_part_slice([5, 6, 7], func)
85 self.lleg_slice = get_part_slice([16, 17, 18, 19, 20], func)
86 self.rleg_slice = get_part_slice([11, 12, 13, 14, 15], func)
87 self.root_slice = get_part_slice([0], func)
88 self.body_slice = get_part_slice([_ for _ in range(21)], func)
89 else:
90 raise ValueError()
91
92 self.head_embed = nn.Linear(len(self.head_slice), latent_dim)
93 self.stem_embed = nn.Linear(len(self.stem_slice), latent_dim)
94 self.larm_embed = nn.Linear(len(self.larm_slice), latent_dim)
95 self.rarm_embed = nn.Linear(len(self.rarm_slice), latent_dim)
96 self.lleg_embed = nn.Linear(len(self.lleg_slice), latent_dim)
97 self.rleg_embed = nn.Linear(len(self.rleg_slice), latent_dim)
98 self.root_embed = nn.Linear(len(self.root_slice), latent_dim)
99 self.body_embed = nn.Linear(len(self.body_slice), latent_dim)
100
101 assert len(set(self.body_slice)) == input_dim
102
103 def forward(self, motion):
104 head_feat = self.head_embed(motion[:, :, self.head_slice].contiguous())
105 stem_feat = self.stem_embed(motion[:, :, self.stem_slice].contiguous())
106 larm_feat = self.larm_embed(motion[:, :, self.larm_slice].contiguous())
107 rarm_feat = self.rarm_embed(motion[:, :, self.rarm_slice].contiguous())
108 lleg_feat = self.lleg_embed(motion[:, :, self.lleg_slice].contiguous())
109 rleg_feat = self.rleg_embed(motion[:, :, self.rleg_slice].contiguous())
110 root_feat = self.root_embed(motion[:, :, self.root_slice].contiguous())
111 body_feat = self.body_embed(motion[:, :, self.body_slice].contiguous())
112 feat = torch.cat((head_feat, stem_feat, larm_feat, rarm_feat,
113 lleg_feat, rleg_feat, root_feat, body_feat),
114 dim=-1)
115 return feat
116
117
118class PoseDecoder(nn.Module):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected