(self,
dataset_name="human_ml3d",
latent_dim=64,
output_dim=263)
| 118 | class PoseDecoder(nn.Module): |
| 119 | |
| 120 | def __init__(self, |
| 121 | dataset_name="human_ml3d", |
| 122 | latent_dim=64, |
| 123 | output_dim=263): |
| 124 | super().__init__() |
| 125 | self.dataset_name = dataset_name |
| 126 | self.latent_dim = latent_dim |
| 127 | self.output_dim = output_dim |
| 128 | if dataset_name == "human_ml3d": |
| 129 | func = get_t2m_slice |
| 130 | self.head_slice = get_part_slice([12, 15], func) |
| 131 | self.stem_slice = get_part_slice([3, 6, 9], func) |
| 132 | self.larm_slice = get_part_slice([14, 17, 19, 21], func) |
| 133 | self.rarm_slice = get_part_slice([13, 16, 18, 20], func) |
| 134 | self.lleg_slice = get_part_slice([2, 5, 8, 11], func) |
| 135 | self.rleg_slice = get_part_slice([1, 4, 7, 10], func) |
| 136 | self.root_slice = get_part_slice([0], func) |
| 137 | self.body_slice = get_part_slice([_ for _ in range(22)], func) |
| 138 | elif dataset_name == "kit_ml": |
| 139 | func = get_kit_slice |
| 140 | self.head_slice = get_part_slice([4], func) |
| 141 | self.stem_slice = get_part_slice([1, 2, 3], func) |
| 142 | self.larm_slice = get_part_slice([8, 9, 10], func) |
| 143 | self.rarm_slice = get_part_slice([5, 6, 7], func) |
| 144 | self.lleg_slice = get_part_slice([16, 17, 18, 19, 20], func) |
| 145 | self.rleg_slice = get_part_slice([11, 12, 13, 14, 15], func) |
| 146 | self.root_slice = get_part_slice([0], func) |
| 147 | self.body_slice = get_part_slice([_ for _ in range(21)], func) |
| 148 | else: |
| 149 | raise ValueError() |
| 150 | |
| 151 | self.head_out = nn.Linear(latent_dim, len(self.head_slice)) |
| 152 | self.stem_out = nn.Linear(latent_dim, len(self.stem_slice)) |
| 153 | self.larm_out = nn.Linear(latent_dim, len(self.larm_slice)) |
| 154 | self.rarm_out = nn.Linear(latent_dim, len(self.rarm_slice)) |
| 155 | self.lleg_out = nn.Linear(latent_dim, len(self.lleg_slice)) |
| 156 | self.rleg_out = nn.Linear(latent_dim, len(self.rleg_slice)) |
| 157 | self.root_out = nn.Linear(latent_dim, len(self.root_slice)) |
| 158 | self.body_out = nn.Linear(latent_dim, len(self.body_slice)) |
| 159 | |
| 160 | def forward(self, motion): |
| 161 | B, T = motion.shape[:2] |
nothing calls this directly
no test coverage detected