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

Function get_model

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

Source from the content-addressed store, hash-verified

90 'Optimizer {} not understood.'.format(config.optim.optimizer))
91
92def get_model(config, is_train = True, resume = False, resume_path = None):
93 neighbour_matrix = get_neighbour_matrix_from_hand(parents,childrens,num_joints=config.hyponet.num_joints,num_edges=config.hyponet.num_twists,knn=config.hyponet.knn)
94 model = get_hyponet(config, neighbour_matrix, is_train=is_train, use_lora=getattr(config.hyponet, 'use_lora', False))
95 model = model.to(config.device)
96 model_cond = get_pose_net(config,is_train=is_train).to(config.device) # HRNet backbone
97 if config.training.get('dpo', False) or config.training.get('kto', False):
98 ref_model = get_hyponet(config, neighbour_matrix, is_train=is_train)
99 ref_model = ref_model.to(config.device)
100 ref_model_cond = get_pose_net(config, is_train=is_train).to(config.device) # HRNet backbone
101
102 if is_train and getattr(config.hyponet, 'use_lora', False): # For LoRA
103 for param in model.parameters():
104 param.requires_grad = False
105 for name, param in model.blocks.named_parameters():
106 if 'lora_proj' in name:
107 param.requires_grad = True
108
109 ema_helper = ExponentialMovingAverage(model.parameters(), decay=config.hyponet.ema_rate)
110 ema_helper_cond = ExponentialMovingAverage(model_cond.parameters(), decay=config.hyponet.ema_rate)
111
112 optimizer_hyponet, optimizer_hrnet, loss = None, None, None
113 if is_train:
114 optimizer_hyponet = get_optimizer(config, model.parameters(), lr=config.optim.lr_model)
115 backbone_params = list(map(id, model_cond.incre_modules.parameters())) + \
116 list(map(id, model_cond.downsamp_modules.parameters())) + \
117 list(map(id, model_cond.final_feat_layer.parameters())) + \
118 list(map(id, model_cond.pred_beta.parameters())) + \
119 list(map(id, model_cond.fmap_layer.parameters())) + \
120 list(map(id, model_cond.hmap_layer.parameters())) + \
121 list(map(id, model_cond.fmap_layer_local.parameters()))
122 logits_params = filter(lambda p: id(p) not in backbone_params, model_cond.parameters())
123 finetune_params = filter(lambda p: id(p) in backbone_params, model_cond.parameters())
124 optim_list =[{"params":logits_params, "lr":config.optim.lr_hrnet[0]},
125 {"params":finetune_params, "lr":config.optim.lr_hrnet[1]}]
126 optimizer_hrnet = torch.optim.Adam(optim_list)
127
128 loss = SMPL_LOSS(config).to(config.device)
129
130 if config.training.get('dpo', False):
131 loss = DPO_SMPL_LOSS(config).to(config.device)
132 if config.training.get('kto', False):
133 loss = KTO_SMPL_LOSS(config).to(config.device)
134
135 start_epoch, step = 0, 0
136 min_mpjpe_pw3d, min_mpjpe_h36m = 1e10, 1e10
137 if resume:
138 states = torch.load(resume_path, map_location='cpu')
139 start_epoch = states['epoch'] + 1
140 step = states['step']
141 if 'min_mpjpe_pw3d' in states:
142 min_mpjpe_pw3d = states['min_mpjpe_pw3d']
143 if 'min_mpjpe_h36m' in states:
144 min_mpjpe_h36m = states['min_mpjpe_h36m']
145 model.load_state_dict(states['model'], strict=False)
146 model_cond.load_state_dict(states['model_cond'])
147
148 if config.training.get('dpo', False) or config.training.get('kto', False):
149 ref_model.load_state_dict(states['model'], strict=False)

Callers 7

trainMethod · 0.90
validateMethod · 0.90
trainMethod · 0.90
validateMethod · 0.90
validateMethod · 0.85
_make_smplerx_modelMethod · 0.85
_make_modelMethod · 0.85

Calls 15

toMethod · 0.95
load_state_dictMethod · 0.95
copy_toMethod · 0.95
get_hyponetFunction · 0.90
get_pose_netFunction · 0.90
SMPL_LOSSClass · 0.90
DPO_SMPL_LOSSClass · 0.90
KTO_SMPL_LOSSClass · 0.90
get_optimizerFunction · 0.85
printFunction · 0.85

Tested by

no test coverage detected