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
| class LearnToRank(nn.Module): """ 学习排序模块:估计每个变换预测的可靠性 在源域训练,目标域推理时固定不变 """ def __init__(self, feature_dim: int, hidden_dim: int = 128): super().__init__() self.reliability_net = nn.Sequential( nn.Linear(feature_dim * 2, hidden_dim), nn.ReLU(), nn.Dropout(0.3), nn.Linear(hidden_dim, hidden_dim // 2), nn.ReLU(), nn.Linear(hidden_dim // 2, 1), nn.Sigmoid() ) def forward(self, original_feat: torch.Tensor, transformed_feat: torch.Tensor) -> torch.Tensor: """ Args: original_feat: 原始样本特征 (B, D) transformed_feat: 变换后样本特征 (B, D) Returns: reliability: 可靠性权重 (B, 1) """ diff = original_feat - transformed_feat concat = torch.cat([original_feat, diff], dim=-1) reliability = self.reliability_net(concat) return reliability
class BFTInference(nn.Module): """ BFT 完整推理管道 对单个测试样本: 1. 生成 N 个变换版本 2. 每个版本通过模型前向传播 3. 排序模块评估可靠性 4. 加权聚合输出最终预测 """ def __init__(self, model: nn.Module, transformer: BFTTransform, ranker: LearnToRank, feature_layer: str = "features"): super().__init__() self.model = model self.transformer = transformer self.ranker = ranker self.feature_layer = feature_layer def forward(self, x: torch.Tensor) -> torch.Tensor: """ 无需梯度更新的前向推理 Args: x: 输入EEG信号 (B, C, T) Returns: output: 加权聚合后的预测 (B, num_classes) """ original_feat = self.model.extract_features(x) original_pred = self.model.classifier(original_feat) variants = self.transformer(x) all_preds = [original_pred] all_weights = [torch.ones(x.shape[0], 1, device=x.device)] for variant in variants: var_feat = self.model.extract_features(variant) var_pred = self.model.classifier(var_feat) with torch.no_grad(): weight = self.ranker(original_feat.detach(), var_feat.detach()) all_preds.append(var_pred) all_weights.append(weight) all_preds = torch.stack(all_preds, dim=0) all_weights = torch.cat(all_weights, dim=1) all_weights = all_weights / (all_weights.sum(dim=-1, keepdim=True) + 1e-8) output = (all_preds * all_weights.unsqueeze(-1).unsqueeze(0)).sum(dim=0) return output
if __name__ == "__main__": print("=" * 60) print("BFT 完整推理管道测试") print("=" * 60) class SimpleEEGModel(nn.Module): def __init__(self, channels=32, num_classes=3): super().__init__() self.feature_extractor = nn.Sequential( nn.Conv1d(channels, 64, 7, stride=2, padding=3), nn.BatchNorm1d(64), nn.ReLU(), nn.Conv1d(64, 128, 5, stride=2, padding=2), nn.BatchNorm1d(128), nn.ReLU(), nn.AdaptiveAvgPool1d(1) ) self.classifier_head = nn.Sequential( nn.Linear(128, 64), nn.ReLU(), nn.Linear(64, num_classes) ) def extract_features(self, x): feat = self.feature_extractor(x).squeeze(-1) return feat def classifier(self, feat): return self.classifier_head(feat) model = SimpleEEGModel(channels=32, num_classes=3) transformer = BFTTransform(num_transforms=8) ranker = LearnToRank(feature_dim=128) bft = BFTInference(model, transformer, ranker) test_eeg = torch.randn(4, 32, 256) with torch.no_grad(): traditional_pred = model.classifier(model.extract_features(test_eeg)) bft_pred = bft(test_eeg) print(f"\n输入 shape: {test_eeg.shape}") print(f"传统预测 shape: {traditional_pred.shape}") print(f"BFT预测 shape: {bft_pred.shape}") print(f"\n传统预测 (sample 0): {traditional_pred[0].tolist()}") print(f"BFT预测 (sample 0): {bft_pred[0].tolist()}") print(f"\n预测差异: {(bft_pred - traditional_pred).abs().mean():.4f}") print("✅ BFT 通过前向传播完成适应,无需反向传播")
|