(self, images)
| 219 | return processed_results |
| 220 | |
| 221 | def test_forward(self, images): |
| 222 | start_event = Event(enable_timing=True) |
| 223 | end_event = Event(enable_timing=True) |
| 224 | |
| 225 | start_event.record() |
| 226 | features = self.backbone(images.tensor[:, :, :]) |
| 227 | |
| 228 | all_features = [features[f] for f in self.in_features] |
| 229 | |
| 230 | all_anchors, all_centers = self.anchor_generator(all_features) |
| 231 | |
| 232 | features_whole = [all_features[x] for x in self.layers_whole_test] |
| 233 | features_value = [all_features[x] for x in self.layers_value_test] |
| 234 | features_key = [all_features[x] for x in self.query_layer_test] |
| 235 | |
| 236 | anchors_whole = [all_anchors[x] for x in self.layers_whole_test] |
| 237 | anchors_value = [all_anchors[x] for x in self.layers_value_test] |
| 238 | |
| 239 | det_cls_whole, det_delta_whole = self.det_head(features_whole) |
| 240 | |
| 241 | |
| 242 | if not self.query_infer: |
| 243 | det_cls_query, det_bbox_query = self.det_head(features_value) |
| 244 | det_cls_query = [permute_to_N_HWA_K(x, self.num_classes) for x in det_cls_query] |
| 245 | det_bbox_query = [permute_to_N_HWA_K(x, 4) for x in det_bbox_query] |
| 246 | query_anchors = anchors_value |
| 247 | else: |
| 248 | if not self.qInfer.initialized: |
| 249 | cls_weights, cls_biases, bbox_weights, bbox_biases = self.det_head.get_params() |
| 250 | qcls_weights, qcls_bias = self.query_head.get_params() |
| 251 | params = [cls_weights, cls_biases, bbox_weights, bbox_biases, qcls_weights, qcls_bias] |
| 252 | else: |
| 253 | params = None |
| 254 | |
| 255 | det_cls_query, det_bbox_query, query_anchors = self.qInfer.run_qinfer(params, features_key, features_value, anchors_value) |
| 256 | |
| 257 | results = self.inference(det_cls_whole, det_delta_whole, anchors_whole, |
| 258 | det_cls_query, det_bbox_query, query_anchors, |
| 259 | images.image_sizes) |
| 260 | |
| 261 | end_event.record() |
| 262 | torch.cuda.synchronize() |
| 263 | total_time = start_event.elapsed_time(end_event) |
| 264 | return results, total_time |
| 265 | |
| 266 | # @float_function |
| 267 | def _giou_loss(self, pred_deltas, anchors, gt_boxes): |
no test coverage detected