diff --git a/zntrack/__init__.py b/zntrack/__init__.py index fc4c6d28..d1095be8 100644 --- a/zntrack/__init__.py +++ b/zntrack/__init__.py @@ -13,6 +13,7 @@ from zntrack.fields.fields import ( deps, deps_path, + local_config, metrics, metrics_path, outs, @@ -51,6 +52,7 @@ "outs", "metrics", "params", + "local_config", "deps", "plots", "outs_path", diff --git a/zntrack/fields/fields.py b/zntrack/fields/fields.py index ab235a80..578dd8ae 100644 --- a/zntrack/fields/fields.py +++ b/zntrack/fields/fields.py @@ -2,7 +2,7 @@ from zntrack.fields.dependency import Dependency from zntrack.fields.dvc.options import DVCOption, PlotsOption -from zntrack.fields.zn.options import Output, Params, Plots +from zntrack.fields.zn.options import LocalConfig, Output, Params, Plots # Serialized Fields @@ -53,6 +53,23 @@ def params(*args, **kwargs): return Params(*args, **kwargs) +def local_config(*args, **kwargs): + """Define a Node Parameter. + + Parameters + ---------- + args: any + A data object that is used as a parameter. + Typically, this should be a string or number. + The object is serialized and deserialized by ZnTrack + and stored in params.yaml. + see https://dvc.org/doc/command-reference/stage/add#-p + kwargs: dict + Additional keyword arguments. + """ + return LocalConfig(*args, **kwargs) + + def deps(*data): """Define a Node Dependency. diff --git a/zntrack/fields/zn/options.py b/zntrack/fields/zn/options.py index edb9dcc5..36691fe9 100644 --- a/zntrack/fields/zn/options.py +++ b/zntrack/fields/zn/options.py @@ -166,6 +166,18 @@ def get_stage_add_argument(self, instance: "Node") -> typing.List[tuple]: return [(f"--{self.dvc_option}", f"{file}:{instance.name}")] +class LocalConfig(Params): + def get_files(self, instance: "Node") -> list: + """Get the list of files affected by this field. + + Returns + ------- + list + A list of file paths. + """ + return [config.files.local_config] + + class Output(LazyField): """A field that is saved to disk.""" diff --git a/zntrack/utils/config.py b/zntrack/utils/config.py index 4ceae6cd..f21f22e0 100644 --- a/zntrack/utils/config.py +++ b/zntrack/utils/config.py @@ -21,6 +21,11 @@ class Files: params: Path = Path("params.yaml") dvc: Path = Path("dvc.yaml") + @property + def local_config(self) -> Path: + Path(".zntrack/").mkdir(parents=True, exist_ok=True) + return Path(".zntrack/config.local") + @dataclasses.dataclass class Config: