1use crate::metric::{MetricDefinition, MetricId};
2use crate::renderer::tui::TuiSplit;
3use crate::renderer::{
4 EvaluationName, EvaluationProgress, MetricState, MetricsRenderer, MetricsRendererEvaluation,
5 ProgressType, TrainingProgress,
6};
7use crate::renderer::{MetricsRendererTraining, tui::NumericMetricsState};
8use crate::{Interrupter, LearnerSummary};
9use ratatui::{
10 Terminal,
11 crossterm::{
12 event::{self, Event, KeyCode},
13 execute,
14 terminal::{EnterAlternateScreen, LeaveAlternateScreen, disable_raw_mode, enable_raw_mode},
15 },
16 prelude::*,
17};
18use std::collections::HashMap;
19use std::panic::{set_hook, take_hook};
20use std::sync::mpsc::{Receiver, Sender};
21use std::sync::{Arc, Mutex, mpsc};
22use std::thread::JoinHandle;
23use std::{
24 error::Error,
25 io::{self, Stdout},
26 time::{Duration, Instant},
27};
28
29use super::{
30 Callback, CallbackFn, ControlsView, MetricsView, PopupState, ProgressBarState, StatusState,
31 TextMetricsState, TuiGroup, TuiTag,
32};
33
34pub(crate) type TerminalBackend = CrosstermBackend<Stdout>;
36pub(crate) type TerminalFrame<'a> = ratatui::Frame<'a>;
38
39type PanicHook = Box<dyn Fn(&std::panic::PanicHookInfo<'_>) + 'static + Sync + Send>;
40
41const MAX_REFRESH_RATE_MILLIS: u64 = 100;
42
43enum TuiRendererEvent {
44 MetricRegistration(MetricDefinition),
45 MetricsUpdate((TuiSplit, TuiGroup, MetricState)),
46 StatusUpdateTrain((TuiSplit, TrainingProgress, Vec<ProgressType>)),
47 StatusUpdateTest((EvaluationProgress, Vec<ProgressType>)),
48 ProcessEnd {
49 summary: Option<LearnerSummary>,
50 reset: bool,
52 },
53 ManualClose,
54 Close,
55 Persistent,
56}
57
58pub struct TuiMetricsRendererWrapper {
60 sender: mpsc::Sender<TuiRendererEvent>,
61 interrupter: Interrupter,
62 handle_join: Option<JoinHandle<()>>,
63 kill_signal: Arc<Mutex<Receiver<()>>>,
64}
65
66impl TuiMetricsRendererWrapper {
67 pub fn new(interrupter: Interrupter, checkpoint: Option<usize>) -> Self {
69 let (sender, receiver) = mpsc::channel();
70 let (kill_signal_sender, kill_signal_receiver) = mpsc::channel();
71
72 let interrupter_clone = interrupter.clone();
73 let handle_join = std::thread::Builder::new()
74 .name("train-renderer".into())
75 .spawn(move || {
76 let mut renderer =
77 TuiMetricsRenderer::new(interrupter_clone, checkpoint, kill_signal_sender);
78
79 let tick_rate = Duration::from_millis(MAX_REFRESH_RATE_MILLIS);
80 loop {
81 match receiver.try_recv() {
82 Ok(event) => renderer.handle_event(event),
83 Err(mpsc::TryRecvError::Empty) => (),
84 Err(mpsc::TryRecvError::Disconnected) => {
85 log::error!("Renderer thread disconnected.");
86 break;
87 }
88 }
89
90 if renderer.last_update.elapsed() >= tick_rate
92 && let Err(err) = renderer.render()
93 {
94 log::error!("Render error: {err}");
95 break;
96 }
97
98 if (renderer.manual_close && renderer.interrupter.should_stop())
99 || renderer.close
100 {
101 break;
102 }
103 }
104 })
105 .unwrap();
106
107 Self {
108 sender,
109 interrupter,
110 handle_join: Some(handle_join),
111 kill_signal: Arc::new(Mutex::new(kill_signal_receiver)),
112 }
113 }
114
115 fn send_event(&self, event: TuiRendererEvent) {
116 if self.kill_signal.lock().unwrap().try_recv().is_ok() {
117 panic!("Killing training from user input.")
118 }
119 if let Err(e) = self.sender.send(event) {
120 log::warn!("Failed to send TUI event: {e}");
121 }
122 }
123
124 pub fn persistent(self) -> Self {
126 self.send_event(TuiRendererEvent::Persistent);
127 self
128 }
129}
130
131struct TuiMetricsRenderer {
132 terminal: Terminal<TerminalBackend>,
133 last_update: std::time::Instant,
134 progress: ProgressBarState,
135 metric_definitions: HashMap<MetricId, MetricDefinition>,
136 metrics_numeric: NumericMetricsState,
137 metrics_text: TextMetricsState,
138 status: StatusState,
139 interrupter: Interrupter,
140 popup: PopupState,
141 previous_panic_hook: Option<Arc<PanicHook>>,
142 persistent: bool,
143 manual_close: bool,
144 close: bool,
145 summary: Option<LearnerSummary>,
146 kill_signal: Sender<()>,
147}
148
149impl MetricsRendererEvaluation for TuiMetricsRendererWrapper {
150 fn update_test(&mut self, name: EvaluationName, state: MetricState) {
151 self.send_event(TuiRendererEvent::MetricsUpdate((
152 TuiSplit::Test,
153 TuiGroup::Named(name.name),
154 state,
155 )));
156 }
157
158 fn render_test(&mut self, item: EvaluationProgress, progress_indicators: Vec<ProgressType>) {
159 self.send_event(TuiRendererEvent::StatusUpdateTest((
160 item,
161 progress_indicators,
162 )));
163 }
164
165 fn on_test_end(&mut self, summary: Option<LearnerSummary>) -> Result<(), Box<dyn Error>> {
166 self.send_event(TuiRendererEvent::ProcessEnd {
168 summary,
169 reset: false,
170 });
171 Ok(())
172 }
173}
174
175impl MetricsRenderer for TuiMetricsRendererWrapper {
176 fn manual_close(&mut self) {
177 self.send_event(TuiRendererEvent::ManualClose);
178 let _ = self.handle_join.take().unwrap().join();
179 }
180
181 fn register_metric(&mut self, definition: MetricDefinition) {
182 self.send_event(TuiRendererEvent::MetricRegistration(definition));
183 }
184}
185
186impl MetricsRendererTraining for TuiMetricsRendererWrapper {
187 fn update_train(&mut self, state: MetricState) {
188 self.send_event(TuiRendererEvent::MetricsUpdate((
189 TuiSplit::Train,
190 TuiGroup::Default,
191 state,
192 )));
193 }
194
195 fn update_valid(&mut self, state: MetricState) {
196 self.send_event(TuiRendererEvent::MetricsUpdate((
197 TuiSplit::Valid,
198 TuiGroup::Default,
199 state,
200 )));
201 }
202
203 fn render_train(&mut self, item: TrainingProgress, progress_indicators: Vec<ProgressType>) {
204 self.send_event(TuiRendererEvent::StatusUpdateTrain((
205 TuiSplit::Train,
206 item,
207 progress_indicators,
208 )));
209 }
210
211 fn render_valid(&mut self, item: TrainingProgress, progress_indicators: Vec<ProgressType>) {
212 self.send_event(TuiRendererEvent::StatusUpdateTrain((
213 TuiSplit::Valid,
214 item,
215 progress_indicators,
216 )));
217 }
218
219 fn on_train_end(&mut self, summary: Option<LearnerSummary>) -> Result<(), Box<dyn Error>> {
220 self.interrupter.reset();
222 self.send_event(TuiRendererEvent::ProcessEnd {
224 summary,
225 reset: true,
226 });
227 Ok(())
228 }
229}
230
231impl Drop for TuiMetricsRendererWrapper {
232 fn drop(&mut self) {
233 if !std::thread::panicking() {
234 self.send_event(TuiRendererEvent::Close);
235 let _ = self.handle_join.take().unwrap().join();
236 }
237 }
238}
239
240impl TuiMetricsRenderer {
241 fn update_metric(&mut self, split: TuiSplit, group: TuiGroup, state: MetricState) {
242 match state {
243 MetricState::Generic(entry) => {
244 let name = self
245 .metric_definitions
246 .get(&entry.metric_id)
247 .unwrap()
248 .name
249 .clone()
250 .into();
251 self.metrics_text.update(split, group, entry, name);
252 }
253 MetricState::Numeric(entry, value) => {
254 let name: Arc<String> = self
255 .metric_definitions
256 .get(&entry.metric_id)
257 .unwrap()
258 .name
259 .clone()
260 .into();
261 self.metrics_numeric
262 .push(TuiTag::new(split, group.clone()), name.clone(), value);
263 self.metrics_text.update(split, group, entry, name);
264 }
265 };
266 }
267
268 pub fn new(
269 interrupter: Interrupter,
270 checkpoint: Option<usize>,
271 kill_signal: Sender<()>,
272 ) -> Self {
273 let mut stdout = io::stdout();
274 execute!(stdout, EnterAlternateScreen).unwrap();
275 enable_raw_mode().unwrap();
276 let terminal = Terminal::new(CrosstermBackend::new(stdout)).unwrap();
277
278 let previous_panic_hook = Arc::new(take_hook());
281 set_hook(Box::new({
282 let previous_panic_hook = previous_panic_hook.clone();
283 move |panic_info| {
284 let _ = disable_raw_mode();
285 let _ = execute!(io::stdout(), LeaveAlternateScreen);
286 previous_panic_hook(panic_info);
287 }
288 }));
289
290 Self {
291 terminal,
292 last_update: Instant::now(),
293 progress: ProgressBarState::new(checkpoint),
294 metric_definitions: HashMap::default(),
295 metrics_numeric: NumericMetricsState::default(),
296 metrics_text: TextMetricsState::default(),
297 status: StatusState::default(),
298 interrupter,
299 popup: PopupState::Empty,
300 previous_panic_hook: Some(previous_panic_hook),
301 persistent: false,
302 manual_close: false,
303 close: false,
304 summary: None,
305 kill_signal,
306 }
307 }
308
309 fn handle_event(&mut self, event: TuiRendererEvent) {
310 match event {
311 TuiRendererEvent::MetricRegistration(definition) => {
312 self.metric_definitions
313 .insert(definition.metric_id.clone(), definition);
314 }
315 TuiRendererEvent::MetricsUpdate((split, group, state)) => {
316 self.update_metric(split, group, state);
317 }
318 TuiRendererEvent::StatusUpdateTrain((split, item, status)) => match split {
319 TuiSplit::Train => {
320 self.progress.update_train(&item);
321 self.metrics_numeric.update_progress_train(&item);
322 self.status.update_train(status);
323 }
324 TuiSplit::Valid => {
325 self.progress.update_valid(&item);
326 self.metrics_numeric.update_progress_valid(&item);
327 self.status.update_valid(status);
328 }
329 _ => (),
330 },
331 TuiRendererEvent::StatusUpdateTest((item, status)) => {
332 self.progress.update_test(&item);
333 self.metrics_numeric.update_progress_test(&item);
334 self.status.update_test(status);
335 }
336 TuiRendererEvent::ProcessEnd { summary, reset } => {
337 match (self.summary.take(), summary) {
338 (None, Some(summary)) => {
339 self.summary = Some(summary);
340 }
341 (Some(current), Some(other)) => self.summary = Some(current.merge(other)),
342 (_, _) => { }
343 }
344
345 if reset {
346 self.interrupter.reset();
347 }
348 }
349 TuiRendererEvent::ManualClose => self.manual_close = true,
350 TuiRendererEvent::Persistent => self.persistent = true,
351 TuiRendererEvent::Close => self.close = true,
352 }
353 }
354
355 fn render(&mut self) -> Result<(), Box<dyn Error>> {
356 self.draw()?;
357 self.handle_user_input()?;
358
359 self.last_update = Instant::now();
360
361 Ok(())
362 }
363
364 fn draw(&mut self) -> Result<(), Box<dyn Error>> {
365 self.terminal.draw(|frame| {
366 let size = frame.area();
367
368 match self.popup.view() {
369 Some(view) => view.render(frame, size),
370 None => {
371 let view = MetricsView::new(
372 self.metrics_numeric.view(),
373 self.metrics_text.view(),
374 self.progress.view(),
375 ControlsView,
376 self.status.view(),
377 );
378
379 view.render(frame, size);
380 }
381 };
382 })?;
383
384 Ok(())
385 }
386
387 fn handle_user_input(&mut self) -> Result<(), Box<dyn Error>> {
388 while event::poll(Duration::from_secs(0))? {
389 let event = event::read()?;
390 self.popup.on_event(&event);
391
392 if self.popup.is_empty() {
393 self.metrics_numeric.on_event(&event);
394
395 if let Event::Key(key) = event
396 && let KeyCode::Char('q') = key.code
397 {
398 self.popup = PopupState::Full(
399 "Quit".to_string(),
400 vec![
401 Callback::new(
402 "Stop the training.",
403 "Stop the training immediately. This will break from the \
404 training loop, but any remaining code after the loop will be \
405 executed.",
406 's',
407 QuitPopupAccept(self.interrupter.clone()),
408 ),
409 Callback::new(
410 "Stop the training immediately.",
411 "Kill the program. This will create a panic! which will make \
412 the current training fails. Any code following the training \
413 won't be executed.",
414 'k',
415 KillPopupAccept(self.kill_signal.clone()),
416 ),
417 Callback::new(
418 "Cancel",
419 "Cancel the action, continue the training.",
420 'c',
421 PopupCancel,
422 ),
423 ],
424 );
425 }
426 }
427 }
428
429 Ok(())
430 }
431
432 fn handle_post_training(&mut self) -> Result<(), Box<dyn Error>> {
433 self.popup = PopupState::Full(
434 "Training is done".to_string(),
435 vec![Callback::new(
436 "Training Done",
437 "Press 'x' to close this popup. Press 'q' to exit the application after the \
438 popup is closed.",
439 'x',
440 PopupCancel,
441 )],
442 );
443
444 self.draw().ok();
445
446 loop {
447 if let Ok(true) = event::poll(Duration::from_millis(MAX_REFRESH_RATE_MILLIS)) {
448 match event::read() {
449 Ok(event @ Event::Key(key)) => {
450 if self.popup.is_empty() {
451 self.metrics_numeric.on_event(&event);
452 if let KeyCode::Char('q') = key.code {
453 break;
454 }
455 } else {
456 self.popup.on_event(&event);
457 }
458 self.draw().ok();
459 }
460
461 Ok(Event::Resize(..)) => {
462 self.draw().ok();
463 }
464 Err(err) => {
465 eprintln!("Error reading event: {err}");
466 break;
467 }
468 _ => continue,
469 }
470 }
471 }
472 Ok(())
473 }
474
475 fn reset(&mut self) -> Result<(), Box<dyn Error>> {
477 if self.previous_panic_hook.is_some() {
479 if self.persistent
480 && let Err(err) = self.handle_post_training()
481 {
482 eprintln!("Error in post-training handling: {err}");
483 }
484
485 disable_raw_mode()?;
486 execute!(self.terminal.backend_mut(), LeaveAlternateScreen)?;
487 self.terminal.show_cursor()?;
488
489 let _ = take_hook();
491 if let Some(previous_panic_hook) =
492 Arc::into_inner(self.previous_panic_hook.take().unwrap())
493 {
494 set_hook(previous_panic_hook);
495 }
496 }
497 Ok(())
498 }
499}
500
501struct QuitPopupAccept(Interrupter);
502struct KillPopupAccept(Sender<()>);
503struct PopupCancel;
504
505impl CallbackFn for KillPopupAccept {
506 fn call(&self) -> bool {
507 self.0.send(()).unwrap();
508 panic!("Killing training from user input.");
509 }
510}
511
512impl CallbackFn for QuitPopupAccept {
513 fn call(&self) -> bool {
514 self.0.stop(Some("Stopping training from user input."));
515 true
516 }
517}
518
519impl CallbackFn for PopupCancel {
520 fn call(&self) -> bool {
521 true
522 }
523}
524
525impl Drop for TuiMetricsRenderer {
526 fn drop(&mut self) {
527 if !std::thread::panicking() {
530 self.reset().unwrap();
531
532 if let Some(summary) = &self.summary {
533 println!("{summary}");
534 log::info!("{summary}");
535 }
536 }
537 }
538}