add_model_resnet18 #31

Merged
brian merged 47 commits from add_model_resnet18 into main 2024-04-03 21:04:23 +02:00
Showing only changes of commit 50756450e4 - Show all commits
+94 -5
View File
@@ -2,9 +2,29 @@ from __future__ import annotations
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 .salience import SalienceTail
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):
@@ -13,6 +33,16 @@ class VisualCommunicationModel(nn.Module):
# store other models
self.resnet_head = ResNet18Head()
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):
# generate visual representation
@@ -20,10 +50,69 @@ class VisualCommunicationModel(nn.Module):
# prepare result map
results = {}
# predict visual syntax
visual_syntax_pred = self.visual_syntax_tail(vis_rep).cpu()
# visual_syntax_pred = float()
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