BiGRU_T_version / src /bigru_t /multimodal /video_encoder.py
PowerMachine's picture
Upload folder using huggingface_hub
3275441 verified
Raw History Blame Contribute Delete
1.36 kB
"""Xavante - video_encoder.py - Encoder de video (fusao imagem + audio)."""
from __future__ import annotations
import logging
import torch
import torch.nn as nn
from .image_encoder import ImageEncoder
from .audio_encoder import AudioEncoder
logger = logging.getLogger(__name__)
class VideoEncoder(nn.Module):
"""Combina frames de video + audio em uma representacao unificada."""
def __init__(self, d_model: int = 512, n_frames: int = 8):
super().__init__()
self.n_frames = n_frames
self.frame_encoder = ImageEncoder(d_model)
self.audio_encoder = AudioEncoder(d_model)
# Temporal aggregation
self.temporal = nn.GRU(d_model, d_model, batch_first=True)
self.fuse = nn.Linear(d_model * 2, d_model)
def forward(
self,
frames: torch.Tensor, # [B, T, C, H, W]
audio: torch.Tensor, # [B, C_audio, T_audio]
) -> torch.Tensor:
B, T, C, H, W = frames.shape
frames_flat = frames.view(B * T, C, H, W)
frame_emb = self.frame_encoder(frames_flat) # [B*T, d]
frame_emb = frame_emb.view(B, T, -1)
_, h = self.temporal(frame_emb)
video_emb = h.squeeze(0) # [B, d]
audio_emb = self.audio_encoder(audio)
return torch.tanh(self.fuse(torch.cat([video_emb, audio_emb], dim=-1)))
__all__ = ["VideoEncoder"]