from collections import OrderedDict from dataclasses import dataclass from typing import Optional, Tuple import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel from transformers.modeling_outputs import CausalLMOutput from transformers.utils import ModelOutput, logging from .modeling_msp_audio import MSPAudioModel from .modeling_msp_visual import MSPVisualModel from .configuration_msp import MSPConfig from .modeling_msp_fusion import MSPFusionModel logger = logging.get_logger(__name__) @dataclass class MSPOutput(ModelOutput): loss: Optional[torch.FloatTensor] = None logits: Optional[torch.FloatTensor] = None audio_logits: Optional[torch.FloatTensor] = None visual_logits: Optional[torch.FloatTensor] = None audio_loss: Optional[torch.FloatTensor] = None visual_loss: Optional[torch.FloatTensor] = None last_hidden_state: Optional[torch.FloatTensor] = None audio_hidden_state: Optional[torch.FloatTensor] = None visual_hidden_state: Optional[torch.FloatTensor] = None fusion_padding_mask: Optional[torch.Tensor] = None fusion_input_lengths: Optional[torch.Tensor] = None audio_input_lengths: Optional[torch.Tensor] = None visual_input_lengths: Optional[torch.Tensor] = None attentions: Optional[Tuple[torch.FloatTensor, ...]] = None class MSPPreTrainedModel(PreTrainedModel): config_class = MSPConfig base_model_prefix = "msp" main_input_name = "input_values" input_modalities = ["audio", "video"] supports_gradient_checkpointing = False all_tied_weights_keys = OrderedDict() def _init_weights(self, module): if isinstance(module, nn.Linear): module.weight.data.normal_(mean=0.0, std=0.02) if module.bias is not None: module.bias.data.zero_() elif isinstance(module, nn.LayerNorm): module.bias.data.zero_() module.weight.data.fill_(1.0) def _apply_modality_dropout( self, has_audio: bool, has_visual: bool, ) -> tuple[bool, bool]: if not self.training: return has_audio, has_visual if not has_audio or not has_visual: return has_audio, has_visual if self.config.modality_dropout_prob <= 0.0: return has_audio, has_visual if torch.rand(()) >= self.config.modality_dropout_prob: return has_audio, has_visual audio_drop_prob = self.config.audio_dropout_prob visual_drop_prob = self.config.visual_dropout_prob total = audio_drop_prob + visual_drop_prob if total <= 0: return has_audio, has_visual drop_audio = torch.rand(()) < (audio_drop_prob / total) if drop_audio: return False, True return True, False def _ctc_loss( logits: torch.Tensor, labels: torch.Tensor, input_lengths: torch.Tensor, pad_token_id: int, reduction: str, zero_infinity: bool, ) -> torch.Tensor: """Compute CTC loss from logits, labels, and pre-computed input lengths.""" labels_mask = labels >= 0 target_lengths = labels_mask.sum(-1) flattened_targets = labels.masked_select(labels_mask) log_probs = F.log_softmax(logits, dim=-1, dtype=torch.float32).transpose(0, 1) with torch.backends.cudnn.flags(enabled=False): return F.ctc_loss( log_probs, flattened_targets, input_lengths, target_lengths, blank=pad_token_id, reduction=reduction, zero_infinity=zero_infinity, ) class MSPModel(MSPPreTrainedModel): def __init__(self, config: MSPConfig): super().__init__(config) # Audio encoder and auxiliary CTC head self.audio_model = MSPAudioModel(config.audio_config) self.audio_head = nn.Sequential( nn.Dropout(config.audio_config.final_dropout), nn.Linear(config.audio_config.hidden_size, config.audio_config.vocab_size), ) # Visual encoder and auxiliary CTC head self.visual_model = MSPVisualModel(config.visual_config) self.visual_head = nn.Sequential( nn.Dropout(config.visual_config.final_dropout), nn.Linear( config.visual_config.hidden_size, config.visual_config.vocab_size, ), ) # Bidirectional cross-attention fusion self.fusion_model = MSPFusionModel(config.msp_fusion_config) @property def dummy_inputs(self) -> dict: return { "input_values": torch.zeros(1, 16000, dtype=torch.float32), "pixel_values_videos": torch.zeros(1, 1, 10, 88, 88, dtype=torch.float32), "padding_mask": torch.ones(1, 16000, dtype=torch.long), "padding_mask_videos": torch.ones(1, 10, dtype=torch.long), } def forward( self, input_values: Optional[torch.Tensor] = None, pixel_values_videos: Optional[torch.Tensor] = None, padding_mask: Optional[torch.Tensor] = None, padding_mask_videos: Optional[torch.Tensor] = None, output_attentions: Optional[bool] = None, output_hidden_states: Optional[bool] = None, **kwargs, ) -> MSPOutput: has_audio = input_values is not None has_visual = pixel_values_videos is not None if not has_audio and not has_visual: raise ValueError( "Either input_values or pixel_values_videos must be provided." ) output_attentions = ( output_attentions if output_attentions is not None else False ) use_audio, use_visual = self._apply_modality_dropout(has_audio, has_visual) audio_output, visual_output, fusion_output = None, None, None audio_hidden_states, visual_hidden_states = None, None audio_logits, visual_logits = None, None audio_input_lengths, visual_input_lengths, fusion_input_lengths = ( None, None, None, ) if use_audio: audio_output = self.audio_model( input_values=input_values, padding_mask=padding_mask, output_attentions=output_attentions, output_hidden_states=output_hidden_states, ) padding_mask = ( padding_mask if padding_mask is not None else torch.ones_like( input_values, dtype=torch.long, device=input_values.device ) ) audio_input_lengths = self.audio_model._get_feat_extract_output_lengths( padding_mask.sum(-1) ).to(torch.long) audio_hidden_states = audio_output.last_hidden_state audio_logits = self.audio_head(audio_hidden_states) if use_visual: visual_output = self.visual_model( pixel_values_videos=pixel_values_videos, padding_mask_videos=padding_mask_videos, output_attentions=output_attentions, output_hidden_states=output_hidden_states, ) padding_mask_videos = ( visual_output.padding_mask_videos if padding_mask_videos is not None else torch.ones( (pixel_values_videos.shape[0], pixel_values_videos.shape[2]), dtype=torch.long, device=pixel_values_videos.device, ) ) visual_input_lengths = ( padding_mask_videos.sum(-1) .to(torch.long) .to(pixel_values_videos.device) ) visual_hidden_states = visual_output.last_hidden_state visual_logits = self.visual_head(visual_hidden_states) fusion_output = self.fusion_model.forward( audio_hidden_states=audio_hidden_states, visual_hidden_states=visual_hidden_states, audio_key_padding_mask=padding_mask, visual_key_padding_mask=padding_mask_videos, output_attentions=output_attentions, ) fusion_input_lengths = ( fusion_output.fusion_padding_mask.sum(-1) .to(torch.long) .to(fusion_output.fusion_padding_mask.device) if fusion_output.fusion_padding_mask is not None and not use_audio else audio_input_lengths ) return MSPOutput( last_hidden_state=fusion_output.last_hidden_state, audio_hidden_state=fusion_output.audio_hidden_state, visual_hidden_state=fusion_output.visual_hidden_state, fusion_padding_mask=fusion_output.fusion_padding_mask, audio_logits=audio_logits, visual_logits=visual_logits, audio_input_lengths=audio_input_lengths, visual_input_lengths=visual_input_lengths, fusion_input_lengths=fusion_input_lengths, attentions=fusion_output.attentions, ) class MSPForCTC(MSPPreTrainedModel): def __init__(self, config: MSPConfig): super().__init__(config) if config.vocab_size is None: raise ValueError( "vocab_size must be set in MSPConfig to instantiate MSPForCTC." ) self.msp = MSPModel(config) # Final CTC head for the fused representation self.msp_head = nn.Sequential( nn.Dropout(config.final_dropout), nn.Linear(config.msp_fusion_config.fusion_hidden_size, config.vocab_size) ) @property def dummy_inputs(self) -> dict: return { "input_values": torch.zeros(1, 16000, dtype=torch.float32), "pixel_values_videos": torch.zeros(1, 1, 10, 88, 88, dtype=torch.float32), "padding_mask": torch.ones(1, 16000, dtype=torch.long), "padding_mask_videos": torch.ones(1, 10, dtype=torch.long), "labels": torch.ones(1, 5, dtype=torch.long), } # --- Freeze helpers --- def freeze_feature_encoder(self) -> None: """Freeze feature extractors of both encoders (for end-to-end fine-tuning).""" self.msp.audio_model.feature_extractor._freeze_parameters() for param in self.msp.visual_model.feature_extractor_video.parameters(): param.requires_grad = False for param in self.msp.visual_model.feature_extractor_audio.parameters(): param.requires_grad = False def freeze_base_model(self) -> None: """Freeze both encoders (for fusion-only training).""" for param in self.msp.audio_model.parameters(): param.requires_grad = False for param in self.msp.visual_model.parameters(): param.requires_grad = False def freeze_audio_branch(self) -> None: """Freeze audio encoder and its CTC head (for fusion-only training).""" for param in self.msp.audio_model.parameters(): param.requires_grad = False for param in self.msp.audio_head.parameters(): param.requires_grad = False def freeze_visual_branch(self) -> None: """Freeze visual encoder and its CTC head (for fusion-only training).""" for param in self.msp.visual_model.parameters(): param.requires_grad = False for param in self.msp.visual_head.parameters(): param.requires_grad = False def forward( self, input_values: Optional[torch.Tensor] = None, pixel_values_videos: Optional[torch.Tensor] = None, padding_mask: Optional[torch.Tensor] = None, padding_mask_videos: Optional[torch.Tensor] = None, labels: Optional[torch.Tensor] = None, output_attentions: Optional[bool] = None, output_hidden_states: Optional[bool] = None, **kwargs, ) -> CausalLMOutput: if input_values is None and pixel_values_videos is None: raise ValueError( "Either input_values or pixel_values_videos must be provided." ) msp_out = self.msp( input_values=input_values, pixel_values_videos=pixel_values_videos, padding_mask=padding_mask, padding_mask_videos=padding_mask_videos, output_attentions=output_attentions, output_hidden_states=output_hidden_states, ) # Final CTC logits from the fused representation logits = self.msp_head(msp_out.last_hidden_state) loss = None if labels is not None: valid_labels = labels[labels >= 0] if ( valid_labels.numel() > 0 and valid_labels.max() >= self.config.vocab_size ): raise ValueError( f"Label value {valid_labels.max()} >= vocab_size={self.config.vocab_size}." ) # Audio CTC loss ctc_audio = None if ( msp_out.audio_input_lengths is not None and self.config.ctc_loss_audio_weight != 0.0 ): audio_lengths = msp_out.audio_input_lengths ctc_audio = _ctc_loss( logits=msp_out.audio_logits, labels=labels, input_lengths=audio_lengths, pad_token_id=self.config.pad_token_id, reduction=self.config.ctc_loss_reduction, zero_infinity=self.config.ctc_zero_infinity, ) # Visual CTC loss ctc_visual = None if ( msp_out.visual_input_lengths is not None and self.config.ctc_loss_visual_weight != 0.0 ): visual_lengths = msp_out.visual_input_lengths ctc_visual = _ctc_loss( logits=msp_out.visual_logits, labels=labels, input_lengths=visual_lengths, pad_token_id=self.config.pad_token_id, reduction=self.config.ctc_loss_reduction, zero_infinity=self.config.ctc_zero_infinity, ) # Fusion CTC loss ctc_msp = None if ( msp_out.fusion_input_lengths is not None and self.config.ctc_loss_msp_weight != 0.0 ): msp_lengths = msp_out.fusion_input_lengths ctc_msp = _ctc_loss( logits=logits, labels=labels, input_lengths=msp_lengths, pad_token_id=self.config.pad_token_id, reduction=self.config.ctc_loss_reduction, zero_infinity=self.config.ctc_zero_infinity, ) # Weighted combination loss = self.config.ctc_loss_msp_weight * ctc_msp if ctc_audio is not None: loss = loss + self.config.ctc_loss_audio_weight * ctc_audio if ctc_visual is not None: loss = loss + self.config.ctc_loss_visual_weight * ctc_visual return CausalLMOutput( loss=loss, logits=logits )