segment-everything 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- segment_everything/__init__.py +5 -0
- segment_everything/augmentation/albumentations_helper.py +0 -0
- segment_everything/detect_and_segment.py +131 -0
- segment_everything/napari_helper.py +15 -0
- segment_everything/prompt_generator.py +188 -0
- segment_everything/py.typed +5 -0
- segment_everything/stacked_label_dataset.py +113 -0
- segment_everything/stacked_labels.py +428 -0
- segment_everything/vendored/PromptGuidedDecoder/Prompt_guided_Mask_Decoder.pt +0 -0
- segment_everything/vendored/__init__.py +5 -0
- segment_everything/vendored/dice.py +158 -0
- segment_everything/vendored/efficientvit/__init__.py +0 -0
- segment_everything/vendored/efficientvit/apps/__init__.py +0 -0
- segment_everything/vendored/efficientvit/apps/data_provider/__init__.py +7 -0
- segment_everything/vendored/efficientvit/apps/data_provider/augment/__init__.py +6 -0
- segment_everything/vendored/efficientvit/apps/data_provider/augment/bbox.py +30 -0
- segment_everything/vendored/efficientvit/apps/data_provider/augment/color_aug.py +78 -0
- segment_everything/vendored/efficientvit/apps/data_provider/base.py +254 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/__init__.py +6 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/_data_loader.py +1538 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/_data_worker.py +357 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/controller.py +100 -0
- segment_everything/vendored/efficientvit/apps/setup.py +150 -0
- segment_everything/vendored/efficientvit/apps/trainer/__init__.py +6 -0
- segment_everything/vendored/efficientvit/apps/trainer/base.py +318 -0
- segment_everything/vendored/efficientvit/apps/trainer/run_config.py +129 -0
- segment_everything/vendored/efficientvit/apps/utils/__init__.py +12 -0
- segment_everything/vendored/efficientvit/apps/utils/dist.py +32 -0
- segment_everything/vendored/efficientvit/apps/utils/ema.py +52 -0
- segment_everything/vendored/efficientvit/apps/utils/export.py +45 -0
- segment_everything/vendored/efficientvit/apps/utils/init.py +66 -0
- segment_everything/vendored/efficientvit/apps/utils/lr.py +52 -0
- segment_everything/vendored/efficientvit/apps/utils/metric.py +43 -0
- segment_everything/vendored/efficientvit/apps/utils/misc.py +101 -0
- segment_everything/vendored/efficientvit/apps/utils/opt.py +28 -0
- segment_everything/vendored/efficientvit/cls_model_zoo.py +79 -0
- segment_everything/vendored/efficientvit/clscore/__init__.py +0 -0
- segment_everything/vendored/efficientvit/clscore/data_provider/__init__.py +5 -0
- segment_everything/vendored/efficientvit/clscore/data_provider/imagenet.py +142 -0
- segment_everything/vendored/efficientvit/clscore/trainer/__init__.py +6 -0
- segment_everything/vendored/efficientvit/clscore/trainer/cls_run_config.py +18 -0
- segment_everything/vendored/efficientvit/clscore/trainer/cls_trainer.py +265 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/__init__.py +7 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/label_smooth.py +18 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/metric.py +23 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/mixup.py +67 -0
- segment_everything/vendored/efficientvit/models/__init__.py +0 -0
- segment_everything/vendored/efficientvit/models/efficientvit/__init__.py +8 -0
- segment_everything/vendored/efficientvit/models/efficientvit/backbone.py +380 -0
- segment_everything/vendored/efficientvit/models/efficientvit/cls.py +188 -0
- segment_everything/vendored/efficientvit/models/efficientvit/sam.py +181 -0
- segment_everything/vendored/efficientvit/models/efficientvit/seg.py +373 -0
- segment_everything/vendored/efficientvit/models/nn/__init__.py +8 -0
- segment_everything/vendored/efficientvit/models/nn/act.py +30 -0
- segment_everything/vendored/efficientvit/models/nn/drop.py +104 -0
- segment_everything/vendored/efficientvit/models/nn/norm.py +164 -0
- segment_everything/vendored/efficientvit/models/nn/ops.py +597 -0
- segment_everything/vendored/efficientvit/models/utils/__init__.py +7 -0
- segment_everything/vendored/efficientvit/models/utils/list.py +53 -0
- segment_everything/vendored/efficientvit/models/utils/network.py +73 -0
- segment_everything/vendored/efficientvit/models/utils/random.py +65 -0
- segment_everything/vendored/efficientvit/sam_model_zoo.py +45 -0
- segment_everything/vendored/efficientvit/seg_model_zoo.py +70 -0
- segment_everything/vendored/get_object_aware.py +26 -0
- segment_everything/vendored/mobilesamv2/__init__.py +16 -0
- segment_everything/vendored/mobilesamv2/automatic_mask_generator.py +415 -0
- segment_everything/vendored/mobilesamv2/build_sam.py +246 -0
- segment_everything/vendored/mobilesamv2/modeling/__init__.py +11 -0
- segment_everything/vendored/mobilesamv2/modeling/common.py +43 -0
- segment_everything/vendored/mobilesamv2/modeling/image_encoder.py +394 -0
- segment_everything/vendored/mobilesamv2/modeling/mask_decoder.py +213 -0
- segment_everything/vendored/mobilesamv2/modeling/prompt_encoder.py +217 -0
- segment_everything/vendored/mobilesamv2/modeling/sam.py +203 -0
- segment_everything/vendored/mobilesamv2/modeling/transformer.py +240 -0
- segment_everything/vendored/mobilesamv2/predictor.py +384 -0
- segment_everything/vendored/mobilesamv2/utils/__init__.py +5 -0
- segment_everything/vendored/mobilesamv2/utils/amg.py +347 -0
- segment_everything/vendored/mobilesamv2/utils/onnx.py +144 -0
- segment_everything/vendored/mobilesamv2/utils/transforms.py +103 -0
- segment_everything/vendored/object_detection/__init__.py +0 -0
- segment_everything/vendored/object_detection/ultralytics/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/nn/__init__.py +9 -0
- segment_everything/vendored/object_detection/ultralytics/nn/autobackend.py +658 -0
- segment_everything/vendored/object_detection/ultralytics/nn/autoshape.py +397 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/__init__.py +110 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/block.py +304 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/conv.py +297 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/head.py +468 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/transformer.py +378 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/utils.py +78 -0
- segment_everything/vendored/object_detection/ultralytics/nn/tasks.py +1049 -0
- segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/__init__.py +6 -0
- segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/model.py +104 -0
- segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/predict.py +95 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/cfg/__init__.py +588 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/cfg/default.yaml +117 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/__init__.py +9 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/annotator.py +53 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/augment.py +899 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/base.py +286 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/build.py +213 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/converter.py +358 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataloaders/__init__.py +0 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataloaders/stream_loaders.py +459 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataset.py +274 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataset_wrappers.py +53 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/download_weights.sh +18 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_coco.sh +60 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_coco128.sh +17 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_imagenet.sh +51 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/utils.py +716 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/__init__.py +0 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/exporter.py +1214 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/model.py +641 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/predictor.py +461 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/results.py +741 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/__init__.py +893 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/autobatch.py +108 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/callbacks/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/callbacks/base.py +212 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/checks.py +547 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/dist.py +67 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/downloads.py +353 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/errors.py +12 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/files.py +100 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/instance.py +391 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/loss.py +579 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/metrics.py +1189 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/ops.py +870 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/patches.py +45 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/plotting.py +767 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/tal.py +276 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/torch_utils.py +684 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/tuner.py +54 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/v8/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/v8/detect/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/v8/detect/predict.py +69 -0
- segment_everything/vendored/tinyvit/__init__.py +2 -0
- segment_everything/vendored/tinyvit/tiny_vit.py +867 -0
- segment_everything/weights_helper.py +124 -0
- segment_everything-0.1.0.dist-info/METADATA +53 -0
- segment_everything-0.1.0.dist-info/RECORD +145 -0
- segment_everything-0.1.0.dist-info/WHEEL +4 -0
- segment_everything-0.1.0.dist-info/licenses/LICENSE +28 -0
|
@@ -0,0 +1,304 @@
|
|
|
1
|
+
# Ultralytics YOLO 🚀, AGPL-3.0 license
|
|
2
|
+
"""
|
|
3
|
+
Block modules
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
import torch
|
|
7
|
+
import torch.nn as nn
|
|
8
|
+
import torch.nn.functional as F
|
|
9
|
+
|
|
10
|
+
from .conv import Conv, DWConv, GhostConv, LightConv, RepConv
|
|
11
|
+
from .transformer import TransformerBlock
|
|
12
|
+
|
|
13
|
+
__all__ = ('DFL', 'HGBlock', 'HGStem', 'SPP', 'SPPF', 'C1', 'C2', 'C3', 'C2f', 'C3x', 'C3TR', 'C3Ghost',
|
|
14
|
+
'GhostBottleneck', 'Bottleneck', 'BottleneckCSP', 'Proto', 'RepC3')
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class DFL(nn.Module):
|
|
18
|
+
"""
|
|
19
|
+
Integral module of Distribution Focal Loss (DFL).
|
|
20
|
+
Proposed in Generalized Focal Loss https://ieeexplore.ieee.org/document/9792391
|
|
21
|
+
"""
|
|
22
|
+
|
|
23
|
+
def __init__(self, c1=16):
|
|
24
|
+
"""Initialize a convolutional layer with a given number of input channels."""
|
|
25
|
+
super().__init__()
|
|
26
|
+
self.conv = nn.Conv2d(c1, 1, 1, bias=False).requires_grad_(False)
|
|
27
|
+
x = torch.arange(c1, dtype=torch.float)
|
|
28
|
+
self.conv.weight.data[:] = nn.Parameter(x.view(1, c1, 1, 1))
|
|
29
|
+
self.c1 = c1
|
|
30
|
+
|
|
31
|
+
def forward(self, x):
|
|
32
|
+
"""Applies a transformer layer on input tensor 'x' and returns a tensor."""
|
|
33
|
+
b, c, a = x.shape # batch, channels, anchors
|
|
34
|
+
return self.conv(x.view(b, 4, self.c1, a).transpose(2, 1).softmax(1)).view(b, 4, a)
|
|
35
|
+
# return self.conv(x.view(b, self.c1, 4, a).softmax(1)).view(b, 4, a)
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
class Proto(nn.Module):
|
|
39
|
+
"""YOLOv8 mask Proto module for segmentation models."""
|
|
40
|
+
|
|
41
|
+
def __init__(self, c1, c_=256, c2=32): # ch_in, number of protos, number of masks
|
|
42
|
+
super().__init__()
|
|
43
|
+
self.cv1 = Conv(c1, c_, k=3)
|
|
44
|
+
self.upsample = nn.ConvTranspose2d(c_, c_, 2, 2, 0, bias=True) # nn.Upsample(scale_factor=2, mode='nearest')
|
|
45
|
+
self.cv2 = Conv(c_, c_, k=3)
|
|
46
|
+
self.cv3 = Conv(c_, c2)
|
|
47
|
+
|
|
48
|
+
def forward(self, x):
|
|
49
|
+
"""Performs a forward pass through layers using an upsampled input image."""
|
|
50
|
+
return self.cv3(self.cv2(self.upsample(self.cv1(x))))
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class HGStem(nn.Module):
|
|
54
|
+
"""StemBlock of PPHGNetV2 with 5 convolutions and one maxpool2d.
|
|
55
|
+
https://github.com/PaddlePaddle/PaddleDetection/blob/develop/ppdet/modeling/backbones/hgnet_v2.py
|
|
56
|
+
"""
|
|
57
|
+
|
|
58
|
+
def __init__(self, c1, cm, c2):
|
|
59
|
+
super().__init__()
|
|
60
|
+
self.stem1 = Conv(c1, cm, 3, 2, act=nn.ReLU())
|
|
61
|
+
self.stem2a = Conv(cm, cm // 2, 2, 1, 0, act=nn.ReLU())
|
|
62
|
+
self.stem2b = Conv(cm // 2, cm, 2, 1, 0, act=nn.ReLU())
|
|
63
|
+
self.stem3 = Conv(cm * 2, cm, 3, 2, act=nn.ReLU())
|
|
64
|
+
self.stem4 = Conv(cm, c2, 1, 1, act=nn.ReLU())
|
|
65
|
+
self.pool = nn.MaxPool2d(kernel_size=2, stride=1, padding=0, ceil_mode=True)
|
|
66
|
+
|
|
67
|
+
def forward(self, x):
|
|
68
|
+
"""Forward pass of a PPHGNetV2 backbone layer."""
|
|
69
|
+
x = self.stem1(x)
|
|
70
|
+
x = F.pad(x, [0, 1, 0, 1])
|
|
71
|
+
x2 = self.stem2a(x)
|
|
72
|
+
x2 = F.pad(x2, [0, 1, 0, 1])
|
|
73
|
+
x2 = self.stem2b(x2)
|
|
74
|
+
x1 = self.pool(x)
|
|
75
|
+
x = torch.cat([x1, x2], dim=1)
|
|
76
|
+
x = self.stem3(x)
|
|
77
|
+
x = self.stem4(x)
|
|
78
|
+
return x
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
class HGBlock(nn.Module):
|
|
82
|
+
"""HG_Block of PPHGNetV2 with 2 convolutions and LightConv.
|
|
83
|
+
https://github.com/PaddlePaddle/PaddleDetection/blob/develop/ppdet/modeling/backbones/hgnet_v2.py
|
|
84
|
+
"""
|
|
85
|
+
|
|
86
|
+
def __init__(self, c1, cm, c2, k=3, n=6, lightconv=False, shortcut=False, act=nn.ReLU()):
|
|
87
|
+
super().__init__()
|
|
88
|
+
block = LightConv if lightconv else Conv
|
|
89
|
+
self.m = nn.ModuleList(block(c1 if i == 0 else cm, cm, k=k, act=act) for i in range(n))
|
|
90
|
+
self.sc = Conv(c1 + n * cm, c2 // 2, 1, 1, act=act) # squeeze conv
|
|
91
|
+
self.ec = Conv(c2 // 2, c2, 1, 1, act=act) # excitation conv
|
|
92
|
+
self.add = shortcut and c1 == c2
|
|
93
|
+
|
|
94
|
+
def forward(self, x):
|
|
95
|
+
"""Forward pass of a PPHGNetV2 backbone layer."""
|
|
96
|
+
y = [x]
|
|
97
|
+
y.extend(m(y[-1]) for m in self.m)
|
|
98
|
+
y = self.ec(self.sc(torch.cat(y, 1)))
|
|
99
|
+
return y + x if self.add else y
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
class SPP(nn.Module):
|
|
103
|
+
"""Spatial Pyramid Pooling (SPP) layer https://arxiv.org/abs/1406.4729."""
|
|
104
|
+
|
|
105
|
+
def __init__(self, c1, c2, k=(5, 9, 13)):
|
|
106
|
+
"""Initialize the SPP layer with input/output channels and pooling kernel sizes."""
|
|
107
|
+
super().__init__()
|
|
108
|
+
c_ = c1 // 2 # hidden channels
|
|
109
|
+
self.cv1 = Conv(c1, c_, 1, 1)
|
|
110
|
+
self.cv2 = Conv(c_ * (len(k) + 1), c2, 1, 1)
|
|
111
|
+
self.m = nn.ModuleList([nn.MaxPool2d(kernel_size=x, stride=1, padding=x // 2) for x in k])
|
|
112
|
+
|
|
113
|
+
def forward(self, x):
|
|
114
|
+
"""Forward pass of the SPP layer, performing spatial pyramid pooling."""
|
|
115
|
+
x = self.cv1(x)
|
|
116
|
+
return self.cv2(torch.cat([x] + [m(x) for m in self.m], 1))
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
class SPPF(nn.Module):
|
|
120
|
+
"""Spatial Pyramid Pooling - Fast (SPPF) layer for YOLOv5 by Glenn Jocher."""
|
|
121
|
+
|
|
122
|
+
def __init__(self, c1, c2, k=5): # equivalent to SPP(k=(5, 9, 13))
|
|
123
|
+
super().__init__()
|
|
124
|
+
c_ = c1 // 2 # hidden channels
|
|
125
|
+
self.cv1 = Conv(c1, c_, 1, 1)
|
|
126
|
+
self.cv2 = Conv(c_ * 4, c2, 1, 1)
|
|
127
|
+
self.m = nn.MaxPool2d(kernel_size=k, stride=1, padding=k // 2)
|
|
128
|
+
|
|
129
|
+
def forward(self, x):
|
|
130
|
+
"""Forward pass through Ghost Convolution block."""
|
|
131
|
+
x = self.cv1(x)
|
|
132
|
+
y1 = self.m(x)
|
|
133
|
+
y2 = self.m(y1)
|
|
134
|
+
return self.cv2(torch.cat((x, y1, y2, self.m(y2)), 1))
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
class C1(nn.Module):
|
|
138
|
+
"""CSP Bottleneck with 1 convolution."""
|
|
139
|
+
|
|
140
|
+
def __init__(self, c1, c2, n=1): # ch_in, ch_out, number
|
|
141
|
+
super().__init__()
|
|
142
|
+
self.cv1 = Conv(c1, c2, 1, 1)
|
|
143
|
+
self.m = nn.Sequential(*(Conv(c2, c2, 3) for _ in range(n)))
|
|
144
|
+
|
|
145
|
+
def forward(self, x):
|
|
146
|
+
"""Applies cross-convolutions to input in the C3 module."""
|
|
147
|
+
y = self.cv1(x)
|
|
148
|
+
return self.m(y) + y
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
class C2(nn.Module):
|
|
152
|
+
"""CSP Bottleneck with 2 convolutions."""
|
|
153
|
+
|
|
154
|
+
def __init__(self, c1, c2, n=1, shortcut=True, g=1, e=0.5): # ch_in, ch_out, number, shortcut, groups, expansion
|
|
155
|
+
super().__init__()
|
|
156
|
+
self.c = int(c2 * e) # hidden channels
|
|
157
|
+
self.cv1 = Conv(c1, 2 * self.c, 1, 1)
|
|
158
|
+
self.cv2 = Conv(2 * self.c, c2, 1) # optional act=FReLU(c2)
|
|
159
|
+
# self.attention = ChannelAttention(2 * self.c) # or SpatialAttention()
|
|
160
|
+
self.m = nn.Sequential(*(Bottleneck(self.c, self.c, shortcut, g, k=((3, 3), (3, 3)), e=1.0) for _ in range(n)))
|
|
161
|
+
|
|
162
|
+
def forward(self, x):
|
|
163
|
+
"""Forward pass through the CSP bottleneck with 2 convolutions."""
|
|
164
|
+
a, b = self.cv1(x).chunk(2, 1)
|
|
165
|
+
return self.cv2(torch.cat((self.m(a), b), 1))
|
|
166
|
+
|
|
167
|
+
|
|
168
|
+
class C2f(nn.Module):
|
|
169
|
+
"""CSP Bottleneck with 2 convolutions."""
|
|
170
|
+
|
|
171
|
+
def __init__(self, c1, c2, n=1, shortcut=False, g=1, e=0.5): # ch_in, ch_out, number, shortcut, groups, expansion
|
|
172
|
+
super().__init__()
|
|
173
|
+
self.c = int(c2 * e) # hidden channels
|
|
174
|
+
self.cv1 = Conv(c1, 2 * self.c, 1, 1)
|
|
175
|
+
self.cv2 = Conv((2 + n) * self.c, c2, 1) # optional act=FReLU(c2)
|
|
176
|
+
self.m = nn.ModuleList(Bottleneck(self.c, self.c, shortcut, g, k=((3, 3), (3, 3)), e=1.0) for _ in range(n))
|
|
177
|
+
|
|
178
|
+
def forward(self, x):
|
|
179
|
+
"""Forward pass through C2f layer."""
|
|
180
|
+
y = list(self.cv1(x).chunk(2, 1))
|
|
181
|
+
y.extend(m(y[-1]) for m in self.m)
|
|
182
|
+
return self.cv2(torch.cat(y, 1))
|
|
183
|
+
|
|
184
|
+
def forward_split(self, x):
|
|
185
|
+
"""Forward pass using split() instead of chunk()."""
|
|
186
|
+
y = list(self.cv1(x).split((self.c, self.c), 1))
|
|
187
|
+
y.extend(m(y[-1]) for m in self.m)
|
|
188
|
+
return self.cv2(torch.cat(y, 1))
|
|
189
|
+
|
|
190
|
+
|
|
191
|
+
class C3(nn.Module):
|
|
192
|
+
"""CSP Bottleneck with 3 convolutions."""
|
|
193
|
+
|
|
194
|
+
def __init__(self, c1, c2, n=1, shortcut=True, g=1, e=0.5): # ch_in, ch_out, number, shortcut, groups, expansion
|
|
195
|
+
super().__init__()
|
|
196
|
+
c_ = int(c2 * e) # hidden channels
|
|
197
|
+
self.cv1 = Conv(c1, c_, 1, 1)
|
|
198
|
+
self.cv2 = Conv(c1, c_, 1, 1)
|
|
199
|
+
self.cv3 = Conv(2 * c_, c2, 1) # optional act=FReLU(c2)
|
|
200
|
+
self.m = nn.Sequential(*(Bottleneck(c_, c_, shortcut, g, k=((1, 1), (3, 3)), e=1.0) for _ in range(n)))
|
|
201
|
+
|
|
202
|
+
def forward(self, x):
|
|
203
|
+
"""Forward pass through the CSP bottleneck with 2 convolutions."""
|
|
204
|
+
return self.cv3(torch.cat((self.m(self.cv1(x)), self.cv2(x)), 1))
|
|
205
|
+
|
|
206
|
+
|
|
207
|
+
class C3x(C3):
|
|
208
|
+
"""C3 module with cross-convolutions."""
|
|
209
|
+
|
|
210
|
+
def __init__(self, c1, c2, n=1, shortcut=True, g=1, e=0.5):
|
|
211
|
+
"""Initialize C3TR instance and set default parameters."""
|
|
212
|
+
super().__init__(c1, c2, n, shortcut, g, e)
|
|
213
|
+
self.c_ = int(c2 * e)
|
|
214
|
+
self.m = nn.Sequential(*(Bottleneck(self.c_, self.c_, shortcut, g, k=((1, 3), (3, 1)), e=1) for _ in range(n)))
|
|
215
|
+
|
|
216
|
+
|
|
217
|
+
class RepC3(nn.Module):
|
|
218
|
+
"""Rep C3."""
|
|
219
|
+
|
|
220
|
+
def __init__(self, c1, c2, n=3, e=1.0):
|
|
221
|
+
super().__init__()
|
|
222
|
+
c_ = int(c2 * e) # hidden channels
|
|
223
|
+
self.cv1 = Conv(c1, c2, 1, 1)
|
|
224
|
+
self.cv2 = Conv(c1, c2, 1, 1)
|
|
225
|
+
self.m = nn.Sequential(*[RepConv(c_, c_) for _ in range(n)])
|
|
226
|
+
self.cv3 = Conv(c_, c2, 1, 1) if c_ != c2 else nn.Identity()
|
|
227
|
+
|
|
228
|
+
def forward(self, x):
|
|
229
|
+
"""Forward pass of RT-DETR neck layer."""
|
|
230
|
+
return self.cv3(self.m(self.cv1(x)) + self.cv2(x))
|
|
231
|
+
|
|
232
|
+
|
|
233
|
+
class C3TR(C3):
|
|
234
|
+
"""C3 module with TransformerBlock()."""
|
|
235
|
+
|
|
236
|
+
def __init__(self, c1, c2, n=1, shortcut=True, g=1, e=0.5):
|
|
237
|
+
"""Initialize C3Ghost module with GhostBottleneck()."""
|
|
238
|
+
super().__init__(c1, c2, n, shortcut, g, e)
|
|
239
|
+
c_ = int(c2 * e)
|
|
240
|
+
self.m = TransformerBlock(c_, c_, 4, n)
|
|
241
|
+
|
|
242
|
+
|
|
243
|
+
class C3Ghost(C3):
|
|
244
|
+
"""C3 module with GhostBottleneck()."""
|
|
245
|
+
|
|
246
|
+
def __init__(self, c1, c2, n=1, shortcut=True, g=1, e=0.5):
|
|
247
|
+
"""Initialize 'SPP' module with various pooling sizes for spatial pyramid pooling."""
|
|
248
|
+
super().__init__(c1, c2, n, shortcut, g, e)
|
|
249
|
+
c_ = int(c2 * e) # hidden channels
|
|
250
|
+
self.m = nn.Sequential(*(GhostBottleneck(c_, c_) for _ in range(n)))
|
|
251
|
+
|
|
252
|
+
|
|
253
|
+
class GhostBottleneck(nn.Module):
|
|
254
|
+
"""Ghost Bottleneck https://github.com/huawei-noah/ghostnet."""
|
|
255
|
+
|
|
256
|
+
def __init__(self, c1, c2, k=3, s=1): # ch_in, ch_out, kernel, stride
|
|
257
|
+
super().__init__()
|
|
258
|
+
c_ = c2 // 2
|
|
259
|
+
self.conv = nn.Sequential(
|
|
260
|
+
GhostConv(c1, c_, 1, 1), # pw
|
|
261
|
+
DWConv(c_, c_, k, s, act=False) if s == 2 else nn.Identity(), # dw
|
|
262
|
+
GhostConv(c_, c2, 1, 1, act=False)) # pw-linear
|
|
263
|
+
self.shortcut = nn.Sequential(DWConv(c1, c1, k, s, act=False), Conv(c1, c2, 1, 1,
|
|
264
|
+
act=False)) if s == 2 else nn.Identity()
|
|
265
|
+
|
|
266
|
+
def forward(self, x):
|
|
267
|
+
"""Applies skip connection and concatenation to input tensor."""
|
|
268
|
+
return self.conv(x) + self.shortcut(x)
|
|
269
|
+
|
|
270
|
+
|
|
271
|
+
class Bottleneck(nn.Module):
|
|
272
|
+
"""Standard bottleneck."""
|
|
273
|
+
|
|
274
|
+
def __init__(self, c1, c2, shortcut=True, g=1, k=(3, 3), e=0.5): # ch_in, ch_out, shortcut, groups, kernels, expand
|
|
275
|
+
super().__init__()
|
|
276
|
+
c_ = int(c2 * e) # hidden channels
|
|
277
|
+
self.cv1 = Conv(c1, c_, k[0], 1)
|
|
278
|
+
self.cv2 = Conv(c_, c2, k[1], 1, g=g)
|
|
279
|
+
self.add = shortcut and c1 == c2
|
|
280
|
+
|
|
281
|
+
def forward(self, x):
|
|
282
|
+
"""'forward()' applies the YOLOv5 FPN to input data."""
|
|
283
|
+
return x + self.cv2(self.cv1(x)) if self.add else self.cv2(self.cv1(x))
|
|
284
|
+
|
|
285
|
+
|
|
286
|
+
class BottleneckCSP(nn.Module):
|
|
287
|
+
"""CSP Bottleneck https://github.com/WongKinYiu/CrossStagePartialNetworks."""
|
|
288
|
+
|
|
289
|
+
def __init__(self, c1, c2, n=1, shortcut=True, g=1, e=0.5): # ch_in, ch_out, number, shortcut, groups, expansion
|
|
290
|
+
super().__init__()
|
|
291
|
+
c_ = int(c2 * e) # hidden channels
|
|
292
|
+
self.cv1 = Conv(c1, c_, 1, 1)
|
|
293
|
+
self.cv2 = nn.Conv2d(c1, c_, 1, 1, bias=False)
|
|
294
|
+
self.cv3 = nn.Conv2d(c_, c_, 1, 1, bias=False)
|
|
295
|
+
self.cv4 = Conv(2 * c_, c2, 1, 1)
|
|
296
|
+
self.bn = nn.BatchNorm2d(2 * c_) # applied to cat(cv2, cv3)
|
|
297
|
+
self.act = nn.SiLU()
|
|
298
|
+
self.m = nn.Sequential(*(Bottleneck(c_, c_, shortcut, g, e=1.0) for _ in range(n)))
|
|
299
|
+
|
|
300
|
+
def forward(self, x):
|
|
301
|
+
"""Applies a CSP bottleneck with 3 convolutions."""
|
|
302
|
+
y1 = self.cv3(self.m(self.cv1(x)))
|
|
303
|
+
y2 = self.cv2(x)
|
|
304
|
+
return self.cv4(self.act(self.bn(torch.cat((y1, y2), 1))))
|
|
@@ -0,0 +1,297 @@
|
|
|
1
|
+
# Ultralytics YOLO 🚀, AGPL-3.0 license
|
|
2
|
+
"""
|
|
3
|
+
Convolution modules
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
import math
|
|
7
|
+
|
|
8
|
+
import numpy as np
|
|
9
|
+
import torch
|
|
10
|
+
import torch.nn as nn
|
|
11
|
+
|
|
12
|
+
__all__ = ('Conv', 'LightConv', 'DWConv', 'DWConvTranspose2d', 'ConvTranspose', 'Focus', 'GhostConv',
|
|
13
|
+
'ChannelAttention', 'SpatialAttention', 'CBAM', 'Concat', 'RepConv')
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def autopad(k, p=None, d=1): # kernel, padding, dilation
|
|
17
|
+
"""Pad to 'same' shape outputs."""
|
|
18
|
+
if d > 1:
|
|
19
|
+
k = d * (k - 1) + 1 if isinstance(k, int) else [d * (x - 1) + 1 for x in k] # actual kernel-size
|
|
20
|
+
if p is None:
|
|
21
|
+
p = k // 2 if isinstance(k, int) else [x // 2 for x in k] # auto-pad
|
|
22
|
+
return p
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class Conv(nn.Module):
|
|
26
|
+
"""Standard convolution with args(ch_in, ch_out, kernel, stride, padding, groups, dilation, activation)."""
|
|
27
|
+
default_act = nn.SiLU() # default activation
|
|
28
|
+
|
|
29
|
+
def __init__(self, c1, c2, k=1, s=1, p=None, g=1, d=1, act=True):
|
|
30
|
+
"""Initialize Conv layer with given arguments including activation."""
|
|
31
|
+
super().__init__()
|
|
32
|
+
self.conv = nn.Conv2d(c1, c2, k, s, autopad(k, p, d), groups=g, dilation=d, bias=False)
|
|
33
|
+
self.bn = nn.BatchNorm2d(c2)
|
|
34
|
+
self.act = self.default_act if act is True else act if isinstance(act, nn.Module) else nn.Identity()
|
|
35
|
+
|
|
36
|
+
def forward(self, x):
|
|
37
|
+
"""Apply convolution, batch normalization and activation to input tensor."""
|
|
38
|
+
return self.act(self.bn(self.conv(x)))
|
|
39
|
+
|
|
40
|
+
def forward_fuse(self, x):
|
|
41
|
+
"""Perform transposed convolution of 2D data."""
|
|
42
|
+
return self.act(self.conv(x))
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
class Conv2(Conv):
|
|
46
|
+
"""Simplified RepConv module with Conv fusing."""
|
|
47
|
+
|
|
48
|
+
def __init__(self, c1, c2, k=3, s=1, p=None, g=1, d=1, act=True):
|
|
49
|
+
"""Initialize Conv layer with given arguments including activation."""
|
|
50
|
+
super().__init__(c1, c2, k, s, p, g=g, d=d, act=act)
|
|
51
|
+
self.cv2 = nn.Conv2d(c1, c2, 1, s, autopad(1, p, d), groups=g, dilation=d, bias=False) # add 1x1 conv
|
|
52
|
+
|
|
53
|
+
def forward(self, x):
|
|
54
|
+
"""Apply convolution, batch normalization and activation to input tensor."""
|
|
55
|
+
return self.act(self.bn(self.conv(x) + self.cv2(x)))
|
|
56
|
+
|
|
57
|
+
def fuse_convs(self):
|
|
58
|
+
"""Fuse parallel convolutions."""
|
|
59
|
+
w = torch.zeros_like(self.conv.weight.data)
|
|
60
|
+
i = [x // 2 for x in w.shape[2:]]
|
|
61
|
+
w[:, :, i[0]:i[0] + 1, i[1]:i[1] + 1] = self.cv2.weight.data.clone()
|
|
62
|
+
self.conv.weight.data += w
|
|
63
|
+
self.__delattr__('cv2')
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
class LightConv(nn.Module):
|
|
67
|
+
"""Light convolution with args(ch_in, ch_out, kernel).
|
|
68
|
+
https://github.com/PaddlePaddle/PaddleDetection/blob/develop/ppdet/modeling/backbones/hgnet_v2.py
|
|
69
|
+
"""
|
|
70
|
+
|
|
71
|
+
def __init__(self, c1, c2, k=1, act=nn.ReLU()):
|
|
72
|
+
"""Initialize Conv layer with given arguments including activation."""
|
|
73
|
+
super().__init__()
|
|
74
|
+
self.conv1 = Conv(c1, c2, 1, act=False)
|
|
75
|
+
self.conv2 = DWConv(c2, c2, k, act=act)
|
|
76
|
+
|
|
77
|
+
def forward(self, x):
|
|
78
|
+
"""Apply 2 convolutions to input tensor."""
|
|
79
|
+
return self.conv2(self.conv1(x))
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
class DWConv(Conv):
|
|
83
|
+
"""Depth-wise convolution."""
|
|
84
|
+
|
|
85
|
+
def __init__(self, c1, c2, k=1, s=1, d=1, act=True): # ch_in, ch_out, kernel, stride, dilation, activation
|
|
86
|
+
super().__init__(c1, c2, k, s, g=math.gcd(c1, c2), d=d, act=act)
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
class DWConvTranspose2d(nn.ConvTranspose2d):
|
|
90
|
+
"""Depth-wise transpose convolution."""
|
|
91
|
+
|
|
92
|
+
def __init__(self, c1, c2, k=1, s=1, p1=0, p2=0): # ch_in, ch_out, kernel, stride, padding, padding_out
|
|
93
|
+
super().__init__(c1, c2, k, s, p1, p2, groups=math.gcd(c1, c2))
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
class ConvTranspose(nn.Module):
|
|
97
|
+
"""Convolution transpose 2d layer."""
|
|
98
|
+
default_act = nn.SiLU() # default activation
|
|
99
|
+
|
|
100
|
+
def __init__(self, c1, c2, k=2, s=2, p=0, bn=True, act=True):
|
|
101
|
+
"""Initialize ConvTranspose2d layer with batch normalization and activation function."""
|
|
102
|
+
super().__init__()
|
|
103
|
+
self.conv_transpose = nn.ConvTranspose2d(c1, c2, k, s, p, bias=not bn)
|
|
104
|
+
self.bn = nn.BatchNorm2d(c2) if bn else nn.Identity()
|
|
105
|
+
self.act = self.default_act if act is True else act if isinstance(act, nn.Module) else nn.Identity()
|
|
106
|
+
|
|
107
|
+
def forward(self, x):
|
|
108
|
+
"""Applies transposed convolutions, batch normalization and activation to input."""
|
|
109
|
+
return self.act(self.bn(self.conv_transpose(x)))
|
|
110
|
+
|
|
111
|
+
def forward_fuse(self, x):
|
|
112
|
+
"""Applies activation and convolution transpose operation to input."""
|
|
113
|
+
return self.act(self.conv_transpose(x))
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
class Focus(nn.Module):
|
|
117
|
+
"""Focus wh information into c-space."""
|
|
118
|
+
|
|
119
|
+
def __init__(self, c1, c2, k=1, s=1, p=None, g=1, act=True): # ch_in, ch_out, kernel, stride, padding, groups
|
|
120
|
+
super().__init__()
|
|
121
|
+
self.conv = Conv(c1 * 4, c2, k, s, p, g, act=act)
|
|
122
|
+
# self.contract = Contract(gain=2)
|
|
123
|
+
|
|
124
|
+
def forward(self, x): # x(b,c,w,h) -> y(b,4c,w/2,h/2)
|
|
125
|
+
return self.conv(torch.cat((x[..., ::2, ::2], x[..., 1::2, ::2], x[..., ::2, 1::2], x[..., 1::2, 1::2]), 1))
|
|
126
|
+
# return self.conv(self.contract(x))
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
class GhostConv(nn.Module):
|
|
130
|
+
"""Ghost Convolution https://github.com/huawei-noah/ghostnet."""
|
|
131
|
+
|
|
132
|
+
def __init__(self, c1, c2, k=1, s=1, g=1, act=True): # ch_in, ch_out, kernel, stride, groups
|
|
133
|
+
super().__init__()
|
|
134
|
+
c_ = c2 // 2 # hidden channels
|
|
135
|
+
self.cv1 = Conv(c1, c_, k, s, None, g, act=act)
|
|
136
|
+
self.cv2 = Conv(c_, c_, 5, 1, None, c_, act=act)
|
|
137
|
+
|
|
138
|
+
def forward(self, x):
|
|
139
|
+
"""Forward propagation through a Ghost Bottleneck layer with skip connection."""
|
|
140
|
+
y = self.cv1(x)
|
|
141
|
+
return torch.cat((y, self.cv2(y)), 1)
|
|
142
|
+
|
|
143
|
+
|
|
144
|
+
class RepConv(nn.Module):
|
|
145
|
+
"""RepConv is a basic rep-style block, including training and deploy status
|
|
146
|
+
This code is based on https://github.com/DingXiaoH/RepVGG/blob/main/repvgg.py
|
|
147
|
+
"""
|
|
148
|
+
default_act = nn.SiLU() # default activation
|
|
149
|
+
|
|
150
|
+
def __init__(self, c1, c2, k=3, s=1, p=1, g=1, d=1, act=True, bn=False, deploy=False):
|
|
151
|
+
super().__init__()
|
|
152
|
+
assert k == 3 and p == 1
|
|
153
|
+
self.g = g
|
|
154
|
+
self.c1 = c1
|
|
155
|
+
self.c2 = c2
|
|
156
|
+
self.act = self.default_act if act is True else act if isinstance(act, nn.Module) else nn.Identity()
|
|
157
|
+
|
|
158
|
+
self.bn = nn.BatchNorm2d(num_features=c1) if bn and c2 == c1 and s == 1 else None
|
|
159
|
+
self.conv1 = Conv(c1, c2, k, s, p=p, g=g, act=False)
|
|
160
|
+
self.conv2 = Conv(c1, c2, 1, s, p=(p - k // 2), g=g, act=False)
|
|
161
|
+
|
|
162
|
+
def forward_fuse(self, x):
|
|
163
|
+
"""Forward process"""
|
|
164
|
+
return self.act(self.conv(x))
|
|
165
|
+
|
|
166
|
+
def forward(self, x):
|
|
167
|
+
"""Forward process"""
|
|
168
|
+
id_out = 0 if self.bn is None else self.bn(x)
|
|
169
|
+
return self.act(self.conv1(x) + self.conv2(x) + id_out)
|
|
170
|
+
|
|
171
|
+
def get_equivalent_kernel_bias(self):
|
|
172
|
+
kernel3x3, bias3x3 = self._fuse_bn_tensor(self.conv1)
|
|
173
|
+
kernel1x1, bias1x1 = self._fuse_bn_tensor(self.conv2)
|
|
174
|
+
kernelid, biasid = self._fuse_bn_tensor(self.bn)
|
|
175
|
+
return kernel3x3 + self._pad_1x1_to_3x3_tensor(kernel1x1) + kernelid, bias3x3 + bias1x1 + biasid
|
|
176
|
+
|
|
177
|
+
def _avg_to_3x3_tensor(self, avgp):
|
|
178
|
+
channels = self.c1
|
|
179
|
+
groups = self.g
|
|
180
|
+
kernel_size = avgp.kernel_size
|
|
181
|
+
input_dim = channels // groups
|
|
182
|
+
k = torch.zeros((channels, input_dim, kernel_size, kernel_size))
|
|
183
|
+
k[np.arange(channels), np.tile(np.arange(input_dim), groups), :, :] = 1.0 / kernel_size ** 2
|
|
184
|
+
return k
|
|
185
|
+
|
|
186
|
+
def _pad_1x1_to_3x3_tensor(self, kernel1x1):
|
|
187
|
+
if kernel1x1 is None:
|
|
188
|
+
return 0
|
|
189
|
+
else:
|
|
190
|
+
return torch.nn.functional.pad(kernel1x1, [1, 1, 1, 1])
|
|
191
|
+
|
|
192
|
+
def _fuse_bn_tensor(self, branch):
|
|
193
|
+
if branch is None:
|
|
194
|
+
return 0, 0
|
|
195
|
+
if isinstance(branch, Conv):
|
|
196
|
+
kernel = branch.conv.weight
|
|
197
|
+
running_mean = branch.bn.running_mean
|
|
198
|
+
running_var = branch.bn.running_var
|
|
199
|
+
gamma = branch.bn.weight
|
|
200
|
+
beta = branch.bn.bias
|
|
201
|
+
eps = branch.bn.eps
|
|
202
|
+
elif isinstance(branch, nn.BatchNorm2d):
|
|
203
|
+
if not hasattr(self, 'id_tensor'):
|
|
204
|
+
input_dim = self.c1 // self.g
|
|
205
|
+
kernel_value = np.zeros((self.c1, input_dim, 3, 3), dtype=np.float32)
|
|
206
|
+
for i in range(self.c1):
|
|
207
|
+
kernel_value[i, i % input_dim, 1, 1] = 1
|
|
208
|
+
self.id_tensor = torch.from_numpy(kernel_value).to(branch.weight.device)
|
|
209
|
+
kernel = self.id_tensor
|
|
210
|
+
running_mean = branch.running_mean
|
|
211
|
+
running_var = branch.running_var
|
|
212
|
+
gamma = branch.weight
|
|
213
|
+
beta = branch.bias
|
|
214
|
+
eps = branch.eps
|
|
215
|
+
std = (running_var + eps).sqrt()
|
|
216
|
+
t = (gamma / std).reshape(-1, 1, 1, 1)
|
|
217
|
+
return kernel * t, beta - running_mean * gamma / std
|
|
218
|
+
|
|
219
|
+
def fuse_convs(self):
|
|
220
|
+
if hasattr(self, 'conv'):
|
|
221
|
+
return
|
|
222
|
+
kernel, bias = self.get_equivalent_kernel_bias()
|
|
223
|
+
self.conv = nn.Conv2d(in_channels=self.conv1.conv.in_channels,
|
|
224
|
+
out_channels=self.conv1.conv.out_channels,
|
|
225
|
+
kernel_size=self.conv1.conv.kernel_size,
|
|
226
|
+
stride=self.conv1.conv.stride,
|
|
227
|
+
padding=self.conv1.conv.padding,
|
|
228
|
+
dilation=self.conv1.conv.dilation,
|
|
229
|
+
groups=self.conv1.conv.groups,
|
|
230
|
+
bias=True).requires_grad_(False)
|
|
231
|
+
self.conv.weight.data = kernel
|
|
232
|
+
self.conv.bias.data = bias
|
|
233
|
+
for para in self.parameters():
|
|
234
|
+
para.detach_()
|
|
235
|
+
self.__delattr__('conv1')
|
|
236
|
+
self.__delattr__('conv2')
|
|
237
|
+
if hasattr(self, 'nm'):
|
|
238
|
+
self.__delattr__('nm')
|
|
239
|
+
if hasattr(self, 'bn'):
|
|
240
|
+
self.__delattr__('bn')
|
|
241
|
+
if hasattr(self, 'id_tensor'):
|
|
242
|
+
self.__delattr__('id_tensor')
|
|
243
|
+
|
|
244
|
+
|
|
245
|
+
class ChannelAttention(nn.Module):
|
|
246
|
+
"""Channel-attention module https://github.com/open-mmlab/mmdetection/tree/v3.0.0rc1/configs/rtmdet."""
|
|
247
|
+
|
|
248
|
+
def __init__(self, channels: int) -> None:
|
|
249
|
+
super().__init__()
|
|
250
|
+
self.pool = nn.AdaptiveAvgPool2d(1)
|
|
251
|
+
self.fc = nn.Conv2d(channels, channels, 1, 1, 0, bias=True)
|
|
252
|
+
self.act = nn.Sigmoid()
|
|
253
|
+
|
|
254
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
255
|
+
return x * self.act(self.fc(self.pool(x)))
|
|
256
|
+
|
|
257
|
+
|
|
258
|
+
class SpatialAttention(nn.Module):
|
|
259
|
+
"""Spatial-attention module."""
|
|
260
|
+
|
|
261
|
+
def __init__(self, kernel_size=7):
|
|
262
|
+
"""Initialize Spatial-attention module with kernel size argument."""
|
|
263
|
+
super().__init__()
|
|
264
|
+
assert kernel_size in (3, 7), 'kernel size must be 3 or 7'
|
|
265
|
+
padding = 3 if kernel_size == 7 else 1
|
|
266
|
+
self.cv1 = nn.Conv2d(2, 1, kernel_size, padding=padding, bias=False)
|
|
267
|
+
self.act = nn.Sigmoid()
|
|
268
|
+
|
|
269
|
+
def forward(self, x):
|
|
270
|
+
"""Apply channel and spatial attention on input for feature recalibration."""
|
|
271
|
+
return x * self.act(self.cv1(torch.cat([torch.mean(x, 1, keepdim=True), torch.max(x, 1, keepdim=True)[0]], 1)))
|
|
272
|
+
|
|
273
|
+
|
|
274
|
+
class CBAM(nn.Module):
|
|
275
|
+
"""Convolutional Block Attention Module."""
|
|
276
|
+
|
|
277
|
+
def __init__(self, c1, kernel_size=7): # ch_in, kernels
|
|
278
|
+
super().__init__()
|
|
279
|
+
self.channel_attention = ChannelAttention(c1)
|
|
280
|
+
self.spatial_attention = SpatialAttention(kernel_size)
|
|
281
|
+
|
|
282
|
+
def forward(self, x):
|
|
283
|
+
"""Applies the forward pass through C1 module."""
|
|
284
|
+
return self.spatial_attention(self.channel_attention(x))
|
|
285
|
+
|
|
286
|
+
|
|
287
|
+
class Concat(nn.Module):
|
|
288
|
+
"""Concatenate a list of tensors along dimension."""
|
|
289
|
+
|
|
290
|
+
def __init__(self, dimension=1):
|
|
291
|
+
"""Concatenates a list of tensors along a specified dimension."""
|
|
292
|
+
super().__init__()
|
|
293
|
+
self.d = dimension
|
|
294
|
+
|
|
295
|
+
def forward(self, x):
|
|
296
|
+
"""Forward pass for the YOLOv8 mask Proto module."""
|
|
297
|
+
return torch.cat(x, self.d)
|