mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-07-27 23:47:46 +00:00
fix: fix ExtendedSettingsConfigDict does not work (#660)
* refactor: Replace ExtendedSettingsConfigDict with SettingsConfigDict * lint * lint
This commit is contained in:
+1
-2
@@ -65,8 +65,7 @@ coverage.xml
|
||||
|
||||
# Django stuff:
|
||||
*.log
|
||||
/log/
|
||||
log*/
|
||||
/log*/
|
||||
local_settings.py
|
||||
db.sqlite3
|
||||
db.sqlite3-journal
|
||||
|
||||
@@ -1,11 +1,12 @@
|
||||
from pathlib import Path
|
||||
|
||||
from pydantic_settings import SettingsConfigDict
|
||||
|
||||
from rdagent.components.workflow.conf import BasePropSetting
|
||||
from rdagent.core.conf import ExtendedSettingsConfigDict
|
||||
|
||||
|
||||
class MedBasePropSetting(BasePropSetting):
|
||||
model_config = ExtendedSettingsConfigDict(env_prefix="DM_", protected_namespaces=())
|
||||
model_config = SettingsConfigDict(env_prefix="DM_", protected_namespaces=())
|
||||
|
||||
# 1) overriding the default
|
||||
scen: str = "rdagent.scenarios.data_mining.experiment.model_experiment.DMModelScenario"
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
from pydantic_settings import SettingsConfigDict
|
||||
|
||||
from rdagent.app.kaggle.conf import KaggleBasePropSetting
|
||||
from rdagent.core.conf import ExtendedSettingsConfigDict
|
||||
|
||||
|
||||
class DataScienceBasePropSetting(KaggleBasePropSetting):
|
||||
model_config = ExtendedSettingsConfigDict(env_prefix="DS_", protected_namespaces=())
|
||||
model_config = SettingsConfigDict(env_prefix="DS_", protected_namespaces=())
|
||||
|
||||
# Main components
|
||||
## Scen
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
from rdagent.core.conf import ExtendedBaseSettings, ExtendedSettingsConfigDict
|
||||
from pydantic_settings import SettingsConfigDict
|
||||
|
||||
from rdagent.core.conf import ExtendedBaseSettings
|
||||
|
||||
|
||||
class KaggleBasePropSetting(ExtendedBaseSettings):
|
||||
model_config = ExtendedSettingsConfigDict(env_prefix="KG_", protected_namespaces=())
|
||||
model_config = SettingsConfigDict(env_prefix="KG_", protected_namespaces=())
|
||||
|
||||
# 1) overriding the default
|
||||
scen: str = "rdagent.scenarios.kaggle.experiment.scenario.KGScenario"
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
from pydantic_settings import SettingsConfigDict
|
||||
|
||||
from rdagent.components.workflow.conf import BasePropSetting
|
||||
from rdagent.core.conf import ExtendedSettingsConfigDict
|
||||
|
||||
|
||||
class ModelBasePropSetting(BasePropSetting):
|
||||
model_config = ExtendedSettingsConfigDict(env_prefix="QLIB_MODEL_", protected_namespaces=())
|
||||
model_config = SettingsConfigDict(env_prefix="QLIB_MODEL_", protected_namespaces=())
|
||||
|
||||
# 1) override base settings
|
||||
scen: str = "rdagent.scenarios.qlib.experiment.model_experiment.QlibModelScenario"
|
||||
@@ -29,7 +30,7 @@ class ModelBasePropSetting(BasePropSetting):
|
||||
|
||||
|
||||
class FactorBasePropSetting(BasePropSetting):
|
||||
model_config = ExtendedSettingsConfigDict(env_prefix="QLIB_FACTOR_", protected_namespaces=())
|
||||
model_config = SettingsConfigDict(env_prefix="QLIB_FACTOR_", protected_namespaces=())
|
||||
|
||||
# 1) override base settings
|
||||
scen: str = "rdagent.scenarios.qlib.experiment.factor_experiment.QlibFactorScenario"
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
from pydantic_settings import SettingsConfigDict
|
||||
|
||||
from rdagent.components.coder.CoSTEER.config import CoSTEERSettings
|
||||
from rdagent.core.conf import ExtendedSettingsConfigDict
|
||||
|
||||
|
||||
class FactorCoSTEERSettings(CoSTEERSettings):
|
||||
model_config = ExtendedSettingsConfigDict(env_prefix="FACTOR_CoSTEER_")
|
||||
model_config = SettingsConfigDict(env_prefix="FACTOR_CoSTEER_")
|
||||
|
||||
data_folder: str = "git_ignore_folder/factor_implementation_source_data"
|
||||
"""Path to the folder containing financial data (default is fundamental data in Qlib)"""
|
||||
|
||||
+25
-32
@@ -2,53 +2,46 @@ from __future__ import annotations
|
||||
|
||||
# TODO: use pydantic for other modules in Qlib
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pydantic.fields import FieldInfo
|
||||
from typing import cast
|
||||
|
||||
from pydantic_settings import (
|
||||
BaseSettings,
|
||||
EnvSettingsSource,
|
||||
PydanticBaseSettingsSource,
|
||||
SettingsConfigDict,
|
||||
)
|
||||
|
||||
|
||||
class ExtendedEnvSettingsSource(EnvSettingsSource):
|
||||
def get_field_value(self, field: FieldInfo, field_name: str) -> tuple[Any, str, bool]:
|
||||
# Dynamically gather prefixes from the current and parent classes
|
||||
prefixes = [self.config.get("env_prefix", "")]
|
||||
if hasattr(self.settings_cls, "__bases__"):
|
||||
for base in self.settings_cls.__bases__:
|
||||
if hasattr(base, "model_config"):
|
||||
parent_prefix = base.model_config.get("env_prefix")
|
||||
if parent_prefix and parent_prefix not in prefixes:
|
||||
prefixes.append(parent_prefix)
|
||||
for prefix in prefixes:
|
||||
self.env_prefix = prefix
|
||||
env_val, field_key, value_is_complex = super().get_field_value(field, field_name)
|
||||
if env_val is not None:
|
||||
return env_val, field_key, value_is_complex
|
||||
|
||||
return super().get_field_value(field, field_name)
|
||||
|
||||
|
||||
class ExtendedSettingsConfigDict(SettingsConfigDict, total=False): ...
|
||||
|
||||
|
||||
class ExtendedBaseSettings(BaseSettings):
|
||||
|
||||
@classmethod
|
||||
def settings_customise_sources(
|
||||
cls,
|
||||
settings_cls: type[BaseSettings],
|
||||
init_settings: PydanticBaseSettingsSource, # noqa
|
||||
env_settings: PydanticBaseSettingsSource, # noqa
|
||||
dotenv_settings: PydanticBaseSettingsSource, # noqa
|
||||
file_secret_settings: PydanticBaseSettingsSource, # noqa
|
||||
init_settings: PydanticBaseSettingsSource,
|
||||
env_settings: PydanticBaseSettingsSource,
|
||||
dotenv_settings: PydanticBaseSettingsSource,
|
||||
file_secret_settings: PydanticBaseSettingsSource,
|
||||
) -> tuple[PydanticBaseSettingsSource, ...]:
|
||||
return (ExtendedEnvSettingsSource(settings_cls),)
|
||||
# 1) walk from base class
|
||||
def base_iter(settings_cls: type[ExtendedBaseSettings]) -> list[type[ExtendedBaseSettings]]:
|
||||
bases = []
|
||||
for cl in settings_cls.__bases__:
|
||||
if issubclass(cl, ExtendedBaseSettings) and cl is not ExtendedBaseSettings:
|
||||
bases.append(cl)
|
||||
bases.extend(base_iter(cl))
|
||||
return bases
|
||||
|
||||
# 2) Build EnvSettingsSource from base classes, so we can add parent Env Sources
|
||||
parent_env_settings = [
|
||||
EnvSettingsSource(
|
||||
base_cls,
|
||||
case_sensitive=base_cls.model_config.get("case_sensitive"),
|
||||
env_prefix=base_cls.model_config.get("env_prefix"),
|
||||
env_nested_delimiter=base_cls.model_config.get("env_nested_delimiter"),
|
||||
)
|
||||
for base_cls in base_iter(cast(type[ExtendedBaseSettings], settings_cls))
|
||||
]
|
||||
return init_settings, env_settings, *parent_env_settings, dotenv_settings, file_secret_settings
|
||||
|
||||
|
||||
class RDAgentSettings(ExtendedBaseSettings):
|
||||
|
||||
@@ -26,13 +26,14 @@ import docker.models # type: ignore[import-untyped]
|
||||
import docker.models.containers # type: ignore[import-untyped]
|
||||
import docker.types # type: ignore[import-untyped]
|
||||
from pydantic import BaseModel
|
||||
from pydantic_settings import SettingsConfigDict
|
||||
from rich import print
|
||||
from rich.console import Console
|
||||
from rich.progress import Progress, SpinnerColumn, TextColumn
|
||||
from rich.rule import Rule
|
||||
from rich.table import Table
|
||||
|
||||
from rdagent.core.conf import ExtendedBaseSettings, ExtendedSettingsConfigDict
|
||||
from rdagent.core.conf import ExtendedBaseSettings
|
||||
from rdagent.core.experiment import RD_AGENT_SETTINGS
|
||||
from rdagent.log import rdagent_logger as logger
|
||||
from rdagent.oai.llm_utils import md5_hash
|
||||
@@ -186,7 +187,7 @@ class DockerConf(ExtendedBaseSettings):
|
||||
|
||||
|
||||
class QlibDockerConf(DockerConf):
|
||||
model_config = ExtendedSettingsConfigDict(env_prefix="QLIB_DOCKER_")
|
||||
model_config = SettingsConfigDict(env_prefix="QLIB_DOCKER_")
|
||||
|
||||
build_from_dockerfile: bool = True
|
||||
dockerfile_folder_path: Path = Path(__file__).parent.parent / "scenarios" / "qlib" / "docker"
|
||||
@@ -199,7 +200,7 @@ class QlibDockerConf(DockerConf):
|
||||
|
||||
|
||||
class DMDockerConf(DockerConf):
|
||||
model_config = ExtendedSettingsConfigDict(env_prefix="DM_DOCKER_")
|
||||
model_config = SettingsConfigDict(env_prefix="DM_DOCKER_")
|
||||
|
||||
build_from_dockerfile: bool = True
|
||||
dockerfile_folder_path: Path = Path(__file__).parent.parent / "scenarios" / "data_mining" / "docker"
|
||||
@@ -218,7 +219,7 @@ class DMDockerConf(DockerConf):
|
||||
|
||||
|
||||
class KGDockerConf(DockerConf):
|
||||
model_config = ExtendedSettingsConfigDict(env_prefix="KG_DOCKER_")
|
||||
model_config = SettingsConfigDict(env_prefix="KG_DOCKER_")
|
||||
|
||||
build_from_dockerfile: bool = True
|
||||
dockerfile_folder_path: Path = Path(__file__).parent.parent / "scenarios" / "kaggle" / "docker" / "kaggle_docker"
|
||||
@@ -238,7 +239,7 @@ class KGDockerConf(DockerConf):
|
||||
|
||||
|
||||
class DSDockerConf(DockerConf):
|
||||
model_config = ExtendedSettingsConfigDict(env_prefix="DS_DOCKER_")
|
||||
model_config = SettingsConfigDict(env_prefix="DS_DOCKER_")
|
||||
|
||||
build_from_dockerfile: bool = False
|
||||
image: str = "gcr.io/kaggle-gpu-images/python:latest"
|
||||
@@ -252,7 +253,7 @@ class DSDockerConf(DockerConf):
|
||||
|
||||
|
||||
class MLEBDockerConf(DockerConf):
|
||||
model_config = ExtendedSettingsConfigDict(env_prefix="MLEB_DOCKER_")
|
||||
model_config = SettingsConfigDict(env_prefix="MLEB_DOCKER_")
|
||||
|
||||
build_from_dockerfile: bool = True
|
||||
dockerfile_folder_path: Path = Path(__file__).parent.parent / "scenarios" / "kaggle" / "docker" / "mle_bench_docker"
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
import unittest
|
||||
|
||||
|
||||
class ConfUtils(unittest.TestCase):
|
||||
|
||||
def test_conf(self):
|
||||
import os
|
||||
|
||||
from rdagent.utils.env import QlibDockerConf
|
||||
|
||||
os.environ["MEM_LIMIT"] = "200g"
|
||||
assert QlibDockerConf().mem_limit == "200g" # base class will affect subclasses
|
||||
os.environ["QLIB_DOCKER_MEM_LIMIT"] = "300g"
|
||||
assert QlibDockerConf().mem_limit == "300g" # more accurate subclass will override the base class
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user