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
| """ AttentiveGaze 注意力增强眼部特征提取器 空间注意力 + 通道注意力,学习眼部最诊断性的区域 """
import torch import torch.nn as nn import torch.nn.functional as F
class SpatialAttention(nn.Module): """空间注意力:学习眼部图像中哪些空间位置最重要""" def __init__(self, kernel_size: int = 7): super().__init__() self.conv = nn.Conv2d(2, 1, kernel_size, padding=kernel_size//2) def forward(self, x): avg_pool = torch.mean(x, dim=1, keepdim=True) max_pool, _ = torch.max(x, dim=1, keepdim=True) combined = torch.cat([avg_pool, max_pool], dim=1) attention = torch.sigmoid(self.conv(combined)) return x * attention
class ChannelAttention(nn.Module): """通道注意力:学习哪些特征通道最重要""" def __init__(self, channels: int, reduction: int = 16): super().__init__() self.avg_pool = nn.AdaptiveAvgPool2d(1) self.max_pool = nn.AdaptiveMaxPool2d(1) self.mlp = nn.Sequential( nn.Linear(channels, channels // reduction), nn.ReLU(), nn.Linear(channels // reduction, channels) ) def forward(self, x): b, c, _, _ = x.shape avg_out = self.mlp(self.avg_pool(x).view(b, c)) max_out = self.mlp(self.max_pool(x).view(b, c)) attention = torch.sigmoid(avg_out + max_out) return x * attention.view(b, c, 1, 1)
class AttentionEnhancedEyeEncoder(nn.Module): """ AttentiveGaze 眼部特征提取器 结合空间注意力和通道注意力, 自适应聚焦虹膜边界、角膜反射等诊断性区域 """ def __init__(self, in_channels: int = 3, feat_channels: int = 64): super().__init__() self.conv1 = nn.Conv2d(in_channels, feat_channels, 3, padding=1) self.bn1 = nn.BatchNorm2d(feat_channels) self.conv2 = nn.Conv2d(feat_channels, feat_channels*2, 3, stride=2, padding=1) self.bn2 = nn.BatchNorm2d(feat_channels*2) self.conv3 = nn.Conv2d(feat_channels*2, feat_channels*4, 3, stride=2, padding=1) self.bn3 = nn.BatchNorm2d(feat_channels*4) self.ca1 = ChannelAttention(feat_channels) self.sa1 = SpatialAttention() self.ca2 = ChannelAttention(feat_channels*2) self.sa2 = SpatialAttention() self.ca3 = ChannelAttention(feat_channels*4) self.sa3 = SpatialAttention() def forward(self, x): x = F.relu(self.bn1(self.conv1(x))) x = self.ca1(x) x = self.sa1(x) x = F.relu(self.bn2(self.conv2(x))) x = self.ca2(x) x = self.sa2(x) x = F.relu(self.bn3(self.conv3(x))) x = self.ca3(x) x = self.sa3(x) x = F.adaptive_avg_pool2d(x, 1).flatten(1) return x
if __name__ == "__main__": encoder = AttentionEnhancedEyeEncoder(in_channels=3, feat_channels=32) eye_img = torch.randn(4, 3, 64, 64) features = encoder(eye_img) print(f"输入: {eye_img.shape}") print(f"输出特征: {features.shape}") print(f"参数量: {sum(p.numel() for p in encoder.parameters()):,}")
|