"""Definition of VisualCommunicationModel class.""" from __future__ import annotations from torch import nn from shared.dto import ( AngleData, ContactData, DistanceData, FramingData, InformationValueData, ModalityColorData, ModalityDepthData, ModalityLightingData, PointOfViewData, SalienceData, VisualSyntaxData, ) 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): super().__init__() # store other models self.resnet_head = ResNet18Head() 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): """Calculate model output on data.""" # generate visual representation vis_rep = self.resnet_head(x) # prepare result map results = {} # predict visual syntax results['visual_syntax'] = VisualSyntaxData.from_list( self.visual_syntax_tail( vis_rep, ).cpu(), ) # predict contact results['contact'] = ContactData.from_list( self.contact_tail( vis_rep, ).cpu(), ) # predict angle results['angle'] = AngleData.from_list( self.angle_tail( vis_rep, ).cpu(), ) # predict point of view results['point_of_view'] = PointOfViewData.from_list( self.point_of_view_tail( vis_rep, ).cpu(), ) # predict distance results['distance'] = DistanceData.from_list( self.distance_tail( vis_rep, ).cpu(), ) # predict modality lighting results['modality_lighting'] = ModalityLightingData.from_list( self.modality_lighting_tail( vis_rep, ).cpu(), ) # predict modality color results['modality_color'] = ModalityColorData.from_list( self.modality_color_tail( vis_rep, ).cpu(), ) # predict modality depth results['modality_depth'] = ModalityDepthData.from_list( self.modality_depth_tail( vis_rep, ).cpu(), ) # predict information value results['information_value'] = InformationValueData.from_list( self.information_value_tail( vis_rep, ).cpu(), ) # predict framing results['framing'] = FramingData.from_list( self.framing_tail( vis_rep, ).cpu(), ) # predict salience results['salience'] = SalienceData.from_list( self.salience_tail( vis_rep, ).cpu(), ) return results