File size: 6,088 Bytes
6a5bb7e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
"""注意力注入 — 通过 ultralytics callback 在训练模型创建后注入

核心思路:不在 yaml 里插入 ECA/CBAM 层(会导致层号偏移、权重错位),
而是在训练开始后,用 PyTorch forward_hook 给指定层的输出加注意力。

关键修复:必须用 callback 在 trainer 创建模型之后注入 hooks,
因为 YOLO.train() 内部通过 get_model() 会创建全新的模型对象,
直接在 model 上注册的 hooks 会丢失。

优势:
  1. 层号不变 → model.load() 权重 100% 正确迁移
  2. 不影响模型结构 → 完全兼容 ultralytics
  3. 通过 callback 确保 hooks 在训练模型上生效
"""
import math
import torch
import torch.nn as nn


class AttentionHook:
    """通过 forward hook 给指定层注入注意力,不改变模型结构。"""

    def __init__(self, attention_module):
        self.attn = attention_module
        self.handle = None

    def hook_fn(self, module, input, output):
        return self.attn(output)

    def register(self, target_module):
        self.handle = target_module.register_forward_hook(self.hook_fn)
        return self

    def remove(self):
        if self.handle:
            self.handle.remove()


class ECABlock(nn.Module):
    """ECA 注意力块 — 渐进式残差"""
    def __init__(self, channels, gamma=2, b=1, init_alpha=0.01):
        super().__init__()
        t = int(abs((math.log2(channels) + b) / gamma))
        k = t if t % 2 else t + 1
        self.avg_pool = nn.AdaptiveAvgPool2d(1)
        self.conv = nn.Conv1d(1, 1, kernel_size=k, padding=k // 2, bias=False)
        self.sigmoid = nn.Sigmoid()
        self.alpha = nn.Parameter(torch.full((1,), init_alpha))

    def forward(self, x):
        y = self.avg_pool(x)
        y = self.conv(y.squeeze(-1).transpose(-1, -2))
        y = y.transpose(-1, -2).unsqueeze(-1)
        attn = self.sigmoid(y)
        return x * (1.0 + self.alpha * (attn - 1.0))


class CBAMBlock(nn.Module):
    """CBAM 注意力块 — 渐进式残差"""
    def __init__(self, channels, reduction=16, kernel_size=7, init_alpha=0.01):
        super().__init__()
        mid = max(channels // reduction, 8)
        self.avg_pool = nn.AdaptiveAvgPool2d(1)
        self.max_pool = nn.AdaptiveMaxPool2d(1)
        self.fc = nn.Sequential(
            nn.Conv2d(channels, mid, 1, bias=False),
            nn.ReLU(inplace=True),
            nn.Conv2d(mid, channels, 1, bias=False),
        )
        self.sa_conv = nn.Conv2d(2, 1, kernel_size, padding=kernel_size // 2, bias=False)
        self.sigmoid = nn.Sigmoid()
        self.alpha = nn.Parameter(torch.full((1,), init_alpha))

    def forward(self, x):
        ca = self.sigmoid(self.fc(self.avg_pool(x)) + self.fc(self.max_pool(x)))
        avg_out = torch.mean(x, dim=1, keepdim=True)
        max_out, _ = torch.max(x, dim=1, keepdim=True)
        sa = self.sigmoid(self.sa_conv(torch.cat([avg_out, max_out], dim=1)))
        return x * (1.0 + self.alpha * (ca * sa - 1.0))


def _get_out_channels(layer):
    """获取层的输出通道数。"""
    if hasattr(layer, 'cv2'):
        return layer.cv2.conv.out_channels
    elif hasattr(layer, 'conv'):
        return layer.conv.out_channels
    else:
        for p in layer.parameters():
            return p.shape[0]
    raise ValueError(f"Cannot determine output channels for {layer.__class__.__name__}")


def setup_attention_training(model, layer_indices, attn_type="cbam"):
    """通过 ultralytics callback 注入注意力,确保在训练模型上生效。

    Args:
        model: ultralytics YOLO model
        layer_indices: 要加注意力的层号,如 [2, 4](C3k2 层)
        attn_type: "eca" 或 "cbam"
    """
    _state = {"hooks": [], "attn_modules": None, "injected": False}

    def _on_pretrain_routine_end(trainer):
        """在 trainer 完成模型+优化器设置后注入注意力。"""
        if _state["injected"]:
            return

        training_model = trainer.model
        if hasattr(training_model, 'module'):
            training_model = training_model.module

        device = next(training_model.parameters()).device
        attn_modules = nn.ModuleList()

        for idx in layer_indices:
            layer = training_model.model[idx]
            out_ch = _get_out_channels(layer)

            if attn_type == "cbam":
                attn = CBAMBlock(out_ch)
            else:
                attn = ECABlock(out_ch)

            attn = attn.to(device)
            hook = AttentionHook(attn).register(layer)
            _state["hooks"].append(hook)
            attn_modules.append(attn)

            print(f"[Attention] Layer {idx} ({layer.__class__.__name__}, ch={out_ch}) ← {attn_type.upper()}")

        _state["attn_modules"] = attn_modules

        attn_params = list(attn_modules.parameters())
        if attn_params and trainer.optimizer is not None:
            base_lr = trainer.optimizer.param_groups[0]["lr"]
            initial_lr = trainer.optimizer.param_groups[0].get("initial_lr", base_lr)
            param_group = {
                "params": attn_params,
                "lr": base_lr,
                "initial_lr": initial_lr,
                "weight_decay": trainer.optimizer.param_groups[0].get("weight_decay", 0),
            }
            trainer.optimizer.add_param_group(param_group)

            if hasattr(trainer, 'scheduler') and trainer.scheduler is not None:
                trainer.scheduler.base_lrs.append(initial_lr)
                if hasattr(trainer.scheduler, 'lr_lambdas'):
                    trainer.scheduler.lr_lambdas.append(trainer.scheduler.lr_lambdas[0])

            print(f"[Attention] 已添加 {len(attn_params)} 个注意力参数到优化器")

        _state["injected"] = True

    def _on_train_end(trainer):
        """训练结束后清理 hooks。"""
        for h in _state["hooks"]:
            h.remove()
        print("[Attention] Hooks 已清理")

    model.add_callback("on_pretrain_routine_end", _on_pretrain_routine_end)
    model.add_callback("on_train_end", _on_train_end)

    return _state