(config, is_train = True, resume = False, resume_path = None)
| 90 | 'Optimizer {} not understood.'.format(config.optim.optimizer)) |
| 91 | |
| 92 | def 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) |
no test coverage detected