(self, cfg)
| 71 | Implement Our QueryDet |
| 72 | """ |
| 73 | def __init__(self, cfg): |
| 74 | super().__init__() |
| 75 | |
| 76 | # fmt: off |
| 77 | self.num_classes = cfg.MODEL.RETINANET.NUM_CLASSES |
| 78 | self.in_features = cfg.MODEL.RETINANET.IN_FEATURES |
| 79 | self.query_layer_train = cfg.MODEL.QUERY.Q_FEATURE_TRAIN |
| 80 | self.layers_whole_test = cfg.MODEL.QUERY.FEATURES_WHOLE_TEST |
| 81 | self.layers_value_test = cfg.MODEL.QUERY.FEATURES_VALUE_TEST |
| 82 | self.query_layer_test = cfg.MODEL.QUERY.Q_FEATURE_TEST |
| 83 | # Loss parameters: |
| 84 | self.focal_loss_alpha = cfg.MODEL.CUSTOM.FOCAL_LOSS_ALPHAS |
| 85 | self.focal_loss_gamma = cfg.MODEL.CUSTOM.FOCAL_LOSS_GAMMAS |
| 86 | self.smooth_l1_loss_beta = cfg.MODEL.RETINANET.SMOOTH_L1_LOSS_BETA |
| 87 | self.use_giou_loss = cfg.MODEL.CUSTOM.GIOU_LOSS |
| 88 | self.cls_weights = cfg.MODEL.CUSTOM.CLS_WEIGHTS |
| 89 | self.reg_weights = cfg.MODEL.CUSTOM.REG_WEIGHTS |
| 90 | # training query head |
| 91 | self.small_obj_scale = cfg.MODEL.QUERY.ENCODE_SMALL_OBJ_SCALE |
| 92 | self.query_loss_weights = cfg.MODEL.QUERY.QUERY_LOSS_WEIGHT |
| 93 | self.query_loss_gammas = cfg.MODEL.QUERY.QUERY_LOSS_GAMMA |
| 94 | self.small_center_dis_coeff = cfg.MODEL.QUERY.ENCODE_CENTER_DIS_COEFF |
| 95 | # Inference parameters: |
| 96 | self.score_threshold = cfg.MODEL.RETINANET.SCORE_THRESH_TEST |
| 97 | self.topk_candidates = cfg.MODEL.RETINANET.TOPK_CANDIDATES_TEST |
| 98 | self.use_soft_nms = cfg.MODEL.CUSTOM.USE_SOFT_NMS |
| 99 | self.nms_threshold = cfg.MODEL.RETINANET.NMS_THRESH_TEST |
| 100 | self.max_detections_per_image = cfg.TEST.DETECTIONS_PER_IMAGE |
| 101 | # query inference |
| 102 | self.query_infer = cfg.MODEL.QUERY.QUERY_INFER |
| 103 | self.query_threshold = cfg.MODEL.QUERY.THRESHOLD |
| 104 | self.query_context = cfg.MODEL.QUERY.CONTEXT |
| 105 | # other settings |
| 106 | self.clear_cuda_cache = cfg.MODEL.CUSTOM.CLEAR_CUDA_CACHE |
| 107 | self.anchor_num = len(cfg.MODEL.ANCHOR_GENERATOR.ASPECT_RATIOS[0]) * \ |
| 108 | len(cfg.MODEL.ANCHOR_GENERATOR.SIZES[0]) |
| 109 | self.with_cp = cfg.MODEL.CUSTOM.GRADIENT_CHECKPOINT |
| 110 | # fmt: on |
| 111 | assert 'p2' in self.in_features |
| 112 | |
| 113 | self.backbone = build_backbone(cfg) |
| 114 | if cfg.MODEL.CUSTOM.HEAD_BN: |
| 115 | self.det_head = dh.RetinaNetHead_3x3_MergeBN(cfg, 256, 256, 4, self.anchor_num) |
| 116 | self.query_head = dh.Head_3x3_MergeBN(256, 256, 4, 1) |
| 117 | else: |
| 118 | self.det_head = dh.RetinaNetHead_3x3(cfg, 256, 256, 4, self.anchor_num) |
| 119 | self.query_head = dh.Head_3x3(256, 256, 4, 1) |
| 120 | |
| 121 | self.qInfer = qf.QueryInfer(9, self.num_classes, self.query_threshold, self.query_context) |
| 122 | |
| 123 | backbone_shape = self.backbone.output_shape() |
| 124 | all_det_feature_shapes = [backbone_shape[f] for f in self.in_features] |
| 125 | |
| 126 | self.anchor_generator = build_anchor_generator(cfg, all_det_feature_shapes) |
| 127 | self.query_anchor_generator = AnchorGeneratorWithCenter(sizes=[128], aspect_ratios=[1.0], |
| 128 | strides=[2**(x+2) for x in self.query_layer_train], offset=0.5) |
| 129 | # Matching and loss |
| 130 | self.box2box_transform = Box2BoxTransform(weights=cfg.MODEL.RPN.BBOX_REG_WEIGHTS) |
nothing calls this directly
no test coverage detected