From 50756450e400fae688b990d9ebccfa2e873bf51f Mon Sep 17 00:00:00 2001 From: Brian Bjarke Jensen Date: Sun, 17 Mar 2024 10:29:01 +0100 Subject: [PATCH] defined final model --- model/src/models/visual_communication.py | 99 ++++++++++++++++++++++-- 1 file changed, 94 insertions(+), 5 deletions(-) diff --git a/model/src/models/visual_communication.py b/model/src/models/visual_communication.py index e11a81f..8d666f5 100644 --- a/model/src/models/visual_communication.py +++ b/model/src/models/visual_communication.py @@ -2,9 +2,29 @@ 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 core.data_models import VisualSyntaxData +from core.dto import AngleData +from core.dto import ContactData +from core.dto import DistanceData +from core.dto import FramingData +from core.dto import InformationValueData +from core.dto import ModalityColorData +from core.dto import ModalityDepthData +from core.dto import ModalityLightingData +from core.dto import PointOfViewData +from core.dto import SalienceData +from core.dto import VisualSyntaxData class VisualCommunicationModel(nn.Module): @@ -13,6 +33,16 @@ class VisualCommunicationModel(nn.Module): # 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 @@ -20,10 +50,69 @@ class VisualCommunicationModel(nn.Module): # prepare result map results = {} # predict visual syntax - visual_syntax_pred = self.visual_syntax_tail(vis_rep).cpu() - # visual_syntax_pred = float() results['visual_syntax'] = VisualSyntaxData.from_list( - visual_syntax_pred, + 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