MCPcopy Create free account
hub / github.com/MotrixLab/AiOS / forward

Method forward

detrsmpl/models/architectures/DetrSMPLloss.py:100–228  ·  view source on GitHub ↗

Args: img (Tensor): Input images of shape (N, C, H, W). Typically these should be mean centered and std scaled. img_metas (list[dict]): A List of image info dict where each dict has: 'img_shape', 'scale_factor', 'flip', and may also co

(self, preds, targets)

Source from the content-addressed store, hash-verified

98 pass
99
100 def forward(self, preds, targets):
101 """
102 Args:
103 img (Tensor): Input images of shape (N, C, H, W).
104 Typically these should be mean centered and std scaled.
105 img_metas (list[dict]): A List of image info dict where each dict
106 has: 'img_shape', 'scale_factor', 'flip', and may also contain
107 'filename', 'ori_shape', 'pad_shape', and 'img_norm_cfg'.
108 For details on the values of these keys see
109 :class:`mmdet.datasets.pipelines.Collect`.
110 gt_bboxes (list[Tensor]): Each item are the truth boxes for each
111 image in [tl_x, tl_y, br_x, br_y] format.
112 gt_labels (list[Tensor]): Class indices corresponding to each box
113 gt_bboxes_ignore (None | list[Tensor]): Specify which bounding
114 boxes can be ignored when computing the loss.
115
116 Returns:
117 dict[str, Tensor]: A dictionary of loss components.
118 """
119 # super(SingleStageDetector, self).forward_train(img, img_metas)
120 # NOTE the batched image size information may be useful, e.g.
121 # in DETR, this is needed for the construction of masks, which is
122 # then used for the transformer_head.
123 pred_pose = preds['pred_pose']
124 pred_betas = preds['pred_betas']
125 pred_cameras = preds['pred_cameras']
126 has_smpl = targets['has_smpl']
127 gt_smpl_body_pose = targets[
128 'smpl_body_pose'] # [bs_0: [ins_num, 23, 3]]
129 gt_smpl_global_orient = targets['smpl_global_orient']
130 gt_smpl_body_pose = \
131 [torch.cat((gt_smpl_global_orient[i].view(-1, 1, 3),
132 gt_smpl_body_pose[i]), dim=1).float()
133 for i in range(len(gt_smpl_body_pose))]
134 gt_smpl_betas = targets['smpl_betas']
135 gt_smpl_transl = targets['smpl_transl']
136 gt_keypoints2d = targets['keypoints2d']
137 gt_keypoints3d = targets['keypoints3d'] # [bs_0: [N. K, D], ...]
138 img_metas = targets['img_metas']
139 if 'has_keypoints3d' in targets:
140 has_keypoints3d = targets['has_keypoints3d']
141 else:
142 has_keypoints3d = None
143
144 if 'has_keypoints2d' in targets:
145 has_keypoints2d = targets['has_keypoints2d']
146 else:
147 has_keypoints2d = None
148
149 img = targets['img']
150
151 batch_input_shape = tuple(img[0].size()[-2:])
152 for img_meta in img_metas:
153 img_meta['batch_input_shape'] = batch_input_shape
154
155 L, B, N = pred_pose.shape[:3]
156 if self.body_model_train is not None:
157 pred_output = self.body_model_train(

Callers

nothing calls this directly

Calls 1

multi_applyFunction · 0.90

Tested by

no test coverage detected