(self)
| 69 | trans_scale=0) |
| 70 | |
| 71 | def _create_model(self): |
| 72 | self.model_dict = {} |
| 73 | # Build all image encoders |
| 74 | # Hand encoder only works for right hand, for left hand, flip inputs and flip the results back |
| 75 | self.Encoder = {} |
| 76 | for key in self.cfg.network.encoder.keys(): |
| 77 | if self.cfg.network.encoder.get(key).type == 'resnet50': |
| 78 | self.Encoder[key] = ResnetEncoder().to(self.device) |
| 79 | elif self.cfg.network.encoder.get(key).type == 'hrnet': |
| 80 | self.Encoder[key] = HRNEncoder().to(self.device) |
| 81 | self.model_dict[f'Encoder_{key}'] = self.Encoder[key].state_dict() |
| 82 | |
| 83 | # Build the parameter regressors |
| 84 | self.Regressor = {} |
| 85 | for key in self.cfg.network.regressor.keys(): |
| 86 | n_output = sum(self.param_list_dict[f'{key}_list'].values()) |
| 87 | channels = [2048] + \ |
| 88 | self.cfg.network.regressor.get(key).channels + [n_output] |
| 89 | if self.cfg.network.regressor.get(key).type == 'mlp': |
| 90 | self.Regressor[key] = MLP(channels=channels).to(self.device) |
| 91 | self.model_dict[f'Regressor_{key}'] = self.Regressor[key].state_dict( |
| 92 | ) |
| 93 | |
| 94 | # Build the extractors |
| 95 | # to extract separate head/left hand/right hand feature from body feature |
| 96 | self.Extractor = {} |
| 97 | for key in self.cfg.network.extractor.keys(): |
| 98 | channels = [2048] + \ |
| 99 | self.cfg.network.extractor.get(key).channels + [2048] |
| 100 | if self.cfg.network.extractor.get(key).type == 'mlp': |
| 101 | self.Extractor[key] = MLP(channels=channels).to(self.device) |
| 102 | self.model_dict[f'Extractor_{key}'] = self.Extractor[key].state_dict( |
| 103 | ) |
| 104 | |
| 105 | # Build the moderators |
| 106 | self.Moderator = {} |
| 107 | for key in self.cfg.network.moderator.keys(): |
| 108 | share_part = key.split('_')[0] |
| 109 | detach_inputs = self.cfg.network.moderator.get(key).detach_inputs |
| 110 | detach_feature = self.cfg.network.moderator.get(key).detach_feature |
| 111 | channels = [2048*2] + \ |
| 112 | self.cfg.network.moderator.get(key).channels + [2] |
| 113 | self.Moderator[key] = TempSoftmaxFusion( |
| 114 | detach_inputs=detach_inputs, detach_feature=detach_feature, |
| 115 | channels=channels).to(self.device) |
| 116 | self.model_dict[f'Moderator_{key}'] = self.Moderator[key].state_dict( |
| 117 | ) |
| 118 | |
| 119 | class JointMapper(nn.Module): |
| 120 | def __init__(self, joint_maps=None): |
| 121 | super(JointMapper, self).__init__() |
| 122 | if joint_maps is None: |
| 123 | self.joint_maps = joint_maps |
| 124 | else: |
| 125 | self.register_buffer('joint_maps', |
| 126 | torch.tensor(joint_maps, dtype=torch.long)) |
| 127 | |
| 128 | def forward(self, joints, **kwargs): |
no test coverage detected