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.utils import ModelOutput, logging from .configuration_msp_fusion import MSPFusionConfig logger = logging.get_logger(__name__) @dataclass class MSPFusionOutput(ModelOutput): last_hidden_state: torch.FloatTensor = None fusion_padding_mask: Optional[torch.Tensor] = None audio_hidden_state: Optional[torch.FloatTensor] = None visual_hidden_state: Optional[torch.FloatTensor] = None attentions: Optional[Tuple[torch.FloatTensor, ...]] = None class MSPFusionPreTrainedModel(PreTrainedModel): config_class = MSPFusionConfig base_model_prefix = "msp_fusion" 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=self.config.initializer_range) 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 _align_temporal( self, x: torch.Tensor, target_len: int, mode: str = "nearest" ) -> torch.Tensor: if x.size(1) == target_len: return x return F.interpolate(x.transpose(1, 2), size=target_len, mode=mode).transpose( 1, 2 ) class MSPFusionModel(MSPFusionPreTrainedModel): """Bidirectional cross-attention fusion of audio and visual streams.""" def __init__(self, config: MSPFusionConfig): super().__init__(config) if config.fusion_hidden_size % config.num_attention_heads != 0: raise ValueError( f"fusion_hidden_size ({config.fusion_hidden_size}) must be divisible " f"by num_attention_heads ({config.num_attention_heads})." ) # Modality projection blocks: linear + layer-norm self.audio_proj = nn.Sequential( nn.Linear(config.audio_hidden_size, config.fusion_hidden_size), nn.LayerNorm(config.fusion_hidden_size, eps=config.layer_norm_eps), ) self.visual_proj = nn.Sequential( nn.Linear(config.visual_hidden_size, config.fusion_hidden_size), nn.LayerNorm(config.fusion_hidden_size, eps=config.layer_norm_eps), ) # Bidirectional cross-attention self.audio_to_visual_attn = nn.MultiheadAttention( embed_dim=config.fusion_hidden_size, num_heads=config.num_attention_heads, dropout=config.attention_dropout, batch_first=True, ) self.visual_to_audio_attn = nn.MultiheadAttention( embed_dim=config.fusion_hidden_size, num_heads=config.num_attention_heads, dropout=config.attention_dropout, batch_first=True, ) # Post-attention residual norms self.audio_norm = nn.LayerNorm( config.fusion_hidden_size, eps=config.layer_norm_eps ) self.visual_norm = nn.LayerNorm( config.fusion_hidden_size, eps=config.layer_norm_eps ) # Gated fusion: concat → linear → gelu → dropout → layer-norm self.fusion_gate = nn.Sequential( nn.Linear(config.fusion_hidden_size * 2, config.fusion_hidden_size), nn.GELU(), nn.Dropout(config.dropout), nn.LayerNorm(config.fusion_hidden_size, eps=config.layer_norm_eps), ) @property def dummy_inputs(self) -> dict: return { "audio_hidden_states": torch.zeros(2, 50, self.config.audio_hidden_size), "visual_hidden_states": torch.zeros(2, 10, self.config.visual_hidden_size), } def _get_abs_attention_mask(self, attention_mask, dtype): if attention_mask.dim() == 2: extended_attention_mask = attention_mask[:, None, None, :] elif attention_mask.dim() == 3: extended_attention_mask = attention_mask[:, None, :, :] else: extended_attention_mask = attention_mask extended_attention_mask = extended_attention_mask.to(dtype=dtype) extended_attention_mask = (1.0 - extended_attention_mask) * torch.finfo( dtype ).min return extended_attention_mask def forward( self, audio_hidden_states: Optional[torch.Tensor] = None, visual_hidden_states: Optional[torch.Tensor] = None, audio_key_padding_mask: Optional[torch.Tensor] = None, visual_key_padding_mask: Optional[torch.Tensor] = None, output_attentions: Optional[bool] = None, ) -> MSPFusionOutput: """ Args: audio_hidden_states: (B, T_audio, audio_hidden_size) visual_hidden_states: (B, T_visual, visual_hidden_size) audio_key_padding_mask: (B, T_audio); True = pad position to ignore visual_key_padding_mask: (B, T_visual); True = pad position to ignore output_attentions: return attention weight matrices when True Returns: MSPFusionOutput with last_hidden_state of shape (B, T_audio, fusion_hidden_size) """ has_audio = audio_hidden_states is not None has_visual = visual_hidden_states is not None if not has_audio and not has_visual: raise ValueError( "At least one of audio_hidden_states or visual_hidden_states must be provided." ) output_attentions = ( output_attentions if output_attentions is not None else False ) # Handle cases where only one modality is present # only audio if has_audio and not has_visual: audio_states = self.audio_proj(audio_hidden_states) visual_states = torch.zeros_like(audio_states) fusion_key_padding_mask = audio_key_padding_mask fused = self.fusion_gate(torch.cat([audio_states, visual_states], dim=-1)) return MSPFusionOutput( last_hidden_state=fused, fusion_padding_mask=fusion_key_padding_mask, audio_hidden_state=audio_states, visual_hidden_state=None, attentions=None, ) # only visual if has_visual and not has_audio: visual_states = self.visual_proj(visual_hidden_states) audio_states = torch.zeros_like(visual_states) fusion_key_padding_mask = visual_key_padding_mask fused = self.fusion_gate(torch.cat([audio_states, visual_states], dim=-1)) return MSPFusionOutput( last_hidden_state=fused, fusion_padding_mask=fusion_key_padding_mask, audio_hidden_state=None, visual_hidden_state=visual_states, attentions=None, ) # Both modalities are present T_audio = audio_hidden_states.size(1) visual_hidden_states = self._align_temporal( visual_hidden_states, T_audio, mode="nearest" ) audio_states = self.audio_proj(audio_hidden_states) visual_states = self.visual_proj(visual_hidden_states) # fusion key padding mask for calc ctc loss if audio_key_padding_mask is not None: fusion_key_padding_mask = audio_key_padding_mask elif visual_key_padding_mask is not None: fusion_key_padding_mask = visual_key_padding_mask[:, : audio_states.size(1)] else: fusion_key_padding_mask = None a2v_attn = None v2a_attn = None audio_residual = audio_states visual_residual = visual_states # Bidirectional cross-attention # Audio attends to visual (audio queries, visual keys/values) audio_attended, a2v_attn = self.audio_to_visual_attn.forward( query=audio_states, key=visual_states, value=visual_states, need_weights=output_attentions, average_attn_weights=False, ) audio_states = self.audio_norm(audio_residual + audio_attended) # Visual attends to audio (visual queries, audio keys/values) visual_attended, v2a_attn = self.visual_to_audio_attn.forward( query=visual_states, key=audio_residual, value=audio_residual, need_weights=output_attentions, average_attn_weights=False, ) visual_states = self.visual_norm(visual_residual + visual_attended) # Gated fusion: concatenate both streams and project to fusion space fused = self.fusion_gate(torch.cat([audio_states, visual_states], dim=-1)) attentions = (a2v_attn, v2a_attn) if output_attentions else None return MSPFusionOutput( last_hidden_state=fused, fusion_padding_mask=fusion_key_padding_mask, audio_hidden_state=audio_states, visual_hidden_state=visual_states, attentions=attentions, )