from __future__ import annotations import torch.nn as nn 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 from shared.dto import AngleData from shared.dto import ContactData from shared.dto import DistanceData from shared.dto import FramingData from shared.dto import InformationValueData from shared.dto import ModalityColorData from shared.dto import ModalityDepthData from shared.dto import ModalityLightingData from shared.dto import PointOfViewData from shared.dto import SalienceData from shared.dto import VisualSyntaxData class VisualCommunicationModel(nn.Module): 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): # 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