| 59 | |
| 60 | |
| 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()) |
| 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 | |
| 118 | class PoseDecoder(nn.Module): |