mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-07-27 23:47:46 +00:00
Support LocalEnv (#53)
* Support LocalEnv * resolve issues * resolve issues * resolve issues * resolve issues * resolve issues * resolve issues
This commit is contained in:
+42
-6
@@ -6,12 +6,13 @@ Tries to create uniform environment for the agent to run;
|
||||
|
||||
"""
|
||||
import os
|
||||
from abc import abstractmethod
|
||||
from pathlib import Path
|
||||
from typing import Generic, TypeVar
|
||||
|
||||
import sys
|
||||
import docker
|
||||
import subprocess
|
||||
from abc import abstractmethod
|
||||
from pydantic import BaseModel
|
||||
from typing import Generic, TypeVar, Optional, Dict
|
||||
from pathlib import Path
|
||||
|
||||
ASpecificBaseModel = TypeVar("ASpecificBaseModel", bound=BaseModel)
|
||||
|
||||
@@ -62,15 +63,50 @@ class Env(Generic[ASpecificBaseModel]):
|
||||
|
||||
|
||||
class LocalConf(BaseModel):
|
||||
py_entry: str # where you can find your python path
|
||||
py_bin: str
|
||||
default_entry: str
|
||||
|
||||
|
||||
class LocalEnv(Env[LocalConf]):
|
||||
"""
|
||||
Sometimes local environment may be more convinient for testing
|
||||
"""
|
||||
def prepare(self):
|
||||
if not (Path("~/.qlib/qlib_data/cn_data").expanduser().resolve().exists()):
|
||||
self.run(
|
||||
entry="python -m qlib.run.get_data qlib_data --target_dir ~/.qlib/qlib_data/cn_data --region cn",
|
||||
)
|
||||
else:
|
||||
print("Data already exists. Download skipped.")
|
||||
|
||||
conf: LocalConf
|
||||
def run(self,
|
||||
entry: str | None = None,
|
||||
local_path: Optional[str] = None,
|
||||
env: dict | None = None) -> str:
|
||||
if env is None:
|
||||
env = {}
|
||||
|
||||
if entry is None:
|
||||
entry = self.conf.default_entry
|
||||
|
||||
command = str(Path(self.conf.py_bin).joinpath(entry)).split(" ")
|
||||
|
||||
cwd = None
|
||||
if local_path:
|
||||
cwd = Path(local_path).resolve()
|
||||
print(command)
|
||||
result = subprocess.run(
|
||||
command,
|
||||
cwd=cwd,
|
||||
env={**os.environ, **env},
|
||||
capture_output=True,
|
||||
text=True
|
||||
)
|
||||
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError(f"Error while running the command: {result.stderr}")
|
||||
|
||||
return result.stdout
|
||||
|
||||
|
||||
## Docker Environment -----
|
||||
|
||||
+15
-3
@@ -1,13 +1,11 @@
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.append(str(Path(__file__).resolve().parent.parent))
|
||||
from rdagent.utils.env import QTDockerEnv, LocalEnv, LocalConf
|
||||
import shutil
|
||||
|
||||
from rdagent.utils.env import QTDockerEnv
|
||||
|
||||
DIRNAME = Path(__file__).absolute().resolve().parent
|
||||
|
||||
@@ -23,6 +21,20 @@ class EnvUtils(unittest.TestCase):
|
||||
# shutil.rmtree(mlrun_p)
|
||||
...
|
||||
|
||||
# NOTE: Since I don't know the exact environment in which it will be used, here's just an example.
|
||||
# NOTE: Because you need to download the data during the prepare process. So you need to have pyqlib in your environment.
|
||||
# def test_local(self):
|
||||
# local_conf = LocalConf(
|
||||
# py_bin="/home/v-linlanglv/miniconda3/envs/RD-Agent-310/bin",
|
||||
# default_entry="qrun conf.yaml",
|
||||
# )
|
||||
# qle = LocalEnv(conf=local_conf)
|
||||
# qle.prepare()
|
||||
# conf_path = str(DIRNAME / "env_tpl" / "conf.yaml")
|
||||
# qle.run(entry="qrun " + conf_path)
|
||||
# mlrun_p = DIRNAME / "env_tpl" / "mlruns"
|
||||
# self.assertTrue(mlrun_p.exists(), f"Expected output file {mlrun_p} not found")
|
||||
|
||||
def test_docker(self):
|
||||
"""
|
||||
We will mount `env_tpl` into the docker image.
|
||||
|
||||
Reference in New Issue
Block a user