Skip to main content

burn_train/renderer/tui/
renderer.rs

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
34/// The current terminal backend.
35pub(crate) type TerminalBackend = CrosstermBackend<Stdout>;
36/// The current terminal frame.
37pub(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        /// Interrupter reset.
51        reset: bool,
52    },
53    ManualClose,
54    Close,
55    Persistent,
56}
57
58/// The terminal UI metrics renderer.
59pub 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    /// Create a new terminal UI renderer.
68    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                    // Render
91                    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    /// Set the renderer to persistent mode.
125    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        // Update the summary
167        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        // Reset for following steps.
221        self.interrupter.reset();
222        // Update the summary
223        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        // Reset the terminal to raw mode on panic before running the panic handler
279        // This prevents that the panic message is not visible for the user.
280        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                    (_, _) => { /* nothing to update */ }
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    // Reset the terminal back to raw mode.
476    fn reset(&mut self) -> Result<(), Box<dyn Error>> {
477        // If previous panic hook has already been re-instated, then the terminal was already reset.
478        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            // Reinstall the previous panic hook
490            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        // Reset the terminal back to raw mode. This can be skipped during
528        // panicking because the panic hook has already reset the terminal
529        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}