| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
|
|
| class MiniConvEmbedder(nn.Module): |
| def __init__(self): |
| super(MiniConvEmbedder, self).__init__() |
| |
| |
| self.conv1 = nn.Conv2d(1, 16, kernel_size=3, padding=0) |
| self.conv2 = nn.Conv2d(16, 32, kernel_size=3, padding=0) |
| self.conv3 = nn.Conv2d(32, 64, kernel_size=3, padding=0) |
| self.relu = nn.ReLU(inplace=True) |
| self.gap = nn.AdaptiveAvgPool2d(1) |
|
|
| def forward(self, x): |
| |
| x = self.relu(self.conv1(x)) |
| x = self.relu(self.conv2(x)) |
| x = self.relu(self.conv3(x)) |
| x = self.gap(x) |
| x = torch.flatten(x, 1) |
| return x |
|
|
| class GradientReversalLayer(torch.autograd.Function): |
| @staticmethod |
| def forward(ctx, x, alpha): |
| ctx.alpha = alpha |
| return x.view_as(x) |
|
|
| @staticmethod |
| def backward(ctx, grad_output): |
| return grad_output.neg() * ctx.alpha, None |
|
|
| class AdaptiveLayerNorm(nn.Module): |
| def __init__(self, num_features, num_domains=2): |
| super(AdaptiveLayerNorm, self).__init__() |
| self.num_features = num_features |
| self.norm = nn.LayerNorm(num_features, elementwise_affine=False) |
| self.gamma = nn.Parameter(torch.ones(num_domains, num_features)) |
| self.beta = nn.Parameter(torch.zeros(num_domains, num_features)) |
|
|
| def forward(self, x, domain_id): |
| |
| |
| x = self.norm(x) |
| |
| gamma = self.gamma[domain_id] |
| beta = self.beta[domain_id] |
| return x * gamma + beta |
|
|
| class LIPEV2Student(nn.Module): |
| def __init__(self): |
| super(LIPEV2Student, self).__init__() |
| |
| |
| self.appearance_net = MiniConvEmbedder() |
| |
| |
| self.geo_mlp1 = nn.Linear(956, 256) |
| self.ada_ln = AdaptiveLayerNorm(256, num_domains=2) |
| self.geo_mlp2 = nn.Sequential( |
| nn.ReLU(inplace=True), |
| nn.Dropout(0.05), |
| nn.Linear(256, 256), |
| nn.ReLU(inplace=True) |
| ) |
| |
| |
| self.fusion_mlp = nn.Sequential( |
| nn.Linear(512, 256), |
| nn.ReLU(inplace=True), |
| nn.Dropout(0.05) |
| ) |
| |
| |
| self.pitch_head = nn.Sequential( |
| nn.Linear(256, 64), |
| nn.ReLU(inplace=True), |
| nn.Linear(64, 90) |
| ) |
| |
| self.yaw_head = nn.Sequential( |
| nn.Linear(256, 64), |
| nn.ReLU(inplace=True), |
| nn.Linear(64, 90) |
| ) |
|
|
| |
| self.domain_classifier = nn.Sequential( |
| nn.Linear(256, 128), |
| nn.ReLU(inplace=True), |
| nn.Dropout(0.1), |
| nn.Linear(128, 2) |
| ) |
| |
| |
| self._init_weights() |
|
|
| def _init_weights(self): |
| for m in self.modules(): |
| if isinstance(m, nn.Conv2d) or isinstance(m, nn.Linear): |
| nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu') |
| if m.bias is not None: |
| nn.init.constant_(m.bias, 0) |
|
|
| def forward(self, patches=None, landmarks=None, state='A', alpha=0.0, domain_id=None): |
| """ |
| Asymmetric forward pass. |
| alpha: GRL hyperparameter (used during training for DANN) |
| domain_id: Used for AdaLN (Long tensor of shape Batch) |
| """ |
| if domain_id is None: |
| |
| batch_size = landmarks.shape[0] if landmarks is not None else patches.shape[0] |
| domain_id = torch.zeros(batch_size, dtype=torch.long, device=landmarks.device) |
|
|
| |
| geo_feat = self.geo_mlp1(landmarks) |
| geo_feat = self.ada_ln(geo_feat, domain_id) |
| geo_feat = self.geo_mlp2(geo_feat) |
| |
| if state == 'A' and patches is not None: |
| |
| batch_size = patches.shape[0] |
| patch_h, patch_w = patches.shape[2], patches.shape[3] |
| patches = patches.view(-1, 1, patch_h, patch_w) |
| app_tokens = self.appearance_net(patches) |
| app_feat = app_tokens.view(batch_size, -1) |
| |
| |
| combined = torch.cat([app_feat, geo_feat], dim=1) |
| combined = self.fusion_mlp(combined) |
| else: |
| combined = geo_feat |
| |
| |
| |
| reverse_feature = GradientReversalLayer.apply(combined, alpha) |
| domain_logits = self.domain_classifier(reverse_feature) |
| |
| |
| pitch_logits = self.pitch_head(combined) |
| yaw_logits = self.yaw_head(combined) |
| |
| return pitch_logits, yaw_logits, domain_logits |
|
|
| |
| class LIPEV2StudentBaseline(nn.Module): |
| def __init__(self): |
| super(LIPEV2StudentBaseline, self).__init__() |
| self.appearance_net = MiniConvEmbedder() |
| self.geo_mlp = nn.Sequential( |
| nn.Linear(956, 256), |
| nn.LayerNorm(256), |
| nn.ReLU(inplace=True), |
| nn.Dropout(0.05), |
| nn.Linear(256, 256), |
| nn.ReLU(inplace=True) |
| ) |
| self.pitch_head = nn.Sequential(nn.Linear(256, 64), nn.ReLU(inplace=True), nn.Linear(64, 90)) |
| self.yaw_head = nn.Sequential(nn.Linear(256, 64), nn.ReLU(inplace=True), nn.Linear(64, 90)) |
|
|
| def forward(self, patches=None, landmarks=None, state='A'): |
| geo_feat = self.geo_mlp(landmarks) |
| if state == 'A' and patches is not None: |
| batch_size = patches.shape[0] |
| patches = patches.view(-1, 1, patches.shape[2], patches.shape[3]) |
| app_tokens = self.appearance_net(patches) |
| app_feat = app_tokens.view(batch_size, -1) |
| combined = app_feat + geo_feat |
| else: |
| combined = geo_feat |
| return self.pitch_head(combined), self.yaw_head(combined) |
|
|
| |
| class DualPoolMiniConv(nn.Module): |
| def __init__(self): |
| super(DualPoolMiniConv, self).__init__() |
| self.conv = nn.Sequential( |
| nn.Conv2d(1, 16, kernel_size=3, padding=0), nn.ReLU(inplace=True), |
| nn.Conv2d(16, 32, kernel_size=3, padding=0), nn.ReLU(inplace=True), |
| nn.Conv2d(32, 64, kernel_size=3, padding=0), nn.ReLU(inplace=True) |
| ) |
| self.avg_pool = nn.AdaptiveAvgPool2d(1) |
| self.max_pool = nn.AdaptiveMaxPool2d(1) |
|
|
| def forward(self, x): |
| x = self.conv(x) |
| return torch.cat([self.avg_pool(x), self.max_pool(x)], dim=1).flatten(1) |
|
|
| class LIPEV2StudentGold(nn.Module): |
| def __init__(self): |
| super(LIPEV2StudentGold, self).__init__() |
| self.app_net = DualPoolMiniConv() |
| |
| self.geo_net = nn.Sequential( |
| nn.Linear(956, 256), |
| nn.LayerNorm(256), |
| nn.ReLU(inplace=True), |
| nn.Linear(256, 256), |
| nn.ReLU(inplace=True) |
| ) |
| |
| self.post_concat_bn = nn.BatchNorm1d(512 + 256) |
| |
| self.fusion = nn.Sequential( |
| nn.Linear(768, 256), |
| nn.ReLU(inplace=True), |
| nn.Dropout(0.1), |
| nn.Linear(256, 128), |
| nn.ReLU(inplace=True) |
| ) |
| |
| self.pitch_head = nn.Linear(128, 90) |
| self.yaw_head = nn.Linear(128, 90) |
|
|
| def forward(self, patches, landmarks): |
| batch_size = patches.shape[0] |
| p_h, p_w = patches.shape[2], patches.shape[3] |
| app_feat = self.app_net(patches.view(-1, 1, p_h, p_w)).view(batch_size, -1) |
| geo_feat = self.geo_net(landmarks) |
| |
| combined = torch.cat([app_feat, geo_feat], dim=1) |
| combined = self.post_concat_bn(combined) |
| fused = self.fusion(combined) |
| return self.pitch_head(fused), self.yaw_head(fused) |
|
|
| |
| class LIPEV2StudentGaze360Gold(nn.Module): |
| def __init__(self, num_domains=2): |
| super(LIPEV2StudentGaze360Gold, self).__init__() |
| self.app_net = DualPoolMiniConv() |
| |
| self.geo_mlp1 = nn.Linear(956, 256) |
| self.ada_ln = AdaptiveLayerNorm(256, num_domains=num_domains) |
| self.geo_mlp2 = nn.Sequential( |
| nn.ReLU(inplace=True), |
| nn.Linear(256, 256), |
| nn.ReLU(inplace=True) |
| ) |
| |
| self.post_concat_bn = nn.BatchNorm1d(512 + 256) |
| |
| self.fusion = nn.Sequential( |
| nn.Linear(768, 256), |
| nn.ReLU(inplace=True), |
| nn.Dropout(0.1), |
| nn.Linear(256, 128), |
| nn.ReLU(inplace=True) |
| ) |
| |
| self.pitch_head = nn.Linear(128, 90) |
| self.yaw_head = nn.Linear(128, 90) |
|
|
| self.domain_classifier = nn.Sequential( |
| nn.Linear(128, 128), |
| nn.ReLU(inplace=True), |
| nn.Dropout(0.1), |
| nn.Linear(128, num_domains) |
| ) |
|
|
| def forward(self, patches=None, landmarks=None, state='A', alpha=0.0, domain_id=None): |
| batch_size = landmarks.shape[0] if landmarks is not None else patches.shape[0] |
| if domain_id is None: |
| domain_id = torch.zeros(batch_size, dtype=torch.long, device=landmarks.device) |
|
|
| geo_feat = self.geo_mlp1(landmarks) |
| geo_feat = self.ada_ln(geo_feat, domain_id) |
| geo_feat = self.geo_mlp2(geo_feat) |
| |
| if state == 'A' and patches is not None: |
| p_h, p_w = patches.shape[2], patches.shape[3] |
| app_feat = self.app_net(patches.view(-1, 1, p_h, p_w)).view(batch_size, -1) |
| combined = torch.cat([app_feat, geo_feat], dim=1) |
| combined = self.post_concat_bn(combined) |
| fused = self.fusion(combined) |
| else: |
| fused = self.fusion(torch.cat([torch.zeros(batch_size, 512, device=geo_feat.device), geo_feat], dim=1)) |
| |
|
|
| |
| reverse_feature = GradientReversalLayer.apply(fused, alpha) |
| domain_logits = self.domain_classifier(reverse_feature) |
|
|
| return self.pitch_head(fused), self.yaw_head(fused), domain_logits |
|
|
| |
| class LIPEFinalAppearance(nn.Module): |
| def __init__(self): |
| super(LIPEFinalAppearance, self).__init__() |
| |
| self.conv1 = nn.Conv2d(1, 32, kernel_size=3, padding=1) |
| |
| self.conv2 = nn.Conv2d(32, 64, kernel_size=3, stride=2, padding=1) |
| |
| self.conv3 = nn.Conv2d(64, 128, kernel_size=3, stride=2, padding=1) |
| |
| self.conv4 = nn.Conv2d(128, 256, kernel_size=3, padding=1) |
| |
| self.relu = nn.ReLU(inplace=True) |
| self.gap = nn.AdaptiveAvgPool2d(1) |
| self.gmp = nn.AdaptiveMaxPool2d(1) |
| |
| |
| self.proj = nn.Linear(512, 512) |
|
|
| def forward(self, x): |
| |
| x = self.relu(self.conv1(x)) |
| x = self.relu(self.conv2(x)) |
| x = self.relu(self.conv3(x)) |
| x = self.relu(self.conv4(x)) |
| |
| avg_f = self.gap(x).view(-1, 256) |
| max_f = self.gmp(x).view(-1, 256) |
| combined = torch.cat([avg_f, max_f], dim=1) |
| |
| out = self.proj(combined) |
| return out |
|
|
| class LIPEV2StudentFinal(nn.Module): |
| def __init__(self): |
| super(LIPEV2StudentFinal, self).__init__() |
| self.app_net = LIPEFinalAppearance() |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| self.fusion = nn.Linear(512 + 2, 64) |
| |
| |
| |
| |
| |
| self.regression = nn.Linear(64, 2) |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| self.fusion = nn.Linear(512 + 128, 512) |
| self.regression = nn.Linear(512, 2) |
| |
| def forward(self, patches, landmarks): |
| |
| batch_size = patches.shape[0] |
| app_feat = self.app_net(patches.view(-1, 1, 16, 16)).view(batch_size, -1) |
| |
| |
| geo_feat = torch.zeros(batch_size, 128, device=patches.device) |
| |
| fused = self.fusion(torch.cat([app_feat, geo_feat], dim=1)) |
| out = self.regression(fused) |
| return out |
|
|
| if __name__ == '__main__': |
| |
| model = LIPEV2Student() |
| dummy_patches = torch.randn(8, 4, 8, 8) |
| dummy_landmarks = torch.randn(8, 956) |
| |
| |
| p_a, y_a = model(dummy_patches, dummy_landmarks, state='A') |
| print(f"State A Output Shapes: Pitch {p_a.shape}, Yaw {y_a.shape}") |
| |
| |
| p_b, y_b = model(None, dummy_landmarks, state='B') |
| print(f"State B Output Shapes: Pitch {p_b.shape}, Yaw {y_b.shape}") |
| |
| |
| total_params = sum(p.numel() for p in model.parameters()) |
| print(f"Total Parameters: {total_params:,}") |
|
|