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

Method forward_train

detrsmpl/models/architectures/DetrSMPL.py:125–259  ·  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, img, img_metas, **kwargs)

Source from the content-addressed store, hash-verified

123 return outs
124
125 def forward_train(self, img, img_metas, **kwargs):
126 """
127 Args:
128 img (Tensor): Input images of shape (N, C, H, W).
129 Typically these should be mean centered and std scaled.
130 img_metas (list[dict]): A List of image info dict where each dict
131 has: 'img_shape', 'scale_factor', 'flip', and may also contain
132 'filename', 'ori_shape', 'pad_shape', and 'img_norm_cfg'.
133 For details on the values of these keys see
134 :class:`mmdet.datasets.pipelines.Collect`.
135 gt_bboxes (list[Tensor]): Each item are the truth boxes for each
136 image in [tl_x, tl_y, br_x, br_y] format.
137 gt_labels (list[Tensor]): Class indices corresponding to each box
138 gt_bboxes_ignore (None | list[Tensor]): Specify which bounding
139 boxes can be ignored when computing the loss.
140
141 Returns:
142 dict[str, Tensor]: A dictionary of loss components.
143 """
144 # super(SingleStageDetector, self).forward_train(img, img_metas)
145 # NOTE the batched image size information may be useful, e.g.
146 # in DETR, this is needed for the construction of masks, which is
147 # then used for the transformer_head.
148
149 has_smpl = kwargs['has_smpl']
150 gt_smpl_body_pose = kwargs[
151 'smpl_body_pose'] # [bs_0: [ins_num, 23, 3]]
152 gt_smpl_global_orient = kwargs['smpl_global_orient']
153 gt_smpl_body_pose = \
154 [torch.cat((gt_smpl_global_orient[i].view(-1, 1, 3),
155 gt_smpl_body_pose[i]), dim=1).float()
156 for i in range(len(gt_smpl_body_pose))]
157 gt_smpl_betas = kwargs['smpl_betas']
158 gt_smpl_transl = kwargs['smpl_transl']
159 gt_keypoints2d = kwargs['keypoints2d']
160 gt_keypoints3d = kwargs['keypoints3d'] # [bs_0: [N. K, D], ...]
161
162 if 'has_keypoints3d' in kwargs:
163 has_keypoints3d = kwargs['has_keypoints3d']
164 else:
165 has_keypoints3d = None
166
167 if 'has_keypoints2d' in kwargs:
168 has_keypoints2d = kwargs['has_keypoints2d']
169 else:
170 has_keypoints2d = None
171
172 batch_input_shape = tuple(img[0].size()[-2:])
173 for img_meta in img_metas:
174 img_meta['batch_input_shape'] = batch_input_shape
175
176 # features = self.extract_feat(img)
177 features = self.backbone(img)
178
179 if self.neck is not None:
180 features = self.neck(features)
181
182 # outputs_classes, outputs_coords,

Callers

nothing calls this directly

Calls 1

multi_applyFunction · 0.90

Tested by

no test coverage detected