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
| import torch import torch.nn as nn
class EdgeHARModel(nn.Module): """ EdgeHAR 三码解纠缠模型 适配到DMS场景: - 输入:DMS特征序列(人脸关键点/眼动/头部姿态) - 输出:三个独立潜在码 + 行为分类 优势: - 新用户适配:只调Acquisition-Context Code - 新车型适配:只调Acquisition-Context Code - 活动知识通用:Activity-Semantic Code不变 """ def __init__(self, input_dim: int = 68*2, hidden_dim: int = 128, code_dims: dict = None): super().__init__() if code_dims is None: code_dims = {'semantic': 32, 'dynamics': 32, 'context': 16} self.encoder = nn.Sequential( nn.Conv1d(input_dim, hidden_dim, 3, padding=1), nn.BatchNorm1d(hidden_dim), nn.ReLU(), nn.Conv1d(hidden_dim, hidden_dim, 3, padding=1), nn.BatchNorm1d(hidden_dim), nn.ReLU(), ) self.semantic_head = nn.Linear(hidden_dim, code_dims['semantic']) self.dynamics_head = nn.LSTM(hidden_dim, code_dims['dynamics'], batch_first=True) self.context_head = nn.Linear(hidden_dim, code_dims['context']) self.decoder = nn.Sequential( nn.Linear(code_dims['semantic'] + code_dims['dynamics'] + code_dims['context'], hidden_dim), nn.ReLU(), nn.ConvTranspose1d(hidden_dim, input_dim, 3, padding=1), ) self.classifier = nn.Sequential( nn.Linear(code_dims['semantic'] + code_dims['dynamics'], 64), nn.ReLU(), nn.Linear(64, 7), ) def forward(self, x): """ Args: x: (B, T, input_dim) 时序特征 Returns: reconstruction, classification, codes """ encoded = self.encoder(x.transpose(1, 2)).transpose(1, 2) semantic = self.semantic_head(encoded) dynamics_out, _ = self.dynamics_head(encoded) context = self.context_head(encoded) combined = torch.cat([semantic, dynamics_out, context], dim=-1) reconstructed = self.decoder(combined.transpose(1, 2)).transpose(1, 2) class_input = torch.cat([ semantic.mean(dim=1), dynamics_out.mean(dim=1) ], dim=-1) classification = self.classifier(class_input) return reconstructed, classification, { 'semantic': semantic, 'dynamics': dynamics_out, 'context': context }
|