MSP / modeling_msp_fusion.py
MahmoodAnaam's picture
End of training
9c2f1dd verified
Raw
History Blame Contribute Delete
9.25 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.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,
)