main.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305
  1. import os
  2. import threading
  3. import gc
  4. from typing import Any, Dict
  5. # Must be set before first CUDA context initialization.
  6. os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
  7. import toml
  8. from textual import work
  9. from textual.app import App, ComposeResult
  10. from textual.containers import Horizontal, Vertical
  11. from textual.widgets import (
  12. Button,
  13. Footer,
  14. Header,
  15. Input,
  16. Label,
  17. ProgressBar,
  18. RichLog,
  19. Rule,
  20. )
  21. from tasks import PIPELINE_TASKS, SCENARIOS
  22. from util.config_manager import handle_directory_config
  23. from util.progress import ProgressTracker
  24. from util.screens import OverwriteConfirmScreen
  25. from util.ui_logger import PipelineLogger
  26. class CNNHarnessApp(App):
  27. CSS_PATH = "style.tcss"
  28. def __init__(self):
  29. super().__init__()
  30. self.stop_event = threading.Event()
  31. self.pause_event = threading.Event()
  32. self.current_config: Dict[str, Any] = {}
  33. self.ui_logger = PipelineLogger(self, widget_id="#live_log")
  34. # ==========================================
  35. # UI Layout
  36. # ==========================================
  37. def compose(self) -> ComposeResult:
  38. yield Header()
  39. with Vertical(id="main_container"):
  40. # --- TOP SECTION: Two-Column Configuration ---
  41. with Horizontal(id="config_section"):
  42. # Left Column: Path entry and buttons
  43. with Vertical(id="config_left_col"):
  44. yield Label("Working Directory:")
  45. yield Input(placeholder="./outputs/experiment_1", id="work_dir")
  46. yield Button("Load / Init Dir", variant="primary", id="load_btn")
  47. yield Button("Reload Config", variant="default", id="reload_btn")
  48. # Right Column: Read-only parameters display
  49. with Vertical(id="config_right_col"):
  50. yield Label("Current Parameters:")
  51. yield RichLog(id="config_display", highlight=True, markup=True)
  52. yield Rule()
  53. # --- BOTTOM SECTION: Controls, Progress, and Logs ---
  54. with Horizontal(id="bottom_area"):
  55. # --- Pipeline Controls ---
  56. with Vertical(id="control_panel"):
  57. yield Button("Start Pipeline", variant="success", id="btn_start")
  58. with Horizontal(id="active_controls"):
  59. yield Button("Pause", variant="warning", id="btn_pause")
  60. yield Button("Stop", variant="error", id="btn_stop")
  61. # --- Progress Trackers ---
  62. with Vertical(id="progress_area"):
  63. for i in range(5):
  64. with Horizontal(
  65. id=f"progress_row_{i}", classes="progress_row"
  66. ):
  67. yield Label(
  68. id=f"progress_title_{i}", classes="progress_title"
  69. )
  70. yield ProgressBar(
  71. id=f"progress_bar_{i}",
  72. classes="progress_bar",
  73. show_eta=True,
  74. )
  75. yield Label(
  76. id=f"progress_stats_{i}", classes="progress_stats"
  77. )
  78. # --- Live Log Output ---
  79. with Vertical(id="log_panel"):
  80. with Horizontal(id="log_status_row"):
  81. yield Label("STATUS: READY", id="pipeline_status_label")
  82. with Horizontal(id="log_header_row"):
  83. yield Label("Live Logs:", classes="section_label")
  84. yield RichLog(id="live_log", highlight=True, markup=True)
  85. yield Footer()
  86. def on_mount(self) -> None:
  87. self.query_one("#active_controls").display = False
  88. self.ui_logger.info("Application initialized and ready.")
  89. # ==========================================
  90. # Helper Methods
  91. # ==========================================
  92. def _set_pipeline_status(self, text: str) -> None:
  93. try:
  94. self.query_one("#pipeline_status_label", Label).update(text)
  95. except Exception:
  96. pass
  97. def _update_config_display(self) -> None:
  98. config_log = self.query_one("#config_display", RichLog)
  99. config_log.clear()
  100. if not self.current_config:
  101. config_log.write("[italic]No configuration loaded.[/italic]")
  102. return
  103. toml_string = toml.dumps(self.current_config)
  104. config_log.write(toml_string)
  105. def _toggle_config_inputs(self, disabled: bool) -> None:
  106. self.query_one("#work_dir", Input).disabled = disabled
  107. self.query_one("#load_btn", Button).disabled = disabled
  108. self.query_one("#reload_btn", Button).disabled = disabled
  109. # ==========================================
  110. # Event Handlers
  111. # ==========================================
  112. def on_button_pressed(self, event: Button.Pressed) -> None:
  113. if event.button.id in ("load_btn", "reload_btn"):
  114. work_dir = self.query_one("#work_dir", Input).value
  115. success, message, config_data = handle_directory_config(work_dir)
  116. if success:
  117. self.current_config = config_data
  118. self.current_config["work_dir"] = work_dir
  119. self._update_config_display()
  120. self.ui_logger.info(message)
  121. else:
  122. self.ui_logger.error(message)
  123. elif event.button.id == "btn_start":
  124. if not self.current_config:
  125. self.ui_logger.error(
  126. "No configuration loaded! Please load a directory first."
  127. )
  128. return
  129. scenario_id = self.current_config.get("scenario")
  130. if not scenario_id or scenario_id not in SCENARIOS:
  131. self.ui_logger.error(
  132. f"Invalid or missing scenario '{scenario_id}' in config.toml!"
  133. )
  134. return
  135. work_dir = self.current_config["work_dir"]
  136. # Check if the working directory exists and is non-empty (except for the config.toml file)
  137. if os.path.exists(work_dir) and any(
  138. f for f in os.listdir(work_dir) if f != "config.toml"
  139. ):
  140. def check_overwrite_callback(proceed: bool) -> None:
  141. if proceed:
  142. self._execute_pipeline()
  143. else:
  144. self.ui_logger.error(
  145. "Pipeline start cancelled by user (folder non-empty)."
  146. )
  147. self.app.push_screen(
  148. OverwriteConfirmScreen(work_dir),
  149. check_overwrite_callback, # pyright: ignore
  150. )
  151. else:
  152. self._execute_pipeline()
  153. elif event.button.id == "btn_pause":
  154. if self.pause_event.is_set():
  155. self.pause_event.clear()
  156. event.button.label = "Pause"
  157. event.button.variant = "warning"
  158. self._set_pipeline_status("[bold green]PIPELINE RUNNING[/bold green]")
  159. self.ui_logger.info("Pipeline Resumed.")
  160. else:
  161. self.pause_event.set()
  162. event.button.label = "Resume"
  163. event.button.variant = "success"
  164. self._set_pipeline_status("[bold yellow]PIPELINE PAUSED[/bold yellow]")
  165. self.ui_logger.info(
  166. "Pipeline Paused. Waiting for current operation to yield..."
  167. )
  168. elif event.button.id == "btn_stop":
  169. self._set_pipeline_status("[bold red]STOPPING PIPELINE...[/bold red]")
  170. self.ui_logger.error("Stop requested! Terminating gracefully...")
  171. event.button.disabled = True
  172. self.stop_event.set()
  173. # ==========================================
  174. # Pipeline Execution
  175. # ==========================================
  176. def _execute_pipeline(self) -> None:
  177. self._toggle_config_inputs(disabled=True)
  178. self.query_one("#btn_start").display = False
  179. self.query_one("#active_controls").display = True
  180. self._set_pipeline_status("[bold green]PIPELINE RUNNING[/bold green]")
  181. self.stop_event.clear()
  182. self.pause_event.clear()
  183. btn_pause = self.query_one("#btn_pause", Button)
  184. btn_pause.label = "Pause"
  185. btn_pause.variant = "warning"
  186. self.run_background_pipeline(self.current_config.copy())
  187. @work(exclusive=True, thread=True)
  188. def run_background_pipeline(self, config: Dict[str, Any]) -> None:
  189. scenario_id = config["scenario"]
  190. steps = SCENARIOS[scenario_id]["tasks"]
  191. scenario_label = SCENARIOS[scenario_id]["label"]
  192. root_tracker = ProgressTracker(self, level=0)
  193. pipeline_state: Dict[str, Any] = {
  194. "events": {"stop": self.stop_event, "pause": self.pause_event}
  195. }
  196. try:
  197. self.ui_logger.info(f"Initiating Scenario: {scenario_label}")
  198. total_tasks = len(steps)
  199. root_tracker.reset()
  200. root_tracker.set_title("Pipeline Status")
  201. root_tracker.update(total=total_tasks, advance=0)
  202. for i, task_id in enumerate(steps):
  203. if self.stop_event.is_set():
  204. break
  205. task_name = PIPELINE_TASKS[task_id]["task_name"]
  206. task_func = PIPELINE_TASKS[task_id]["task_func"]
  207. task_logger = self.ui_logger.get_task_logger(task_name)
  208. self.ui_logger.info(f"Starting Task {i + 1}/{total_tasks}: {task_name}")
  209. task_tracker = root_tracker.get_sub_tracker(task_name)
  210. task_func(task_tracker, task_logger, config, pipeline_state)
  211. if not self.stop_event.is_set():
  212. task_logger.info("Task completed.")
  213. root_tracker.update(advance=1)
  214. if self.stop_event.is_set():
  215. self.call_from_thread(
  216. self._set_pipeline_status, "[bold red]PIPELINE STOPPED[/bold red]"
  217. )
  218. self.ui_logger.error("Pipeline stopped by user.")
  219. root_tracker.reset()
  220. else:
  221. self.call_from_thread(
  222. self._set_pipeline_status,
  223. "[bold blue]PIPELINE COMPLETE[/bold blue]",
  224. )
  225. self.ui_logger.info("All tasks finished successfully!")
  226. except InterruptedError as e:
  227. self.ui_logger.error(f"STOPPED: {str(e)}")
  228. self.call_from_thread(
  229. self._set_pipeline_status, "[bold red]PIPELINE STOPPED[/bold red]"
  230. )
  231. root_tracker.reset()
  232. except Exception as e:
  233. self.ui_logger.file_logger.error(
  234. f"Pipeline failed: {str(e)}", exc_info=True
  235. )
  236. self.ui_logger.error(f"ERROR: {str(e)}")
  237. self.call_from_thread(
  238. self._set_pipeline_status, "[bold red]PIPELINE FAILED[/bold red]"
  239. )
  240. finally:
  241. pipeline_state.clear()
  242. gc.collect()
  243. self.call_from_thread(self._reset_ui)
  244. def _reset_ui(self) -> None:
  245. self.query_one("#btn_stop", Button).disabled = False
  246. self.query_one("#btn_start").display = True
  247. self.query_one("#active_controls").display = False
  248. self._toggle_config_inputs(disabled=False)
  249. if __name__ == "__main__":
  250. app = CNNHarnessApp()
  251. app.run()