defined final model

This commit is contained in:
Brian Bjarke Jensen
2024-03-17 10:29:01 +01:00
parent 54eca39b6d
commit 50756450e4
+94 -5
View File
@@ -2,9 +2,29 @@ from __future__ import annotations
import torch.nn as nn 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 .resnet18_head import ResNet18Head
from .salience import SalienceTail
from .visual_syntax import VisualSyntaxTail 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): class VisualCommunicationModel(nn.Module):
@@ -13,6 +33,16 @@ class VisualCommunicationModel(nn.Module):
# store other models # store other models
self.resnet_head = ResNet18Head() self.resnet_head = ResNet18Head()
self.visual_syntax_tail = VisualSyntaxTail() 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): def forward(self, x):
# generate visual representation # generate visual representation
@@ -20,10 +50,69 @@ class VisualCommunicationModel(nn.Module):
# prepare result map # prepare result map
results = {} results = {}
# predict visual syntax # predict visual syntax
visual_syntax_pred = self.visual_syntax_tail(vis_rep).cpu()
# visual_syntax_pred = float()
results['visual_syntax'] = VisualSyntaxData.from_list( 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 return results