image_through_model #49

Merged
brian merged 75 commits from image_through_model into main 2024-10-20 00:10:27 +02:00
Showing only changes of commit 0bfbb1920c - Show all commits
+20 -84
View File
@@ -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