"""Definition of DataModel base class.""" import random from pydantic import BaseModel, ValidationError from torch import Tensor class DataModel(BaseModel): """DataModel base class.""" @classmethod def classname(cls) -> str: """Return classname.""" return cls.__name__ @classmethod def list_fields(cls) -> list[str]: """List options that are stored as attributes.""" return list(cls.model_fields.keys()) @classmethod def from_random(cls): """Instantiate with random numbers.""" kwargs = {field: random.random() for field in cls.list_fields()} return cls(**kwargs) @classmethod def from_choice(cls, option: str): """Instantiate from choice.""" if option is None: 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}" kwargs = {field: 0 for field in cls.list_fields()} kwargs[option] = 1 return cls(**kwargs) @classmethod def from_tensor(cls, tensor: Tensor): """Instantiate from list of values.""" 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 += ')' return model_repr_str def highest_score_field(self) -> str: """Return name of field with highest score.""" model_dict = self.model_dump() return max(model_dict, key=lambda k: model_dict[k]) def highest_score_value(self) -> float: """Return value of field with highest score.""" model_dict = self.model_dump() return max(model_dict.values())