MCPcopy Create free account
hub / github.com/ChenhongyiYang/QueryDet-PyTorch / __init__

Method __init__

models/querydet/detector.py:73–157  ·  view source on GitHub ↗
(self, cfg)

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 3

SoftNMSerClass · 0.90
LoopMatcherClass · 0.90

Tested by

no test coverage detected