| """Project-local default paths shared by ACE command-line entrypoints.""" |
|
|
| from __future__ import annotations |
|
|
| from pathlib import Path |
|
|
|
|
| PROJECT_ROOT = Path(__file__).resolve().parents[1] |
| DATA_DIR = PROJECT_ROOT / "data" |
| GENERATED_DATA_DIR = DATA_DIR / "data" |
| GENERATED_DATA_PATH = GENERATED_DATA_DIR / "ace_pairs.npz" |
| CHECKPOINT_DIR = DATA_DIR / "checkpoint" |
| CHECKPOINT_PATH = CHECKPOINT_DIR / "model_bak.pt" |
| OUTPUT_DIR = PROJECT_ROOT / "output" |
| INFER_DIR = OUTPUT_DIR / "infer" |
| INFER_PATH = INFER_DIR / "rollout.npz" |
| PIC_DIR = OUTPUT_DIR / "pic" |
|
|
|
|
| def resolve_project_path(value: str | Path, *, base: Path = PROJECT_ROOT) -> Path: |
| """Resolve a configured path relative to the ACE project root.""" |
| path = Path(value).expanduser() |
| return path if path.is_absolute() else base / path |
|
|
|
|
| def configured_path(config: dict, key: str, fallback: Path) -> Path: |
| """Read a path from ``config.paths`` while preserving a safe project default.""" |
| value = config.get("paths", {}).get(key) |
| return resolve_project_path(value) if value else fallback |
|
|