66import numpy as np
77import pyqtgraph as pg
88from PySide6 .QtCore import QPointF , QRectF , Qt
9- from PySide6 .QtGui import QBrush , QColor , QPainter , QPen , QPolygonF
9+ from PySide6 .QtGui import QBrush , QColor , QPainter , QPalette , QPen , QPolygonF
1010from PySide6 .QtWidgets import QVBoxLayout , QSizePolicy , QWidget
1111
12+ from interface .theme import DARK_THEME , current_theme
13+
14+
15+ def _apply_plot_theme (plot : pg .PlotWidget , title : str , x_label : str , y_label : str , theme : str ) -> None :
16+ """Apply the application theme to a pyqtgraph plot."""
17+ if theme == DARK_THEME :
18+ background , foreground , axis = "#141414" , "#eeeeee" , "#6a6a6a"
19+ else :
20+ background , foreground , axis = "#ffffff" , "#202020" , "#8a8a8a"
21+ plot .setBackground (background )
22+ plot .setTitle (title , color = foreground , size = "10pt" )
23+ plot .setLabel ("bottom" , x_label , color = foreground )
24+ plot .setLabel ("left" , y_label , color = foreground )
25+ plot .getAxis ("bottom" ).setPen (pg .mkPen (axis ))
26+ plot .getAxis ("left" ).setPen (pg .mkPen (axis ))
27+ plot .getAxis ("bottom" ).setTextPen (pg .mkPen (foreground ))
28+ plot .getAxis ("left" ).setTextPen (pg .mkPen (foreground ))
29+
1230
1331class ModelFlowWidget (QWidget ):
1432 """Transformer flow preview for live training telemetry."""
@@ -25,6 +43,16 @@ def __init__(self) -> None:
2543 self .step = 0
2644 self .loss_value : Optional [float ] = None
2745
46+ def apply_theme (self , _theme : str ) -> None :
47+ """Repaint the custom canvas after an application theme change."""
48+ self .update ()
49+
50+ def _foreground_color (self ) -> QColor :
51+ """Return a readable text color for the selected application theme."""
52+ if current_theme () == DARK_THEME :
53+ return QColor ("#e6e6e6" )
54+ return self .palette ().color (QPalette .WindowText )
55+
2856 def set_state (self , layer_count : int , head_count : int , step : int , loss_value : Optional [float ]) -> None :
2957 """Update flow state from UI-thread training metrics.
3058
@@ -53,7 +81,8 @@ def paintEvent(self, event: Any) -> None:
5381 painter .setRenderHint (QPainter .Antialiasing )
5482 try :
5583 rect = self .rect ().adjusted (18 , 18 , - 18 , - 18 )
56- painter .fillRect (self .rect (), QColor ("#111111" ))
84+ background = QColor ("#111111" ) if current_theme () == DARK_THEME else self .palette ().color (QPalette .Window )
85+ painter .fillRect (self .rect (), background )
5786 self ._draw_grid (painter , rect )
5887 layer_total = min (max (self .layer_count , 4 ), 10 )
5988 node_total = min (max (self .head_count , 5 ), 10 )
@@ -92,7 +121,7 @@ def _draw_layer_box(self, painter: QPainter, x_value: float, top: float, bottom:
92121 painter .setPen (QPen (QColor ("#4a90e2" ), 1 ))
93122 painter .setBrush (QBrush (QColor (35 , 80 , 120 , 55 )))
94123 painter .drawRoundedRect (box , 6 , 6 )
95- painter .setPen (QPen (QColor ( "#e6e6e6" )))
124+ painter .setPen (QPen (self . _foreground_color ( )))
96125 painter .drawText (QRectF (x_value - 42 , top - 54 , 84 , 22 ), Qt .AlignCenter , label )
97126
98127 def _draw_node (self , painter : QPainter , point : QPointF , column : int , active : bool ) -> None :
@@ -179,7 +208,7 @@ def _draw_flow_arrows(self, painter: QPainter, rect: Any) -> None:
179208 painter .drawLine (rect .right () - 34 , backward_y , rect .left () + 34 , backward_y )
180209 painter .drawLine (rect .left () + 42 , backward_y - 8 , rect .left () + 34 , backward_y )
181210 painter .drawLine (rect .left () + 42 , backward_y + 8 , rect .left () + 34 , backward_y )
182- painter .setPen (QPen (QColor ( "#e6e6e6" )))
211+ painter .setPen (QPen (self . _foreground_color ( )))
183212 painter .drawText (rect .left () + 44 , rect .top () + 20 , "MODEL FLOW (LIVE)" )
184213 painter .setPen (QPen (QColor ("#3ed7ff" )))
185214 painter .drawText (rect .left () + 220 , rect .top () + 20 , "Forward pass" )
@@ -193,12 +222,12 @@ def _draw_output_panel(self, painter: QPainter, rect: Any, points: list[QPointF]
193222 painter .setPen (QPen (QColor ("#4f7b48" ), 1 ))
194223 painter .setBrush (QBrush (QColor (25 , 55 , 30 , 160 )))
195224 painter .drawRoundedRect (panel , 5 , 5 )
196- painter .setPen (QPen (QColor ( "#d7d7d7" )))
225+ painter .setPen (QPen (self . _foreground_color ( )))
197226 painter .drawText (QRectF (panel .left (), panel .top () - 44 , panel .width (), 32 ), Qt .AlignCenter , "OUTPUT\n (NEXT TOKEN)" )
198227 samples = [("token" , 0.64 ), ("code" , 0.14 ), ("text" , 0.09 ), ("data" , 0.06 ), ("..." , 0.03 )]
199228 for index , (token , score ) in enumerate (samples ):
200229 y_value = panel .top () + 18 + index * 19
201- painter .setPen (QPen (QColor ("#b6d77a" ) if index == 0 else QColor ( "#d7d7d7" )))
230+ painter .setPen (QPen (QColor ("#b6d77a" ) if index == 0 else self . _foreground_color ( )))
202231 painter .drawText (QRectF (panel .left () + 12 , y_value - 9 , 54 , 18 ), Qt .AlignLeft | Qt .AlignVCenter , token )
203232 painter .drawText (QRectF (panel .right () - 42 , y_value - 9 , 34 , 18 ), Qt .AlignRight | Qt .AlignVCenter , f"{ score :.2f} " )
204233 painter .setPen (QPen (QColor ("#ffb13b" ), 1 ))
@@ -218,10 +247,7 @@ def __init__(self) -> None:
218247 layout = QVBoxLayout (self )
219248 layout .setContentsMargins (0 , 0 , 0 , 0 )
220249 self .plot = pg .PlotWidget ()
221- self .plot .setBackground ("#141414" )
222- self .plot .setTitle ("Prediction distribution" , color = "#eeeeee" , size = "10pt" )
223- self .plot .setLabel ("bottom" , "Vocabulary index" , color = "#d7d7d7" )
224- self .plot .setLabel ("left" , "Probability" , color = "#d7d7d7" )
250+ self .apply_theme (current_theme ())
225251 self .plot .showGrid (x = True , y = True , alpha = 0.22 )
226252 self .plot .setMenuEnabled (False )
227253 self .plot .setMouseEnabled (x = False , y = False )
@@ -230,6 +256,10 @@ def __init__(self) -> None:
230256 self .plot .setYRange (0 , 1.0 )
231257 layout .addWidget (self .plot )
232258
259+ def apply_theme (self , theme : str ) -> None :
260+ """Refresh this custom plot for the selected application theme."""
261+ _apply_plot_theme (self .plot , "Prediction distribution" , "Vocabulary index" , "Probability" , theme )
262+
233263 def update_distribution (self , step : int , loss_value : Optional [float ]) -> None :
234264 """Update synthetic distribution from live training state.
235265
@@ -259,10 +289,7 @@ def __init__(self) -> None:
259289 layout = QVBoxLayout (self )
260290 layout .setContentsMargins (0 , 0 , 0 , 0 )
261291 self .plot = pg .PlotWidget ()
262- self .plot .setBackground ("#141414" )
263- self .plot .setTitle ("Attention" , color = "#eeeeee" , size = "10pt" )
264- self .plot .setLabel ("bottom" , "Key token" , color = "#d7d7d7" )
265- self .plot .setLabel ("left" , "Query token" , color = "#d7d7d7" )
292+ self .apply_theme (current_theme ())
266293 self .plot .setMenuEnabled (False )
267294 self .plot .setMouseEnabled (x = False , y = False )
268295 self .image = pg .ImageItem ()
@@ -272,6 +299,10 @@ def __init__(self) -> None:
272299 layout .addWidget (self .plot )
273300 self .update_heatmap (0 , None )
274301
302+ def apply_theme (self , theme : str ) -> None :
303+ """Refresh this custom plot for the selected application theme."""
304+ _apply_plot_theme (self .plot , "Attention" , "Key token" , "Query token" , theme )
305+
275306 def update_heatmap (self , step : int , grad_norm : Optional [float ]) -> None :
276307 """Update attention proxy heatmap.
277308
@@ -301,10 +332,7 @@ def __init__(self) -> None:
301332 layout = QVBoxLayout (self )
302333 layout .setContentsMargins (0 , 0 , 0 , 0 )
303334 self .plot = pg .PlotWidget ()
304- self .plot .setBackground ("#141414" )
305- self .plot .setTitle ("Activation distribution" , color = "#eeeeee" , size = "10pt" )
306- self .plot .setLabel ("bottom" , "Activation" , color = "#d7d7d7" )
307- self .plot .setLabel ("left" , "Density" , color = "#d7d7d7" )
335+ self .apply_theme (current_theme ())
308336 self .plot .showGrid (x = True , y = True , alpha = 0.22 )
309337 self .plot .setMenuEnabled (False )
310338 self .plot .setMouseEnabled (x = False , y = False )
@@ -314,6 +342,10 @@ def __init__(self) -> None:
314342 self .plot .setYRange (0 , 1.0 )
315343 layout .addWidget (self .plot )
316344
345+ def apply_theme (self , theme : str ) -> None :
346+ """Refresh this custom plot for the selected application theme."""
347+ _apply_plot_theme (self .plot , "Activation distribution" , "Activation" , "Density" , theme )
348+
317349 def update_histogram (self , step : int , tokens_per_second : Optional [float ]) -> None :
318350 """Update activation histogram proxy.
319351
@@ -340,10 +372,7 @@ def __init__(self) -> None:
340372 layout = QVBoxLayout (self )
341373 layout .setContentsMargins (0 , 0 , 0 , 0 )
342374 self .plot = pg .PlotWidget ()
343- self .plot .setBackground ("#141414" )
344- self .plot .setTitle ("Gradient flow" , color = "#eeeeee" , size = "10pt" )
345- self .plot .setLabel ("bottom" , "Gradient norm" , color = "#d7d7d7" )
346- self .plot .setLabel ("left" , "Layer" , color = "#d7d7d7" )
375+ self .apply_theme (current_theme ())
347376 self .plot .showGrid (x = True , y = True , alpha = 0.22 )
348377 self .plot .setMenuEnabled (False )
349378 self .plot .setMouseEnabled (x = False , y = False )
@@ -354,6 +383,10 @@ def __init__(self) -> None:
354383 self .plot .setXRange (0 , 1 )
355384 layout .addWidget (self .plot )
356385
386+ def apply_theme (self , theme : str ) -> None :
387+ """Refresh this custom plot for the selected application theme."""
388+ _apply_plot_theme (self .plot , "Gradient flow" , "Gradient norm" , "Layer" , theme )
389+
357390 def update_flow (self , layer_count : int , grad_norm : Optional [float ], step : int ) -> None :
358391 """Update gradient flow proxy bars.
359392
@@ -374,4 +407,3 @@ def update_flow(self, layer_count: int, grad_norm: Optional[float], step: int) -
374407 self .bars .setOpts (x0 = [0 ] * count , x1 = values , y = self .layers , height = 0.55 , brush = "#ff4ca8" )
375408 self .plot .setYRange (0 , count + 1 )
376409 self .plot .setXRange (0 , max_value * 1.15 )
377-
0 commit comments