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

Method __init__

mogen/models/transformers/finemogen.py:63–101  ·  view source on GitHub ↗
(self,
                 dataset_name="human_ml3d",
                 latent_dim=64,
                 input_dim=263)

Source from the content-addressed store, hash-verified

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())

Callers

nothing calls this directly

Calls 2

get_part_sliceFunction · 0.85
__init__Method · 0.45

Tested by

no test coverage detected