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
| import torch import torch.nn as nn import torch.nn.functional as F from torchvision.models import resnet50
class BYOL(nn.Module): """ Bootstrap Your Own Latent (BYOL) 自监督学习框架,无需负样本: 1. 在线网络:预测目标表示 2. 目标网络:EMA 更新的在线网络 3. 停止梯度 + 预测头 应用于驾驶员分心检测: - 预训练:大量无标签驾驶员图像 - 微调:少量标注数据 + 线性探测 """ def __init__(self, hidden_dim=256, projection_dim=256, prediction_dim=256, ema_decay=0.996): super().__init__() self.ema_decay = ema_decay self.online_encoder = resnet50(pretrained=False) self.online_encoder.fc = nn.Identity() self.online_projector = nn.Sequential( nn.Linear(2048, hidden_dim), nn.BatchNorm1d(hidden_dim), nn.ReLU(inplace=True), nn.Linear(hidden_dim, projection_dim) ) self.online_predictor = nn.Sequential( nn.Linear(projection_dim, hidden_dim), nn.BatchNorm1d(hidden_dim), nn.ReLU(inplace=True), nn.Linear(hidden_dim, projection_dim) ) self.target_encoder = self._copy_model(self.online_encoder) self.target_projector = self._copy_model(self.online_projector) for param in self.target_encoder.parameters(): param.requires_grad = False for param in self.target_projector.parameters(): param.requires_grad = False def _copy_model(self, model): """创建模型副本""" copy = type(model)(*list(model.parameters())[0:0]) copy.load_state_dict(model.state_dict()) return copy def forward(self, x1, x2): """ Args: x1, x2: 同一图像的两个增强视图 Returns: loss: BYOL 损失 """ online_proj1 = self.online_projector(self.online_encoder(x1)) online_pred1 = self.online_predictor(online_proj1) online_proj2 = self.online_projector(self.online_encoder(x2)) online_pred2 = self.online_predictor(online_proj2) with torch.no_grad(): target_proj1 = self.target_projector(self.target_encoder(x1)) target_proj2 = self.target_projector(self.target_encoder(x2)) loss = 2 - 2 * ( F.cosine_similarity(online_pred1, target_proj2.detach(), dim=-1).mean() + F.cosine_similarity(online_pred2, target_proj1.detach(), dim=-1).mean() ) / 2 return loss @torch.no_grad() def update_target(self): """EMA 更新目标网络""" for online_p, target_p in zip( self.online_encoder.parameters(), self.target_encoder.parameters() ): target_p.data.mul_(self.ema_decay).add_( online_p.data, alpha=1 - self.ema_decay ) for online_p, target_p in zip( self.online_projector.parameters(), self.target_projector.parameters() ): target_p.data.mul_(self.ema_decay).add_( online_p.data, alpha=1 - self.ema_decay )
class DriverDistractionClassifier(nn.Module): """ 驾驶员分心分类器(线性探测) 使用 BYOL 预训练的编码器 + 线性分类头 """ def __init__(self, byol_encoder, n_classes=10): super().__init__() self.encoder = byol_encoder for param in self.encoder.parameters(): param.requires_grad = False self.classifier = nn.Linear(2048, n_classes) def forward(self, x): with torch.no_grad(): features = self.encoder(x) return self.classifier(features)
DISTRACTION_CLASSES = [ "safe_driving", "phone_right", "phone_left", "text_right", "text_left", "adjusting_radio", "drinking", "reaching_behind", "hair_makeup", "talking_passenger" ]
if __name__ == "__main__": byol = BYOL() x1 = torch.randn(32, 3, 224, 224) x2 = torch.randn(32, 3, 224, 224) loss = byol(x1, x2) print(f"BYOL Loss: {loss.item():.4f}") byol.update_target() print("Target network updated (EMA)") classifier = DriverDistractionClassifier(byol.online_encoder, n_classes=10) x = torch.randn(16, 3, 224, 224) logits = classifier(x) print(f"Input: {x.shape}") print(f"Output: {logits.shape} (10 类分心行为)") print(f"可训练参数: {sum(p.numel() for p in classifier.parameters() if p.requires_grad):,}")
|