File size: 4,337 Bytes
bc971c7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import torch.nn as nn
import torch
import torchvision.models as models
import torch.nn.functional as F

class SignBartMod(nn.Module):
    def __init__(self, num_classes, num_keypoints, num_face_landmarks, num_emotions,
                 d_model=256, nhead=16, num_layers=2,
                 dim_feedforward=1024, dropout=0.1):
        super().__init__()

        self.n_pose = 7
        self.n_pose_feat = 7
        self.n_hand = 21
        self.n_face = num_face_landmarks

        self.proj_pose_x  = nn.Linear(self.n_pose_feat, d_model)
        self.proj_lhand_x = nn.Linear(self.n_hand, d_model)
        self.proj_rhand_x = nn.Linear(self.n_hand, d_model)
        self.proj_face_x  = nn.Linear(self.n_face, d_model)

        self.proj_pose_y  = nn.Linear(self.n_pose_feat, d_model)
        self.proj_lhand_y = nn.Linear(self.n_hand, d_model)
        self.proj_rhand_y = nn.Linear(self.n_hand, d_model)
        self.proj_face_y  = nn.Linear(self.n_face, d_model)

        self.fusion_x = nn.Linear(4 * d_model, d_model)
        self.fusion_y = nn.Linear(4 * d_model, d_model)
        self.norm_x   = nn.LayerNorm(d_model)
        self.norm_y   = nn.LayerNorm(d_model)

        self.face_x_projector = nn.Linear(self.n_face, d_model)
        self.face_y_projector = nn.Linear(self.n_face, d_model)
        self.pos_encoder = nn.Parameter(torch.randn(1, 500, d_model))

        sign_enc_layer = nn.TransformerEncoderLayer(d_model, nhead, dim_feedforward, dropout, batch_first=True)
        self.encoder = nn.TransformerEncoder(sign_enc_layer, num_layers)

        sign_dec_layer = nn.TransformerDecoderLayer(d_model, nhead, dim_feedforward, dropout, batch_first=True)
        self.decoder = nn.TransformerDecoder(sign_dec_layer, num_layers)

        emo_enc_layer = nn.TransformerEncoderLayer(d_model, nhead, dim_feedforward, dropout, batch_first=True)
        self.emotion_encoder = nn.TransformerEncoder(emo_enc_layer, num_layers=1)

        emo_dec_layer = nn.TransformerDecoderLayer(d_model, nhead, dim_feedforward, dropout, batch_first=True)
        self.emotion_decoder = nn.TransformerDecoder(emo_dec_layer, num_layers=1)

        self.classifier = nn.Linear(d_model, num_classes)
        self.emotion_head = nn.Sequential(
            nn.Linear(d_model, d_model // 2),
            nn.ReLU(),
            nn.Dropout(dropout),
            nn.Linear(d_model // 2, num_emotions)
        )

    def forward(self, x, saliency_mask=None):
        B, T, K, _ = x.shape
        p_feat_end = self.n_pose_feat
        lh_end = p_feat_end + self.n_hand
        rh_end = lh_end + self.n_hand

        pose_x  = self.proj_pose_x (x[:, :, :p_feat_end, 0])
        lhand_x = self.proj_lhand_x(x[:, :, p_feat_end:lh_end, 0])
        rhand_x = self.proj_rhand_x(x[:, :, lh_end:rh_end, 0])
        face_x  = self.proj_face_x (x[:, :, rh_end:, 0])

        x_emb = self.norm_x(self.fusion_x(torch.cat([pose_x, lhand_x, rhand_x, face_x], dim=-1)))
        x_emb = x_emb + self.pos_encoder[:, :T, :]
        memory = self.encoder(x_emb)

        pose_y  = self.proj_pose_y (x[:, :, :p_feat_end, 1])
        lhand_y = self.proj_lhand_y(x[:, :, p_feat_end:lh_end, 1])
        rhand_y = self.proj_rhand_y(x[:, :, lh_end:rh_end, 1])
        face_y  = self.proj_face_y (x[:, :, rh_end:, 1])

        y_emb = self.norm_y(self.fusion_y(torch.cat([pose_y, lhand_y, rhand_y, face_y], dim=-1)))
        y_emb = y_emb + self.pos_encoder[:, :T, :]

        tgt_mask = torch.triu(torch.ones(T, T, device=x.device), diagonal=1).bool()
        output = self.decoder(y_emb, memory, tgt_mask=tgt_mask)
        logits = self.classifier(output.mean(dim=1))

        fx_emb = self.face_x_projector(x[:, :, rh_end:, 0]) + self.pos_encoder[:, :T, :]
        fy_emb = self.face_y_projector(x[:, :, rh_end:, 1]) + self.pos_encoder[:, :T, :]
        f_memory = self.emotion_encoder(fx_emb)
        f_output = self.emotion_decoder(fy_emb, f_memory, tgt_mask=tgt_mask)
        emotion_logits = self.emotion_head(f_output.mean(dim=1))

        saliency_logits = None
        if saliency_mask is not None and self.training:
            mask = saliency_mask.unsqueeze(-1)
            foreground_feat = (output * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1e-6)
            saliency_logits = self.classifier(foreground_feat)

        return logits, emotion_logits, saliency_logits