MSP-Fusion-V0 / modeling_msp_visual.py
MahmoodAnaam's picture
Upload 11 files
02471cf verified
Raw
History Blame Contribute Delete
5.11 kB
import warnings
from collections import OrderedDict
from dataclasses import dataclass
from typing import Optional
import torch
from torch import nn
from transformers import PreTrainedModel
from transformers.modeling_outputs import (
BaseModelOutputWithPooling,
CausalLMOutput,
)
from .configuration_msp_visual import MSPVisualConfig
from .modeling_avhubert import AVHubertModel
warnings.filterwarnings("ignore")
class CTCDecoder(nn.Module):
def __init__(
self,
config: MSPVisualConfig,
):
super().__init__()
self.lm_dropout = torch.nn.Dropout(config.lm_dropout)
self.lm_head = nn.Linear(config.visual_config.hidden_size, config.vocab_size)
self.config = config
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.lm_dropout(hidden_states)
logits = self.lm_head(hidden_states)
return logits
def compute_loss(
self, logits: torch.Tensor, labels: torch.Tensor, input_lengths: torch.Tensor
) -> torch.Tensor:
labels_mask = labels != -100
target_lengths = labels_mask.sum(-1).to(torch.long)
flattened_targets = labels.masked_select(labels_mask)
log_probs = nn.functional.log_softmax(
logits, dim=-1, dtype=torch.float32
).transpose(0, 1)
with torch.backends.cudnn.flags(enabled=False):
loss = nn.functional.ctc_loss(
log_probs=log_probs,
targets=flattened_targets,
input_lengths=input_lengths,
target_lengths=target_lengths,
blank=self.config.pad_token_id,
reduction=self.config.ctc_reduction,
zero_infinity=self.config.ctc_zero_infinity,
)
return loss
class MSPVisualPreTrainedModel(PreTrainedModel):
config_class = MSPVisualConfig
base_model_prefix = "visual_model"
supports_gradient_checkpointing = True
all_tied_weights_keys = OrderedDict()
main_input_name = "pixel_values_videos"
@dataclass
class MSPVisualEncoderOutput(BaseModelOutputWithPooling):
padding_mask: Optional[torch.Tensor] = None
@dataclass
class MSPVisualOutput(CausalLMOutput):
encoder_outputs: Optional[MSPVisualEncoderOutput] = None
class MSPVisualEncoder(MSPVisualPreTrainedModel):
def __init__(self, config: MSPVisualConfig):
super().__init__(config)
self.encoder = AVHubertModel(config.visual_config)
self.encoder.feature_extractor_audio.requires_grad_(False)
def forward(
self,
pixel_values_videos: torch.Tensor,
padding_mask_videos: Optional[torch.Tensor] = None,
**kwargs,
) -> MSPVisualEncoderOutput:
feature, padding_mask = self.encoder.extract_finetune(
source={
"video": pixel_values_videos, # shape [batch_size, num_channels=1, num_frames, height, width]
"audio": None,
},
padding_mask=padding_mask_videos, # shape [batch_size, num_frames]
)
return MSPVisualEncoderOutput(
last_hidden_state=feature, # shape [batch_size, num_frames, hidden_size]
pooler_output=feature.mean(dim=1), # shape [batch_size, hidden_size]
padding_mask=padding_mask, # shape [batch_size, num_frames]
)
class MSPVisualModel(MSPVisualPreTrainedModel):
def __init__(self, config: MSPVisualConfig):
super().__init__(config)
self.visual_encoder = MSPVisualEncoder(config)
self.ctc_decoder = CTCDecoder(config)
def freeze_feature_extractor(self):
self.visual_encoder.encoder.feature_extractor_video.requires_grad_(False)
def freeze_visual_encoder(self):
self.visual_encoder.requires_grad_(False)
def forward(
self,
pixel_values_videos: torch.Tensor,
padding_mask_videos: Optional[torch.Tensor] = None,
labels: Optional[torch.Tensor] = None,
**kwargs,
) -> MSPVisualOutput:
encoder_outputs = self.visual_encoder(
pixel_values_videos=pixel_values_videos,
padding_mask_videos=padding_mask_videos,
)
logits = self.ctc_decoder(encoder_outputs.last_hidden_state)
loss = None
if labels is not None:
if padding_mask_videos is not None:
input_lengths = encoder_outputs.padding_mask.sum(-1).to(torch.long)
else:
input_lengths = torch.full(
(logits.size(0),),
logits.size(1),
dtype=torch.long,
device=logits.device,
)
loss = self.ctc_decoder.compute_loss(
logits=logits, labels=labels, input_lengths=input_lengths
)
return MSPVisualOutput(
loss=loss,
logits=logits,
encoder_outputs=encoder_outputs,
)