MSP-Fusion / modeling_msp.py
MahmoodAnaam's picture
Update modeling_msp.py
3f14f06 verified
Raw
History Blame Contribute Delete
15.3 kB
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
)