Files
NexQuant/rdagent/core/utils.py
T
Linlang 571b5304cb fix mypy error (#91)
* fix mypy error

* fix mypy error

* fix ruff error

* change command

* delete python 3.8&3.9 from CI

* change command

* Some modifications according to the comments

* Add literal type

* Update .github/workflows/ci.yml

* Some modifications according to the comments

* fix ruff error

* fix meta dict

* Fix type

* Some modifications according to the comments

* merge latest code

* Some modifications according to the comments

* Some modifications according to the comments

* fix ci error

* fix ruff error

* Update Makefile

* Update Makefile

---------

Co-authored-by: Ubuntu <debug@debug.qjtqi00gqezu1eqs55bqdrf51f.px.internal.cloudapp.net>
Co-authored-by: Young <afe.young@gmail.com>
Co-authored-by: you-n-g <you-n-g@users.noreply.github.com>
2024-07-25 15:20:04 +08:00

92 lines
2.9 KiB
Python

from __future__ import annotations
import importlib
import json
import multiprocessing as mp
from collections.abc import Callable
from typing import Any, ClassVar, cast
from fuzzywuzzy import fuzz # type: ignore[import-untyped]
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)
kwargs_hash = hash(tuple(sorted(kwargs.items())))
if kwargs_hash not in cls._instance_dict:
cls._instance_dict[kwargs_hash] = super().__new__(cls) # Corrected call
cls._instance_dict[kwargs_hash].__init__(**kwargs) # Ensure __init__ is called
return cls._instance_dict[kwargs_hash]
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]