| 33 | return x |
| 34 | |
| 35 | class DetHeadMixin(BaseMixin): |
| 36 | def __init__(self, args): |
| 37 | super().__init__() |
| 38 | self.num_det_tokens = args.num_det_tokens |
| 39 | self.class_embed = MLP(args.hidden_size, args.hidden_size, args.num_det_classes, 3) |
| 40 | self.bbox_embed = MLP(args.hidden_size, args.hidden_size, 4, 3) |
| 41 | |
| 42 | def final_forward(self, logits, **kw_args): |
| 43 | logits = logits[:, -self.num_det_tokens:] |
| 44 | outputs_class = self.class_embed(logits) |
| 45 | outputs_coord = self.bbox_embed(logits).sigmoid() |
| 46 | out = {'pred_logits': outputs_class, 'pred_boxes': outputs_coord} |
| 47 | return out |
| 48 | |
| 49 | class YOLOS(ViTModel): |
| 50 | def __init__(self, args, transformer=None, layernorm_epsilon=1e-6, **kwargs): |