"""Definition of VisualCommunicationModel class.""" from __future__ import annotations from torch import nn from shared.mongodb.classes import ModelData from .angle import AngleTail from .contact import ContactTail from .distance import DistanceTail from .framing import FramingTail from .information_value import InformationValueTail from .modality_color import ModalityColorTail from .modality_depth import ModalityDepthTail from .modality_lighting import ModalityLightingTail from .point_of_view import PointOfViewTail from .resnet18_head import ResNet18Head from .salience import SalienceTail from .visual_syntax import VisualSyntaxTail class VisualCommunicationModel(nn.Module): """Visual communication model.""" def __init__(self, download_resnet_weights: bool = False): super().__init__() # store other models self.resnet_head = ResNet18Head(download_resnet_weights) self.visual_syntax_tail = VisualSyntaxTail() self.contact_tail = ContactTail() self.angle_tail = AngleTail() self.point_of_view_tail = PointOfViewTail() self.distance_tail = DistanceTail() self.modality_lighting_tail = ModalityLightingTail() self.modality_color_tail = ModalityColorTail() self.modality_depth_tail = ModalityDepthTail() self.information_value_tail = InformationValueTail() self.framing_tail = FramingTail() self.salience_tail = SalienceTail() def forward(self, x) -> ModelData: """Calculate model output on data.""" # generate visual representation features = self.resnet_head(x) # make predictions prediction_dict = { 'visual_syntax': self.visual_syntax_tail(features), 'contact': self.contact_tail(features), 'angle': self.angle_tail(features), 'point_of_view': self.point_of_view_tail(features), 'distance': self.distance_tail(features), 'modality_lighting': self.modality_lighting_tail(features), 'modality_color': self.modality_color_tail(features), 'modality_depth': self.modality_depth_tail(features), 'information_value': self.information_value_tail(features), 'framing': self.framing_tail(features), 'salience': self.salience_tail(features), } # convert to respective classes data = ModelData.from_prediction_dict(prediction_dict) return data