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
| import torch import torch.nn as nn import torch.nn.functional as F
class OmniGazeReward(nn.Module): """ OMNIGAZE 奖励启发的视线估计 论文核心:使用奖励机制引导模型学习跨角度泛化 奖励定义: - 高奖励:大角度场景准确估计 - 低奖励:正面场景准确估计(已容易) """ def __init__(self, backbone: str = "resnet18", reward_weights: dict = None): super().__init__() self.backbone = torch.hub.load('pytorch/vision:v0.10.0', backbone, pretrained=True) self.backbone = nn.Sequential(*list(self.backbone.children())[:-2]) self.gaze_head = nn.Sequential( nn.Linear(512 * 7 * 7, 512), nn.ReLU(inplace=True), nn.Linear(512, 2) ) self.reward_head = nn.Sequential( nn.Linear(512 * 7 * 7, 256), nn.ReLU(inplace=True), nn.Linear(256, 1) ) self.reward_weights = reward_weights or { "frontal": 1.0, "moderate": 2.0, "extreme": 3.0, } def forward(self, x: torch.Tensor) -> dict: """ 前向传播 Args: x: 输入图像, shape=(B, C, H, W) Returns: dict: { "gaze": (pitch, yaw), # 视线方向 "reward": float, # 奖励分数 "reward_weighted_loss": float } """ features = self.backbone(x) features_flat = features.view(features.size(0), -1) gaze = self.gaze_head(features_flat) reward = self.reward_head(features_flat) return { "gaze": gaze, "reward": reward } def compute_reward_weighted_loss(self, gaze_pred: torch.Tensor, gaze_gt: torch.Tensor, head_pose: torch.Tensor, reward: torch.Tensor) -> torch.Tensor: """ 奖励加权损失 Args: gaze_pred: 预测视线, shape=(B, 2) gaze_gt: 真实视线, shape=(B, 2) head_pose: 头部姿态角度, shape=(B, 3) (pitch, yaw, roll) reward: 奖励分数, shape=(B, 1) Returns: reward_weighted_loss: 奖励加权损失 """ gaze_error = F.mse_loss(gaze_pred, gaze_gt, reduction='none') pose_deviation = torch.norm(head_pose[:, :2], dim=1) reward_weight = torch.where( pose_deviation < 10, torch.tensor(self.reward_weights["frontal"]), torch.where( pose_deviation < 30, torch.tensor(self.reward_weights["moderate"]), torch.tensor(self.reward_weights["extreme"]) ) ) weighted_loss = gaze_error * reward_weight.unsqueeze(1) reward_loss = F.mse_loss(reward, reward_weight.unsqueeze(1)) total_loss = weighted_loss.mean() + 0.1 * reward_loss return total_loss
if __name__ == "__main__": model = OmniGazeReward(backbone="resnet18") x = torch.randn(4, 3, 224, 224) gaze_gt = torch.randn(4, 2) * 0.1 head_pose = torch.tensor([ [0, 0, 0], [15, 10, 0], [40, 30, 0], [5, 5, 0] ], dtype=torch.float32) output = model(x) loss = model.compute_reward_weighted_loss( output["gaze"], gaze_gt, head_pose, output["reward"] ) print(f"视线预测: {output['gaze']}") print(f"奖励分数: {output['reward']}") print(f"奖励加权损失: {loss.item():.4f}")
|