(self,
dataset_name="human_ml3d",
latent_dim=64,
input_dim=263)
| 61 | class 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()) |
nothing calls this directly
no test coverage detected