"""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"]