MCPcopy Create free account
hub / github.com/MotrixLab/ADHMR / get_model_score

Function get_model_score

ADHMR/lib/utils/function.py:216–268  ·  view source on GitHub ↗
(config, is_train = True, resume = False, resume_path = None)

Source from the content-addressed store, hash-verified

214 return model
215
216def get_model_score(config, is_train = True, resume = False, resume_path = None):
217 neighbour_matrix = get_neighbour_matrix_from_hand(parents,childrens,num_joints=config.hyponet.num_joints,num_edges=config.hyponet.num_twists,knn=config.scorenet.knn)
218 model_score = get_score_net(config, neighbour_matrix, is_train=is_train).to(config.device)
219 model_score_cond = get_pose_net(config, is_train=is_train, score=True).to(config.device)
220
221 if config.training.scorenet.load_weight:
222 states_load = torch.load(config.training.scorenet.gen_path[-1], map_location='cpu')
223 model_score = load_model(model_score, states_load['model'])
224 model_score_cond = load_model(model_score_cond, states_load['model_cond'])
225
226 ema_score = ExponentialMovingAverage(model_score.parameters(), decay=config.scorenet.ema_rate)
227 ema_score_cond = ExponentialMovingAverage(model_score_cond.parameters(), decay=config.scorenet.ema_rate)
228
229 optimizer_score, optimizer_score_cond, loss = None, None, None
230 if is_train:
231 optimizer_score = get_optimizer(config, model_score.parameters(),lr=config.optim.lr_model)
232 backbone_params = list(map(id,model_score_cond.fmap_layer.parameters())) + list(map(id,model_score_cond.hmap_layer.parameters())) + list(map(id,model_score_cond.fmap_layer_local.parameters()))
233 logits_params = filter(lambda p: id(p) not in backbone_params, model_score_cond.parameters())
234 ft_params = filter(lambda p: id(p) in backbone_params, model_score_cond.parameters())
235 optim_list = [{"params": ft_params,"lr":config.optim.lr_hrnet[1]},
236 {"params":logits_params,"lr":config.optim.lr_hrnet[0]}]
237 optimizer_score_cond = torch.optim.Adam(optim_list)
238 loss = SCORE_LOSS(config).to(config.device)
239
240 start_epoch, step = 0, 0
241 if resume:
242 states = torch.load(resume_path, map_location='cpu')
243 model_score.load_state_dict(states['model_score'])
244 model_score_cond.load_state_dict(states['model_score_cond'])
245 ema_score.load_state_dict(states['ema_score'])
246 ema_score.to(config.device)
247 ema_score_cond.load_state_dict(states['ema_score_cond'])
248 ema_score_cond.to(config.device)
249 if is_train:
250 try:
251 optimizer_score.load_state_dict(states['optimizer_score'])
252 optimizer_score_cond.load_state_dict(states['optimizer_score_cond'])
253 start_epoch = states['epoch'] + 1
254 step = states['step']
255 except:
256 print("Fail in loading optimizer!!!!!!!!!!!!!!!")
257 pass
258 print(f"resume from {resume_path}")
259
260 if is_train:
261 model_score.train()
262 model_score_cond.train()
263 else:
264 model_score.eval()
265 model_score_cond.eval()
266 ema_score.copy_to(model_score.parameters())
267 ema_score_cond.copy_to(model_score_cond.parameters())
268 return model_score, model_score_cond, ema_score, ema_score_cond, optimizer_score, optimizer_score_cond, start_epoch, step, loss
269
270def process_pred(pred, dataset, multi_n, type='H36M', save_path=None, use_score=False):
271 error_dict_all = {'mpjpe': [], 'pa_mpjpe':[],'PVE':[],'score':[]}

Callers 3

trainMethod · 0.90
validateMethod · 0.90
validateMethod · 0.85

Calls 13

toMethod · 0.95
load_state_dictMethod · 0.95
copy_toMethod · 0.95
get_score_netFunction · 0.90
get_pose_netFunction · 0.90
SCORE_LOSSClass · 0.90
get_optimizerFunction · 0.85
printFunction · 0.85
load_modelFunction · 0.70
loadMethod · 0.45

Tested by

no test coverage detected