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