__init__.py 3.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119
  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 -> PLANNED (placeholder)
  10. 5. evaluate_bayesian-> PLANNED (placeholder)
  11. 6. evaluate_noisy -> PLANNED (placeholder)
  12. load_models -> PLANNED (placeholder); loads saved .pt ensembles from
  13. disk into ``state`` so evaluation is decoupled from
  14. training (training deletes models from VRAM after saving).
  15. """
  16. from typing import Any, Dict
  17. from util.progress import ProgressTracker
  18. from util.ui_logger import PipelineLogger
  19. from . import load_data
  20. from . import train_bayesian
  21. from . import train_normal
  22. def _placeholder_task(planned: str):
  23. """Build a no-op task for a pipeline stage that is not implemented yet.
  24. It logs a clear warning and returns cleanly (advancing the pipeline) instead
  25. of crashing, so the scenario wiring can be exercised before the real
  26. evaluation logic lands.
  27. """
  28. def _task(
  29. tracker: ProgressTracker,
  30. logger: PipelineLogger,
  31. config: Dict[str, Any],
  32. state: Dict[str, Any],
  33. ) -> None:
  34. logger.info(f"[NOT IMPLEMENTED] {planned} — skipping (placeholder).")
  35. return _task
  36. PIPELINE_TASKS = {
  37. "load_data": {
  38. "task_name": "Load Image and ADNIMERGE",
  39. "task_func": load_data.load_data_task,
  40. },
  41. "train_regular": {
  42. "task_name": "Train Regular Models",
  43. "task_func": train_normal.train_normal_task,
  44. },
  45. "train_bayesian": {
  46. "task_name": "Train Bayesian Models",
  47. "task_func": train_bayesian.train_bayesian_task,
  48. },
  49. "load_models": {
  50. "task_name": "Load Saved Models",
  51. "task_func": _placeholder_task(
  52. "Load trained normal/Bayesian ensembles from work_dir into state"
  53. ),
  54. },
  55. "evaluate_regular": {
  56. "task_name": "Evaluate Regular Models",
  57. "task_func": _placeholder_task(
  58. "Evaluate normal ensemble and save netCDF (step 4)"
  59. ),
  60. },
  61. "evaluate_bayesian": {
  62. "task_name": "Evaluate Bayesian Models",
  63. "task_func": _placeholder_task(
  64. "Evaluate Bayesian ensemble with MC uncertainty and save netCDF (step 5)"
  65. ),
  66. },
  67. "evaluate_noisy": {
  68. "task_name": "Evaluate Models on Noised Data",
  69. "task_func": _placeholder_task(
  70. "Evaluate both ensembles across Gaussian noise levels and save netCDF (step 6)"
  71. ),
  72. },
  73. }
  74. SCENARIOS = {
  75. "scen_train_all": {
  76. "label": "1. Train, Evaluate, & Noise Analysis",
  77. "tasks": [
  78. "load_data",
  79. "train_regular",
  80. "train_bayesian",
  81. "load_models",
  82. "evaluate_regular",
  83. "evaluate_bayesian",
  84. "evaluate_noisy",
  85. ],
  86. },
  87. "scen_load_all": {
  88. "label": "2. Load, Evaluate, & Noise Analysis",
  89. "tasks": [
  90. "load_data",
  91. "load_models",
  92. "evaluate_regular",
  93. "evaluate_bayesian",
  94. "evaluate_noisy",
  95. ],
  96. },
  97. "scen_load_eval": {
  98. "label": "3. Load & Evaluate (Skip Noise)",
  99. "tasks": [
  100. "load_data",
  101. "load_models",
  102. "evaluate_regular",
  103. "evaluate_bayesian",
  104. ],
  105. },
  106. }