From 541c1ff990bc4747a11c970ef9fb4d9a319c4cb7 Mon Sep 17 00:00:00 2001 From: Linlang <30293408+SunsetWolf@users.noreply.github.com> Date: Tue, 9 Jul 2024 12:45:32 +0800 Subject: [PATCH] Support LocalEnv (#53) * Support LocalEnv * resolve issues * resolve issues * resolve issues * resolve issues * resolve issues * resolve issues --- rdagent/utils/env.py | 48 ++++++++++++++++++++++++++++++++++++------ test/utils/test_env.py | 18 +++++++++++++--- 2 files changed, 57 insertions(+), 9 deletions(-) diff --git a/rdagent/utils/env.py b/rdagent/utils/env.py index e334db4d..d3835943 100644 --- a/rdagent/utils/env.py +++ b/rdagent/utils/env.py @@ -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 ----- diff --git a/test/utils/test_env.py b/test/utils/test_env.py index ba5e44cd..16cdce5e 100644 --- a/test/utils/test_env.py +++ b/test/utils/test_env.py @@ -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.