diff --git a/model/src/models/visual_communication.py b/model/src/models/visual_communication.py index a0b55e1..c7403bb 100644 --- a/model/src/models/visual_communication.py +++ b/model/src/models/visual_communication.py @@ -4,19 +4,7 @@ from __future__ import annotations from torch import nn -from shared.dto import ( - AngleData, - ContactData, - DistanceData, - FramingData, - InformationValueData, - ModalityColorData, - ModalityDepthData, - ModalityLightingData, - PointOfViewData, - SalienceData, - VisualSyntaxData, -) +from shared.dto import ModelData from .angle import AngleTail from .contact import ContactTail @@ -51,76 +39,24 @@ class VisualCommunicationModel(nn.Module): self.framing_tail = FramingTail() self.salience_tail = SalienceTail() - def forward(self, x): + def forward(self, x) -> ModelData: """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 + 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