P3Former.
| 6 | |
| 7 | @MODELS.register_module() |
| 8 | class _P3Former(Cylinder3D): |
| 9 | """P3Former.""" |
| 10 | |
| 11 | def __init__(self, |
| 12 | voxel_encoder: ConfigType, |
| 13 | backbone: ConfigType, |
| 14 | decode_head: ConfigType, |
| 15 | neck: OptConfigType = None, |
| 16 | auxiliary_head: OptConfigType = None, |
| 17 | loss_regularization: OptConfigType = None, |
| 18 | train_cfg: OptConfigType = None, |
| 19 | test_cfg: OptConfigType = None, |
| 20 | data_preprocessor: OptConfigType = None, |
| 21 | init_cfg: OptMultiConfig = None) -> None: |
| 22 | super().__init__(voxel_encoder=voxel_encoder, |
| 23 | backbone=backbone, |
| 24 | decode_head=decode_head, |
| 25 | neck=neck, |
| 26 | auxiliary_head=auxiliary_head, |
| 27 | loss_regularization=loss_regularization, |
| 28 | train_cfg=train_cfg, |
| 29 | test_cfg=test_cfg, |
| 30 | data_preprocessor=data_preprocessor, |
| 31 | init_cfg=init_cfg) |
| 32 | |
| 33 | def loss(self, batch_inputs_dict,batch_data_samples): |
| 34 | """Calculate losses from a batch of inputs and data samples. |
| 35 | |
| 36 | Args: |
| 37 | batch_inputs_dict (dict): Input sample dict which |
| 38 | includes 'points' and 'imgs' keys. |
| 39 | |
| 40 | - points (List[Tensor]): Point cloud of each sample. |
| 41 | - imgs (Tensor, optional): Image tensor has shape (B, C, H, W). |
| 42 | batch_data_samples (List[:obj:`Det3DDataSample`]): The det3d data |
| 43 | samples. It usually includes information such as `metainfo` and |
| 44 | `gt_pts_seg`. |
| 45 | |
| 46 | Returns: |
| 47 | Dict[str, Tensor]: A dictionary of loss components. |
| 48 | """ |
| 49 | |
| 50 | # extract features using backbone |
| 51 | x = self.extract_feat(batch_inputs_dict) |
| 52 | batch_inputs_dict['features'] = x.features |
| 53 | losses = dict() |
| 54 | loss_decode = self._decode_head_forward_train(batch_inputs_dict, batch_data_samples) |
| 55 | losses.update(loss_decode) |
| 56 | |
| 57 | return losses |
| 58 | |
| 59 | def predict(self, batch_inputs_dict, batch_data_samples, **kwargs): |
| 60 | x = self.extract_feat(batch_inputs_dict) |
| 61 | batch_inputs_dict['features'] = x.features |
| 62 | pts_semantic_preds, pts_instance_preds = self.decode_head.predict(batch_inputs_dict, batch_data_samples) |
| 63 | return self.postprocess_result(pts_semantic_preds, pts_instance_preds, batch_data_samples) |
| 64 | |
| 65 | def postprocess_result(self, pts_semantic_preds, pts_instance_preds, batch_data_samples): |
nothing calls this directly
no outgoing calls
no test coverage detected