added method to instantiate from prediction dict

This commit is contained in:
Brian Bjarke Jensen
2024-07-28 23:19:51 +02:00
parent cead3d24c8
commit b415137e25
+57 -24
View File
@@ -1,4 +1,5 @@
"""Definition of ModelData data model.""" """Definition of ModelData data model."""
from __future__ import annotations from __future__ import annotations
from .angle import AngleData from .angle import AngleData
@@ -17,6 +18,7 @@ from .visual_syntax import VisualSyntaxData
class ModelData(DataModel): class ModelData(DataModel):
"""ModelData model for data IO with combined ML model.""" """ModelData model for data IO with combined ML model."""
visual_syntax: VisualSyntaxData visual_syntax: VisualSyntaxData
contact: ContactData contact: ContactData
angle: AngleData angle: AngleData
@@ -34,8 +36,50 @@ class ModelData(DataModel):
"""Instantiate with random numbers.""" """Instantiate with random numbers."""
kwargs = { kwargs = {
field: field_info.annotation.from_random() # type: ignore field: field_info.annotation.from_random() # type: ignore
for field, field_info for field, field_info in cls.model_fields.items()
in cls.model_fields.items() }
return cls(**kwargs)
@classmethod
def from_prediction_dict(
cls,
prediction_dict: dict,
) -> ModelData:
"""Instantiate from prediction dictionary."""
kwargs = {
'visual_syntax': VisualSyntaxData.from_tensor(
prediction_dict['visual_syntax'],
),
'contact': ContactData.from_tensor(
prediction_dict['contact'],
),
'angle': AngleData.from_tensor(
prediction_dict['angle'],
),
'point_of_view': PointOfViewData.from_tensor(
prediction_dict['point_of_view'],
),
'distance': DistanceData.from_tensor(
prediction_dict['distance'],
),
'modality_lighting': ModalityLightingData.from_tensor(
prediction_dict['modality_lighting'],
),
'modality_color': ModalityColorData.from_tensor(
prediction_dict['modality_color'],
),
'modality_depth': ModalityDepthData.from_tensor(
prediction_dict['modality_depth'],
),
'information_value': InformationValueData.from_tensor(
prediction_dict['information_value'],
),
'framing': FramingData.from_tensor(
prediction_dict['framing'],
),
'salience': SalienceData.from_tensor(
prediction_dict['salience'],
),
} }
return cls(**kwargs) return cls(**kwargs)
@@ -56,27 +100,16 @@ class ModelData(DataModel):
) -> ModelData: ) -> ModelData:
"""Instantiate from annotation.""" """Instantiate from annotation."""
kwargs = { kwargs = {
'visual_syntax': VisualSyntaxData 'visual_syntax': VisualSyntaxData.from_choice(visual_syntax),
.from_choice(visual_syntax), 'contact': ContactData.from_choice(contact),
'contact': ContactData 'angle': AngleData.from_choice(angle),
.from_choice(contact), 'point_of_view': PointOfViewData.from_choice(point_of_view),
'angle': AngleData 'distance': DistanceData.from_choice(distance),
.from_choice(angle), 'modality_lighting': ModalityLightingData.from_choice(modality_lighting),
'point_of_view': PointOfViewData 'modality_color': ModalityColorData.from_choice(modality_color),
.from_choice(point_of_view), 'modality_depth': ModalityDepthData.from_choice(modality_depth),
'distance': DistanceData 'information_value': InformationValueData.from_choice(information_value),
.from_choice(distance), 'framing': FramingData.from_choice(framing),
'modality_lighting': ModalityLightingData 'salience': SalienceData.from_choice(salience),
.from_choice(modality_lighting),
'modality_color': ModalityColorData
.from_choice(modality_color),
'modality_depth': ModalityDepthData
.from_choice(modality_depth),
'information_value': InformationValueData
.from_choice(information_value),
'framing': FramingData
.from_choice(framing),
'salience': SalienceData
.from_choice(salience),
} }
return cls(**kwargs) return cls(**kwargs)