MCPcopy Create free account
hub / github.com/InternRobotics/P3Former / _P3Former

Class _P3Former

p3former/segmentors/p3former.py:8–70  ·  view source on GitHub ↗

P3Former.

Source from the content-addressed store, hash-verified

6
7@MODELS.register_module()
8class _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):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected