Files
NexQuant/rdagent/core/utils.py
T

167 lines
5.0 KiB
Python
Raw Normal View History

2024-05-21 22:48:41 +08:00
from __future__ import annotations
2024-05-30 10:33:07 +08:00
import importlib
2024-05-21 22:48:41 +08:00
import json
2024-05-30 10:33:07 +08:00
import multiprocessing as mp
2024-05-21 22:48:41 +08:00
import os
import random
import string
2024-05-30 10:33:07 +08:00
from collections.abc import Callable
2024-05-21 22:48:41 +08:00
from pathlib import Path
2024-06-12 15:12:11 +08:00
from typing import Any, ClassVar
2024-05-21 22:48:41 +08:00
import yaml
from fuzzywuzzy import fuzz
2024-06-12 15:12:11 +08:00
class RDAgentException(Exception): # noqa: N818
2024-05-21 22:48:41 +08:00
pass
class SingletonMeta(type):
2024-06-12 15:12:11 +08:00
_instance_dict: ClassVar[dict] = {}
2024-05-21 22:48:41 +08:00
2024-06-12 15:12:11 +08:00
def __call__(cls, *args: Any, **kwargs: Any) -> Any:
2024-06-05 15:36:15 +08:00
# Since it's hard to align the difference call using args and kwargs, we strictly ask to use kwargs in Singleton
2024-06-12 15:12:11 +08:00
if args:
exception_message = "Please only use kwargs in Singleton to avoid misunderstanding."
raise RDAgentException(exception_message)
2024-06-05 15:36:15 +08:00
kwargs_hash = hash(tuple(sorted(kwargs.items())))
if kwargs_hash not in cls._instance_dict:
2024-06-12 15:12:11 +08:00
cls._instance_dict[kwargs_hash] = super().__call__(**kwargs)
2024-06-05 15:36:15 +08:00
return cls._instance_dict[kwargs_hash]
2024-05-21 22:48:41 +08:00
class SingletonBaseClass(metaclass=SingletonMeta):
"""
2024-06-12 15:12:11 +08:00
Because we try to support defining Singleton with `class A(SingletonBaseClass)`
instead of `A(metaclass=SingletonMeta)` this class becomes necessary.
2024-05-21 22:48:41 +08:00
"""
# TODO: Add move this class to Qlib's general utils.
2024-06-12 15:12:11 +08:00
def parse_json(response: str) -> Any:
2024-05-21 22:48:41 +08:00
try:
return json.loads(response)
except json.decoder.JSONDecodeError:
pass
2024-06-12 15:12:11 +08:00
error_message = f"Failed to parse response: {response}, please report it or help us to fix it."
raise ValueError(error_message)
2024-05-21 22:48:41 +08:00
2024-06-12 15:12:11 +08:00
def similarity(text1: str, text2: str) -> int:
2024-05-21 22:48:41 +08:00
text1 = text1 if isinstance(text1, str) else ""
text2 = text2 if isinstance(text2, str) else ""
# Maybe we can use other similarity algorithm such as tfidf
return fuzz.ratio(text1, text2)
2024-06-12 15:12:11 +08:00
def random_string(length: int = 10) -> str:
2024-05-21 22:48:41 +08:00
letters = string.ascii_letters + string.digits
2024-06-12 15:12:11 +08:00
return "".join(random.SystemRandom().choice(letters) for _ in range(length))
2024-05-21 22:48:41 +08:00
2024-06-12 15:12:11 +08:00
def remove_uncommon_keys(new_dict: dict, org_dict: dict) -> None:
2024-05-21 22:48:41 +08:00
keys_to_remove = []
for key in new_dict:
if key not in org_dict:
keys_to_remove.append(key)
elif isinstance(new_dict[key], dict) and isinstance(org_dict[key], dict):
remove_uncommon_keys(new_dict[key], org_dict[key])
elif isinstance(new_dict[key], dict) and isinstance(org_dict[key], str):
new_dict[key] = org_dict[key]
for key in keys_to_remove:
del new_dict[key]
2024-06-12 15:12:11 +08:00
def crawl_the_folder(folder_path: Path) -> list:
2024-05-21 22:48:41 +08:00
yaml_files = []
for root, _, files in os.walk(folder_path.as_posix()):
for file in files:
2024-06-12 15:12:11 +08:00
if file.endswith((".yaml", ".yml")):
yaml_file_path = Path(root) / file
yaml_files.append(str(yaml_file_path.relative_to(folder_path)))
2024-05-21 22:48:41 +08:00
return sorted(yaml_files)
2024-06-12 15:12:11 +08:00
def compare_yaml(file1: Path | str, file2: Path | str) -> bool:
with Path(file1).open() as stream:
2024-05-21 22:48:41 +08:00
data1 = yaml.safe_load(stream)
2024-06-12 15:12:11 +08:00
with Path(file2).open() as stream:
2024-05-21 22:48:41 +08:00
data2 = yaml.safe_load(stream)
return data1 == data2
2024-06-12 15:12:11 +08:00
def remove_keys(valid_keys: set[Any], ori_dict: dict[Any, Any]) -> dict[Any, Any]:
2024-05-21 22:48:41 +08:00
for key in list(ori_dict.keys()):
if key not in valid_keys:
ori_dict.pop(key)
return ori_dict
class YamlConfigCache(SingletonBaseClass):
def __init__(self) -> None:
super().__init__()
2024-06-12 15:12:11 +08:00
self.path_to_config = {}
2024-05-21 22:48:41 +08:00
2024-06-12 15:12:11 +08:00
def load(self, path: str) -> None:
with Path(path).open() as stream:
2024-05-21 22:48:41 +08:00
data = yaml.safe_load(stream)
self.path_to_config[path] = data
2024-06-12 15:12:11 +08:00
def __getitem__(self, path: str) -> Any:
2024-05-21 22:48:41 +08:00
if path not in self.path_to_config:
self.load(path)
return self.path_to_config[path]
def import_class(class_path: str) -> Any:
"""
Parameters
----------
class_path : str
class path like"scripts.factor_implementation.baselines.naive.one_shot.OneshotFactorGen"
Returns
-------
class of `class_path`
"""
module_path, class_name = class_path.rsplit(".", 1)
module = importlib.import_module(module_path)
return getattr(module, class_name)
def multiprocessing_wrapper(func_calls: list[tuple[Callable, tuple]], n: int) -> list:
"""It will use multiprocessing to call the functions in func_calls with the given parameters.
The results equals to `return [f(*args) for f, args in func_calls]`
It will not call multiprocessing if `n=1`
Parameters
----------
func_calls : List[Tuple[Callable, Tuple]]
the list of functions and their parameters
n : int
the number of subprocesses
Returns
-------
list
"""
if n == 1:
return [f(*args) for f, args in func_calls]
with mp.Pool(processes=n) as pool:
results = [pool.apply_async(f, args) for f, args in func_calls]
return [result.get() for result in results]
# You can test the above function
# def f(x):
# return x**2
#
# if __name__ == "__main__":
# print(multiprocessing_wrapper([(f, (i,)) for i in range(10)], 4))