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