__init__.py 6.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195
  1. """Pipeline task registry and scenario definitions.
  2. A *task* is a callable ``task(tracker, logger, config, state) -> None``. A
  3. *scenario* is an ordered list of task ids the pipeline runs. Tasks communicate
  4. through the mutable ``state`` dict (dataloaders, model paths, control events).
  5. Pipeline stages (see ai/ARCHITECTURE.md):
  6. 1. load_data -> implemented
  7. 2. train_regular -> implemented
  8. 3. train_bayesian -> implemented
  9. 4. evaluate_regular -> implemented
  10. 5. evaluate_bayesian-> implemented
  11. 6. evaluate_noisy -> implemented
  12. 7. run_analysis -> skeleton (step 7); reads evaluations, writes analysis/.
  13. load_models -> implemented; discovers saved .pt ensembles on disk and
  14. records their paths in ``state`` so evaluation is
  15. decoupled from training (training frees models from VRAM
  16. after saving). Evaluation loads one model at a time.
  17. """
  18. import os
  19. from typing import Any, Callable, Dict, List, Tuple, TypedDict
  20. from . import analysis
  21. from . import evaluate
  22. from . import load_data
  23. from . import load_models
  24. from . import train_bayesian
  25. from . import train_normal
  26. class TaskEntry(TypedDict):
  27. task_name: str
  28. task_func: Callable[..., None]
  29. class ScenarioEntry(TypedDict):
  30. label: str
  31. tasks: List[str]
  32. PIPELINE_TASKS: Dict[str, TaskEntry] = {
  33. "load_data": {
  34. "task_name": "Load Image and ADNIMERGE",
  35. "task_func": load_data.load_data_task,
  36. },
  37. "train_regular": {
  38. "task_name": "Train Regular Models",
  39. "task_func": train_normal.train_normal_task,
  40. },
  41. "train_bayesian": {
  42. "task_name": "Train Bayesian Models",
  43. "task_func": train_bayesian.train_bayesian_task,
  44. },
  45. "load_models": {
  46. "task_name": "Load Saved Models",
  47. "task_func": load_models.load_models_task,
  48. },
  49. "evaluate_regular": {
  50. "task_name": "Evaluate Regular Models",
  51. "task_func": evaluate.evaluate_normal_task,
  52. },
  53. "evaluate_bayesian": {
  54. "task_name": "Evaluate Bayesian Models",
  55. "task_func": evaluate.evaluate_bayesian_task,
  56. },
  57. "evaluate_noisy": {
  58. "task_name": "Evaluate Models on Noised Data",
  59. "task_func": evaluate.evaluate_noisy_task,
  60. },
  61. "run_analysis": {
  62. "task_name": "Analyze Evaluations",
  63. "task_func": analysis.run_analysis_task,
  64. },
  65. }
  66. # --------------------------------------------------------------------------
  67. # Output artifact categories and overwrite protection
  68. # --------------------------------------------------------------------------
  69. # Each category maps to the work_dir subdirectories it owns and whether it is
  70. # "protected" -- i.e. expensive/slow to regenerate, so overwriting it warrants a
  71. # confirmation prompt. Analysis artifacts are cheap to regenerate and therefore
  72. # unprotected: reruns of an analysis-only scenario proceed without a warning.
  73. class ArtifactCategory(TypedDict):
  74. dirs: Tuple[str, ...]
  75. protected: bool
  76. desc: str
  77. ARTIFACT_CATEGORIES: Dict[str, ArtifactCategory] = {
  78. "models": {
  79. "dirs": ("normal_models", "bayesian_models"),
  80. "protected": True,
  81. "desc": "trained model checkpoints",
  82. },
  83. "evaluations": {
  84. "dirs": ("evaluations",),
  85. "protected": True,
  86. "desc": "model evaluation outputs (netCDF)",
  87. },
  88. "analysis": {
  89. "dirs": (analysis.ANALYSIS_SUBDIR,),
  90. "protected": False,
  91. "desc": "analysis artifacts (regeneratable)",
  92. },
  93. }
  94. # Which artifact category each task writes to disk. Tasks not listed here (e.g.
  95. # load_data, load_models) produce no persistent, overwriteable output.
  96. TASK_WRITES = {
  97. "train_regular": ("models",),
  98. "train_bayesian": ("models",),
  99. "evaluate_regular": ("evaluations",),
  100. "evaluate_bayesian": ("evaluations",),
  101. "evaluate_noisy": ("evaluations",),
  102. "run_analysis": ("analysis",),
  103. }
  104. def scenario_output_categories(scenario_id: str) -> List[str]:
  105. """Return the artifact categories a scenario writes, in first-seen order."""
  106. categories: List[str] = []
  107. for task_id in SCENARIOS[scenario_id]["tasks"]:
  108. for category in TASK_WRITES.get(task_id, ()):
  109. if category not in categories:
  110. categories.append(category)
  111. return categories
  112. def protected_output_conflicts(
  113. work_dir: str, scenario_id: str
  114. ) -> List[Tuple[str, str]]:
  115. """Return ``(category, subdir)`` for PROTECTED outputs the scenario writes
  116. that already exist and are non-empty. Empty result => safe to run without a
  117. destructive-overwrite prompt (only unprotected/regeneratable outputs, if any,
  118. would be replaced)."""
  119. conflicts: List[Tuple[str, str]] = []
  120. for category in scenario_output_categories(scenario_id):
  121. meta = ARTIFACT_CATEGORIES.get(category)
  122. if not meta or not meta["protected"]:
  123. continue
  124. for subdir in meta["dirs"]:
  125. path = os.path.join(work_dir, subdir)
  126. if os.path.isdir(path) and os.listdir(path):
  127. conflicts.append((category, subdir))
  128. return conflicts
  129. def describe_protected_conflicts(work_dir: str, scenario_id: str) -> List[str]:
  130. """Human-readable one-liners for each protected overwrite conflict."""
  131. lines: List[str] = []
  132. for category, subdir in protected_output_conflicts(work_dir, scenario_id):
  133. desc = ARTIFACT_CATEGORIES[category]["desc"]
  134. lines.append(f"{subdir}/ — {desc}")
  135. return lines
  136. SCENARIOS: Dict[str, ScenarioEntry] = {
  137. "scen_train_all": {
  138. "label": "1. Train, Evaluate, & Noise Analysis",
  139. "tasks": [
  140. "load_data",
  141. "train_regular",
  142. "train_bayesian",
  143. "load_models",
  144. "evaluate_regular",
  145. "evaluate_bayesian",
  146. "evaluate_noisy",
  147. ],
  148. },
  149. "scen_load_all": {
  150. "label": "2. Load, Evaluate, & Noise Analysis",
  151. "tasks": [
  152. "load_data",
  153. "load_models",
  154. "evaluate_regular",
  155. "evaluate_bayesian",
  156. "evaluate_noisy",
  157. ],
  158. },
  159. "scen_load_eval": {
  160. "label": "3. Load & Evaluate (Skip Noise)",
  161. "tasks": [
  162. "load_data",
  163. "load_models",
  164. "evaluate_regular",
  165. "evaluate_bayesian",
  166. ],
  167. },
  168. "scen_analyze": {
  169. "label": "4. Analyze Existing Evaluations",
  170. "tasks": [
  171. "run_analysis",
  172. ],
  173. },
  174. }