Automatic Speech Recognition
Transformers
TensorBoard
Safetensors
msp
Generated from Trainer
custom_code
Instructions to use MahmoodAnaam/MSP-Fusion with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use MahmoodAnaam/MSP-Fusion with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("automatic-speech-recognition", model="MahmoodAnaam/MSP-Fusion", trust_remote_code=True)# Load model directly from transformers import AutoModelForCTC model = AutoModelForCTC.from_pretrained("MahmoodAnaam/MSP-Fusion", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
| 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__) | |
| 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), | |
| ) | |
| 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, | |
| ) | |