fixed conversion of classes
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user