diff --git a/shared/dto/model_data.py b/shared/dto/model_data.py index 157f788..5bbbf07 100644 --- a/shared/dto/model_data.py +++ b/shared/dto/model_data.py @@ -1,4 +1,5 @@ """Definition of ModelData data model.""" + from __future__ import annotations from .angle import AngleData @@ -17,6 +18,7 @@ from .visual_syntax import VisualSyntaxData class ModelData(DataModel): """ModelData model for data IO with combined ML model.""" + visual_syntax: VisualSyntaxData contact: ContactData angle: AngleData @@ -34,8 +36,50 @@ class ModelData(DataModel): """Instantiate with random numbers.""" kwargs = { field: field_info.annotation.from_random() # type: ignore - for field, field_info - in cls.model_fields.items() + for field, field_info 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) @@ -56,27 +100,16 @@ class ModelData(DataModel): ) -> ModelData: """Instantiate from annotation.""" kwargs = { - 'visual_syntax': VisualSyntaxData - .from_choice(visual_syntax), - 'contact': ContactData - .from_choice(contact), - 'angle': AngleData - .from_choice(angle), - 'point_of_view': PointOfViewData - .from_choice(point_of_view), - 'distance': DistanceData - .from_choice(distance), - 'modality_lighting': ModalityLightingData - .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), + 'visual_syntax': VisualSyntaxData.from_choice(visual_syntax), + 'contact': ContactData.from_choice(contact), + 'angle': AngleData.from_choice(angle), + 'point_of_view': PointOfViewData.from_choice(point_of_view), + 'distance': DistanceData.from_choice(distance), + 'modality_lighting': ModalityLightingData.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)