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 |