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.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__) | |
| 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) | |
| 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) | |
| ) | |
| 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 | |
| ) | |