Files
NexQuant/rdagent/core/utils.py
T
Xu Yang 768229427d feat: use unified pickle cacher & move llm config into a isolated config (#424)
* simplify RDAgent conf

* add unified cacher(untested)

* fix small bugs

* fix a bug

* fix a small bug in runner

* use hash_key = None to skip cache

* fix CI

* in factor execution, ignore cache when raise exception

* add file locker to avoid mp calling

* fix CI

* use function __module__ name as folder in cache
2024-10-14 17:34:09 +08:00

159 lines
5.8 KiB
Python

from __future__ import annotations
import functools
import importlib
import json
import multiprocessing as mp
import pickle
from collections.abc import Callable
from pathlib import Path
from typing import Any, ClassVar, NoReturn, cast
from filelock import FileLock
from fuzzywuzzy import fuzz # type: ignore[import-untyped]
from rdagent.core.conf import RD_AGENT_SETTINGS
class RDAgentException(Exception): # noqa: N818
pass
class SingletonBaseClass:
"""
Because we try to support defining Singleton with `class A(SingletonBaseClass)`
instead of `A(metaclass=SingletonMeta)` this class becomes necessary.
"""
_instance_dict: ClassVar[dict] = {}
def __new__(cls, *args: Any, **kwargs: Any) -> Any:
# Since it's hard to align the difference call using args and kwargs, we strictly ask to use kwargs in Singleton
if args:
# TODO: this restriction can be solved.
exception_message = "Please only use kwargs in Singleton to avoid misunderstanding."
raise RDAgentException(exception_message)
class_name = [(-1, f"{cls.__module__}.{cls.__name__}")]
args_l = [(i, args[i]) for i in args]
kwargs_l = sorted(kwargs.items())
all_args = class_name + args_l + kwargs_l
kwargs_hash = hash(tuple(all_args))
if kwargs_hash not in cls._instance_dict:
cls._instance_dict[kwargs_hash] = super().__new__(cls) # Corrected call
return cls._instance_dict[kwargs_hash]
def __reduce__(self) -> NoReturn:
"""
NOTE:
When loading an object from a pickle, the __new__ method does not receive the `kwargs`
it was initialized with. This makes it difficult to retrieve the correct singleton object.
Therefore, we have made it unpickable.
"""
msg = f"Instances of {self.__class__.__name__} cannot be pickled"
raise pickle.PicklingError(msg)
def parse_json(response: str) -> Any:
try:
return json.loads(response)
except json.decoder.JSONDecodeError:
pass
error_message = f"Failed to parse response: {response}, please report it or help us to fix it."
raise ValueError(error_message)
def similarity(text1: str, text2: str) -> int:
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 cast(int, fuzz.ratio(text1, text2)) # mypy does not reguard it as int
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]
def cache_with_pickle(hash_func: Callable, post_process_func: Callable | None = None) -> Callable:
"""
This decorator will cache the return value of the function with pickle.
The cache key is generated by the hash_func. The hash function returns a string or None.
If it returns None, the cache will not be used. The cache will be stored in the folder
specified by RD_AGENT_SETTINGS.pickle_cache_folder_path_str with name hash_key.pkl.
The post_process_func will be called with the original arguments and the cached result
to give each caller a chance to process the cached result. The post_process_func should
return the final result.
"""
def cache_decorator(func: Callable) -> Callable:
@functools.wraps(func)
def cache_wrapper(*args: Any, **kwargs: Any) -> Any:
if RD_AGENT_SETTINGS.cache_with_pickle:
target_folder = Path(RD_AGENT_SETTINGS.pickle_cache_folder_path_str) / func.__module__
target_folder.mkdir(parents=True, exist_ok=True)
hash_key = hash_func(*args, **kwargs)
if hash_key is not None and (target_folder / (hash_key + ".pkl")).exists():
with Path.open(
target_folder / (hash_key + ".pkl"),
"rb",
) as f:
cached_res = pickle.load(f)
return (
post_process_func(*args, cached_res=cached_res, **kwargs)
if post_process_func is not None
else cached_res
)
if hash_key is not None and RD_AGENT_SETTINGS.use_file_lock:
with FileLock(target_folder / (hash_key + ".lock")):
result = func(*args, **kwargs)
if hash_key is not None:
with Path.open(
target_folder / (hash_key + ".pkl"),
"wb",
) as f:
pickle.dump(result, f)
else:
result = func(*args, **kwargs)
return result
return cache_wrapper
return cache_decorator