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
| class MultiDatasetTrainer: """ 多数据集联合训练策略 数据集: - ETH-XGaze: 极端头部姿态 - MPIIFaceGaze: 真实场景 - EYEDIAP: 屏幕-眼动协调 """ def __init__(self, model, datasets, batch_size=64): self.model = model self.datasets = datasets self.batch_size = batch_size total_samples = sum(len(d) for d in datasets) self.weights = [len(d) / total_samples for d in datasets] def train_epoch(self, optimizer): """ 一个epoch的训练 Returns: avg_loss: 平均损失 metrics: 各数据集指标 """ self.model.train() total_loss = 0 metrics = {f'dataset_{i}': {} for i in range(len(self.datasets))} iterators = [iter(torch.utils.data.DataLoader( d, batch_size=self.batch_size, shuffle=True, drop_last=True )) for d in self.datasets] num_batches = max(len(d) // self.batch_size for d in self.datasets) for batch_idx in range(num_batches): optimizer.zero_grad() batch_loss = 0 for i, (iterator, weight) in enumerate(zip(iterators, self.weights)): try: images, gaze_gt = next(iterator) except StopIteration: continue gaze_pred = self.model(images) loss = self._angular_loss(gaze_pred, gaze_gt) batch_loss += weight * loss with torch.no_grad(): angle_error = self._compute_angular_error(gaze_pred, gaze_gt) metrics[f'dataset_{i}']['angle_error'] = angle_error batch_loss.backward() optimizer.step() total_loss += batch_loss.item() return total_loss / num_batches, metrics def _angular_loss(self, pred, gt): """ 角度损失函数 L = arccos(pred · gt) """ cos_sim = F.cosine_similarity(pred, gt, dim=-1) cos_sim = torch.clamp(cos_sim, -1.0 + 1e-7, 1.0 - 1e-7) return torch.mean(torch.acos(cos_sim)) def _compute_angular_error(self, pred, gt): """计算平均角度误差(度)""" cos_sim = F.cosine_similarity(pred, gt, dim=-1) cos_sim = torch.clamp(cos_sim, -1.0, 1.0) angle_rad = torch.acos(cos_sim) angle_deg = torch.rad2deg(angle_rad) return angle_deg.mean().item()
|