"""Pipeline task registry and scenario definitions. A *task* is a callable ``task(tracker, logger, config, state) -> None``. A *scenario* is an ordered list of task ids the pipeline runs. Tasks communicate through the mutable ``state`` dict (dataloaders, model paths, control events). Pipeline stages (see ai/ARCHITECTURE.md): 1. load_data -> implemented 2. train_regular -> implemented 3. train_bayesian -> implemented 4. evaluate_regular -> implemented 5. evaluate_bayesian-> implemented 6. evaluate_noisy -> implemented 7. run_analysis -> skeleton (step 7); reads evaluations, writes analysis/. load_models -> implemented; discovers saved .pt ensembles on disk and records their paths in ``state`` so evaluation is decoupled from training (training frees models from VRAM after saving). Evaluation loads one model at a time. """ import os from typing import Any, Callable, Dict, List, Tuple, TypedDict from . import analysis from . import evaluate from . import load_data from . import load_models from . import train_bayesian from . import train_normal class TaskEntry(TypedDict): task_name: str task_func: Callable[..., None] class ScenarioEntry(TypedDict): label: str tasks: List[str] PIPELINE_TASKS: Dict[str, TaskEntry] = { "load_data": { "task_name": "Load Image and ADNIMERGE", "task_func": load_data.load_data_task, }, "train_regular": { "task_name": "Train Regular Models", "task_func": train_normal.train_normal_task, }, "train_bayesian": { "task_name": "Train Bayesian Models", "task_func": train_bayesian.train_bayesian_task, }, "load_models": { "task_name": "Load Saved Models", "task_func": load_models.load_models_task, }, "evaluate_regular": { "task_name": "Evaluate Regular Models", "task_func": evaluate.evaluate_normal_task, }, "evaluate_bayesian": { "task_name": "Evaluate Bayesian Models", "task_func": evaluate.evaluate_bayesian_task, }, "evaluate_noisy": { "task_name": "Evaluate Models on Noised Data", "task_func": evaluate.evaluate_noisy_task, }, "run_analysis": { "task_name": "Analyze Evaluations", "task_func": analysis.run_analysis_task, }, } # -------------------------------------------------------------------------- # Output artifact categories and overwrite protection # -------------------------------------------------------------------------- # Each category maps to the work_dir subdirectories it owns and whether it is # "protected" -- i.e. expensive/slow to regenerate, so overwriting it warrants a # confirmation prompt. Analysis artifacts are cheap to regenerate and therefore # unprotected: reruns of an analysis-only scenario proceed without a warning. class ArtifactCategory(TypedDict): dirs: Tuple[str, ...] protected: bool desc: str ARTIFACT_CATEGORIES: Dict[str, ArtifactCategory] = { "models": { "dirs": ("normal_models", "bayesian_models"), "protected": True, "desc": "trained model checkpoints", }, "evaluations": { "dirs": ("evaluations",), "protected": True, "desc": "model evaluation outputs (netCDF)", }, "analysis": { "dirs": (analysis.ANALYSIS_SUBDIR,), "protected": False, "desc": "analysis artifacts (regeneratable)", }, } # Which artifact category each task writes to disk. Tasks not listed here (e.g. # load_data, load_models) produce no persistent, overwriteable output. TASK_WRITES = { "train_regular": ("models",), "train_bayesian": ("models",), "evaluate_regular": ("evaluations",), "evaluate_bayesian": ("evaluations",), "evaluate_noisy": ("evaluations",), "run_analysis": ("analysis",), } def scenario_output_categories(scenario_id: str) -> List[str]: """Return the artifact categories a scenario writes, in first-seen order.""" categories: List[str] = [] for task_id in SCENARIOS[scenario_id]["tasks"]: for category in TASK_WRITES.get(task_id, ()): if category not in categories: categories.append(category) return categories def protected_output_conflicts( work_dir: str, scenario_id: str ) -> List[Tuple[str, str]]: """Return ``(category, subdir)`` for PROTECTED outputs the scenario writes that already exist and are non-empty. Empty result => safe to run without a destructive-overwrite prompt (only unprotected/regeneratable outputs, if any, would be replaced).""" conflicts: List[Tuple[str, str]] = [] for category in scenario_output_categories(scenario_id): meta = ARTIFACT_CATEGORIES.get(category) if not meta or not meta["protected"]: continue for subdir in meta["dirs"]: path = os.path.join(work_dir, subdir) if os.path.isdir(path) and os.listdir(path): conflicts.append((category, subdir)) return conflicts def describe_protected_conflicts(work_dir: str, scenario_id: str) -> List[str]: """Human-readable one-liners for each protected overwrite conflict.""" lines: List[str] = [] for category, subdir in protected_output_conflicts(work_dir, scenario_id): desc = ARTIFACT_CATEGORIES[category]["desc"] lines.append(f"{subdir}/ — {desc}") return lines SCENARIOS: Dict[str, ScenarioEntry] = { "scen_train_all": { "label": "1. Train, Evaluate, & Noise Analysis", "tasks": [ "load_data", "train_regular", "train_bayesian", "load_models", "evaluate_regular", "evaluate_bayesian", "evaluate_noisy", ], }, "scen_load_all": { "label": "2. Load, Evaluate, & Noise Analysis", "tasks": [ "load_data", "load_models", "evaluate_regular", "evaluate_bayesian", "evaluate_noisy", ], }, "scen_load_eval": { "label": "3. Load & Evaluate (Skip Noise)", "tasks": [ "load_data", "load_models", "evaluate_regular", "evaluate_bayesian", ], }, "scen_analyze": { "label": "4. Analyze Existing Evaluations", "tasks": [ "run_analysis", ], }, }