MCPcopy Create free account
hub / github.com/Pixel-Talk/EHM-Tracker / _create_model

Method _create_model

src/modules/pixie/pixie.py:71–189  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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

Callers 1

__init__Method · 0.95

Calls 11

SMPLXClass · 0.90
ResnetEncoderClass · 0.85
HRNEncoderClass · 0.85
MLPClass · 0.85
TempSoftmaxFusionClass · 0.85
existsMethod · 0.80
loadMethod · 0.80
evalMethod · 0.80
keysMethod · 0.45
getMethod · 0.45
toMethod · 0.45

Tested by

no test coverage detected