fixed conversion of classes

This commit is contained in:
Brian Bjarke Jensen
2024-07-28 23:12:06 +02:00
parent 164102557e
commit 0bfbb1920c
+20 -84
View File
@@ -4,19 +4,7 @@ from __future__ import annotations
from torch import nn from torch import nn
from shared.dto import ( from shared.dto import ModelData
AngleData,
ContactData,
DistanceData,
FramingData,
InformationValueData,
ModalityColorData,
ModalityDepthData,
ModalityLightingData,
PointOfViewData,
SalienceData,
VisualSyntaxData,
)
from .angle import AngleTail from .angle import AngleTail
from .contact import ContactTail from .contact import ContactTail
@@ -51,76 +39,24 @@ class VisualCommunicationModel(nn.Module):
self.framing_tail = FramingTail() self.framing_tail = FramingTail()
self.salience_tail = SalienceTail() self.salience_tail = SalienceTail()
def forward(self, x): def forward(self, x) -> ModelData:
"""Calculate model output on data.""" """Calculate model output on data."""
# generate visual representation # generate visual representation
vis_rep = self.resnet_head(x) features = self.resnet_head(x)
# prepare result map # make predictions
results = {} prediction_dict = {
# predict visual syntax 'visual_syntax': self.visual_syntax_tail(features),
results['visual_syntax'] = VisualSyntaxData.from_list( 'contact': self.contact_tail(features),
self.visual_syntax_tail( 'angle': self.angle_tail(features),
vis_rep, 'point_of_view': self.point_of_view_tail(features),
).cpu(), 'distance': self.distance_tail(features),
) 'modality_lighting': self.modality_lighting_tail(features),
# predict contact 'modality_color': self.modality_color_tail(features),
results['contact'] = ContactData.from_list( 'modality_depth': self.modality_depth_tail(features),
self.contact_tail( 'information_value': self.information_value_tail(features),
vis_rep, 'framing': self.framing_tail(features),
).cpu(), 'salience': self.salience_tail(features),
) }
# predict angle # convert to respective classes
results['angle'] = AngleData.from_list( data = ModelData.from_prediction_dict(prediction_dict)
self.angle_tail( return data
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