From cead3d24c8ba7d115b29f65c68cc7d4fe9ab45c6 Mon Sep 17 00:00:00 2001 From: Brian Bjarke Jensen Date: Sun, 28 Jul 2024 23:13:15 +0200 Subject: [PATCH] added from_tensor method --- shared/dto/data_model.py | 26 +++++++++++++------------- 1 file changed, 13 insertions(+), 13 deletions(-) diff --git a/shared/dto/data_model.py b/shared/dto/data_model.py index 70bb79f..f19b0c8 100644 --- a/shared/dto/data_model.py +++ b/shared/dto/data_model.py @@ -1,10 +1,9 @@ """Definition of DataModel base class.""" -from __future__ import annotations import random -from pydantic import BaseModel -from pydantic import ValidationError +from pydantic import BaseModel, ValidationError +from torch import Tensor class DataModel(BaseModel): @@ -33,26 +32,27 @@ class DataModel(BaseModel): raise ValidationError() assert isinstance(option, str), 'option is not a string' allowed_options_list = cls.list_fields() - assert option in allowed_options_list, \ - f"{option} is not among allowed fields {allowed_options_list}" + assert ( + option in allowed_options_list + ), f"{option} is not among allowed fields {allowed_options_list}" kwargs = {field: 0 for field in cls.list_fields()} kwargs[option] = 1 return cls(**kwargs) @classmethod - def from_list(cls, data_list: list[float]): + def from_tensor(cls, tensor: Tensor): """Instantiate from list of values.""" - kwargs = {key: val for key, val in zip(cls.list_fields(), data_list)} + assert tensor.size(dim=0) == 1, f'tensor batch larger than 1: {tensor}' + data_list = [float(t.item()) for t in tensor[0]] + kwargs = dict(zip(cls.list_fields(), data_list)) return cls(**kwargs) def __repr__(self) -> str: model_dict = self.model_dump() - model_repr_str = f"{self.classname()}(" - model_repr_str += ', '.join([ - f"{field}={value:.3f}" - for field, value - in model_dict.items() - ]) + model_repr_str = f'{self.classname()}(' + model_repr_str += ', '.join( + [f'{field}={value:.3f}' for field, value in model_dict.items()], + ) model_repr_str += ')' return model_repr_str