172 lines
6.3 KiB
Python
172 lines
6.3 KiB
Python
import sys
|
|
import inspect
|
|
import hashlib
|
|
import logging
|
|
from functools import partial
|
|
import random
|
|
from typing import Callable
|
|
import importlib
|
|
|
|
import numpy as np
|
|
|
|
|
|
def get_logger(module_name: str, filename: str = "main.log") -> logging.Logger:
|
|
"""
|
|
Returns a logger for the given module name and filename.
|
|
:param module_name: name of the module, as string
|
|
:param filename: name of the logging file, as string
|
|
:return: the logger, as logging.Logger object
|
|
"""
|
|
logger = logging.getLogger(module_name)
|
|
logger.setLevel(logging.DEBUG)
|
|
logger.propagate = False
|
|
formatter = logging.Formatter('%(asctime)s - %(name)s - %(levelname)s - %(message)s')
|
|
|
|
file_handler = logging.FileHandler(filename)
|
|
file_handler.setLevel(logging.DEBUG)
|
|
file_handler.setFormatter(formatter)
|
|
|
|
logger.addHandler(file_handler)
|
|
|
|
# add system out handler
|
|
stream_handler = logging.StreamHandler(sys.stdout)
|
|
stream_handler.setLevel(logging.INFO)
|
|
stream_handler.setFormatter(formatter)
|
|
|
|
logger.addHandler(stream_handler)
|
|
|
|
return logger
|
|
|
|
|
|
def get_variable_from_module(module_path: str, variable_name: str):
|
|
"""
|
|
Get a variable from a module by its name during runtime
|
|
|
|
Args:
|
|
module_path: module to fetch variable from
|
|
variable_name: variable to fetch from module
|
|
|
|
Returns:
|
|
variable from module
|
|
|
|
"""
|
|
module = importlib.import_module(module_path)
|
|
variable = getattr(module, variable_name)
|
|
if variable is None:
|
|
raise ValueError(f"Variable {variable_name} not found in module {module_path}")
|
|
return variable
|
|
|
|
|
|
def get_callable_name(callable_obj):
|
|
"""
|
|
Get the name of the callable, handling `functools.partial`.
|
|
|
|
:param callable_obj: callable object
|
|
:return: name of the callable
|
|
"""
|
|
if isinstance(callable_obj, partial):
|
|
func_name = callable_obj.func.__name__
|
|
args = ", ".join(object_to_string(arg) for arg in callable_obj.args)
|
|
kwargs = ", ".join(f"{object_to_string(k)}={object_to_string(v)!r}" for k, v in callable_obj.keywords.items())
|
|
return f"partial({func_name}, {args}, {kwargs})"
|
|
else:
|
|
if hasattr(callable_obj, '__name__'):
|
|
return callable_obj.__name__
|
|
elif hasattr(callable_obj, '__class__'):
|
|
return callable_obj.__class__.__name__
|
|
elif hasattr(callable_obj, '__hash__'):
|
|
return callable_obj.__hash__
|
|
else:
|
|
raise ValueError(f"Could not determine name of callable object {callable_obj}")
|
|
|
|
|
|
def object_to_string(value, skip_types=None):
|
|
"""
|
|
Convert any object to a string representation that avoids memory addresses.
|
|
Handles complex data types recursively.
|
|
|
|
:param value: object to convert
|
|
:param skip_types: types to skip during conversion
|
|
:return: string representation of the object
|
|
"""
|
|
if skip_types is None:
|
|
skip_types = []
|
|
|
|
if any(isinstance(value, t) for t in skip_types):
|
|
return 'skipped_type'
|
|
elif isinstance(value, (str, int, float, bool)): # Handle primitive data types directly
|
|
return repr(value)
|
|
elif isinstance(value, dict):
|
|
return '{' + ', '.join(f"{k}: {object_to_string(v, skip_types)}" for k, v in value.items()) + '}'
|
|
elif isinstance(value, (list, tuple)):
|
|
return '[' + ', '.join(object_to_string(item, skip_types) for item in value) + ']'
|
|
elif isinstance(value, partial):
|
|
return get_callable_name(value)
|
|
elif inspect.isclass(value):
|
|
return f"<class '{value.__name__}'>"
|
|
elif hasattr(value,
|
|
'__class__') and not value.__class__ != "function": # Correct handling for instances of classes, but not functions
|
|
return f"<instance of class '{value.__class__.__name__}'>"
|
|
elif isinstance(value, Callable):
|
|
return f"<callable '{get_callable_name(value)}'>"
|
|
else:
|
|
return repr(value)
|
|
|
|
|
|
def get_config_id(configuration: dict) -> str:
|
|
"""
|
|
Generate a somewhat unique human-readable model name from the model and training parameters.
|
|
:param configuration: dictionary containing model and training parameters
|
|
:return: human-readable model name
|
|
"""
|
|
adjectives = ["autumn", "hidden", "bitter", "misty", "silent", "empty", "dry", "dark", "summer", "icy", "delicate",
|
|
"quiet", "white", "black", "blue", "green", "red", "yellow",
|
|
"purple", "orange", "pink", "golden", "silver", "crimson", "violet", "azure", "amber", "sapphire",
|
|
"emerald", "ruby", "pearl", "topaz", "onyx", "turquoise", "citrine", ]
|
|
nouns = ["waterfall", "river", "breeze", "moon", "rain", "wind", "sea", "morning", "snow", "lake", "sunset", "pine",
|
|
"shadow", "leaf", "dawn", "glitter", "forest", "cloud", "sky", "sun", "butterfly",
|
|
"flower", "bird", "mountain", "valley", "ocean", "star", "night", "dream", "whisper", "echo", "horizon",
|
|
"wave", "petal", "dew", "mist"]
|
|
|
|
# also add dataset config, but skip functions
|
|
base_name = object_to_string(configuration)
|
|
|
|
# hash long name
|
|
basename_hash = hashlib.md5(base_name.encode()).hexdigest()
|
|
|
|
# select adjective and noun based on hash
|
|
random.seed(int(basename_hash, 16))
|
|
model_name = f"{random.choice(adjectives)}_{random.choice(nouns)}_{basename_hash[:5]}"
|
|
|
|
return model_name
|
|
|
|
|
|
def convert_for_json(obj):
|
|
if isinstance(obj, dict):
|
|
return {k: convert_for_json(v) for k, v in obj.items()}
|
|
elif isinstance(obj, list):
|
|
return [convert_for_json(v) for v in obj]
|
|
elif isinstance(obj, np.generic):
|
|
return obj.item()
|
|
else:
|
|
return obj
|
|
|
|
|
|
def recursive_dict_update(current_dict: dict, update_dict: dict) -> dict:
|
|
for key, value in update_dict.items():
|
|
if key not in current_dict:
|
|
current_dict[key] = value
|
|
else:
|
|
if isinstance(value, list):
|
|
if not isinstance(current_dict[key], list):
|
|
raise ValueError(f"Type mismatch, {type(current_dict[key])} is not a list")
|
|
current_dict[key] = current_dict[key] + value
|
|
elif isinstance(value, dict):
|
|
if not isinstance(current_dict[key], dict):
|
|
raise ValueError(f"Type mismatch, {type(current_dict[key])} is not a dict")
|
|
current_dict[key] = recursive_dict_update(current_dict[key], value)
|
|
else:
|
|
current_dict[key] = value
|
|
|
|
return current_dict
|