| 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 |