Skip to content

Commit 8d59955

Browse files
baremetalgoCopilot
andcommitted
Keep theme surfaces coherent
Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
1 parent 782be1e commit 8d59955

9 files changed

Lines changed: 199 additions & 64 deletions

File tree

interface/charts.py

Lines changed: 30 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,24 @@
88
from PySide6.QtGui import QCursor
99
from PySide6.QtWidgets import QSizePolicy, QToolTip, QVBoxLayout, QWidget
1010

11+
from interface.theme import DARK_THEME, current_theme
12+
13+
14+
def _apply_plot_theme(plot: pg.PlotWidget, title: str, x_label: str, y_label: str, theme: str) -> None:
15+
"""Apply colors to a pyqtgraph plot, which does not inherit Qt QSS."""
16+
if theme == DARK_THEME:
17+
background, foreground, axis = "#141414", "#eeeeee", "#6a6a6a"
18+
else:
19+
background, foreground, axis = "#ffffff", "#202020", "#8a8a8a"
20+
plot.setBackground(background)
21+
plot.setTitle(title, color=foreground, size="10pt")
22+
plot.setLabel("bottom", x_label, color=foreground)
23+
plot.setLabel("left", y_label, color=foreground)
24+
plot.getAxis("bottom").setPen(pg.mkPen(axis))
25+
plot.getAxis("left").setPen(pg.mkPen(axis))
26+
plot.getAxis("bottom").setTextPen(pg.mkPen(foreground))
27+
plot.getAxis("left").setTextPen(pg.mkPen(foreground))
28+
1129

1230
class DatasetBarChartWidget(QWidget):
1331
"""Compact bar chart for dataset composition and token statistics."""
@@ -33,22 +51,18 @@ def __init__(
3351
self.labels: list[str] = []
3452
self.values: list[float] = []
3553
self.value_suffix = ""
54+
self.title = title
55+
self.y_label = y_label
3656

3757
layout = QVBoxLayout(self)
3858
layout.setContentsMargins(0, 0, 0, 0)
3959
layout.setSpacing(0)
4060
pg.setConfigOptions(antialias=True)
4161
self.plot = pg.PlotWidget()
42-
self.plot.setBackground("#141414")
43-
self.plot.setTitle(title, color="#eeeeee", size="10pt")
44-
self.plot.setLabel("left", y_label, color="#d7d7d7")
62+
self.apply_theme(current_theme())
4563
self.plot.showGrid(x=False, y=True, alpha=0.24)
4664
self.plot.setMenuEnabled(False)
4765
self.plot.setMouseEnabled(x=False, y=False)
48-
self.plot.getAxis("bottom").setPen(pg.mkPen("#6a6a6a"))
49-
self.plot.getAxis("left").setPen(pg.mkPen("#6a6a6a"))
50-
self.plot.getAxis("bottom").setTextPen(pg.mkPen("#cfcfcf"))
51-
self.plot.getAxis("left").setTextPen(pg.mkPen("#cfcfcf"))
5266
self.plot.getPlotItem().setContentsMargins(8, 8, 8, 8)
5367
self.bar_item = pg.BarGraphItem(x=[], height=[], width=0.58, brush=pg.mkBrush("#f5b041"))
5468
self.plot.addItem(self.bar_item)
@@ -58,6 +72,10 @@ def __init__(
5872
layout.addWidget(self.plot)
5973
self.clear()
6074

75+
def apply_theme(self, theme: str) -> None:
76+
"""Refresh this custom plot for the selected application theme."""
77+
_apply_plot_theme(self.plot, self.title, "", self.y_label, theme)
78+
6179
def clear(self) -> None:
6280
"""Clear chart values."""
6381

@@ -163,17 +181,10 @@ def __init__(
163181
layout.setSpacing(0)
164182
pg.setConfigOptions(antialias=True)
165183
self.plot = pg.PlotWidget()
166-
self.plot.setBackground("#141414")
167-
self.plot.setTitle(title, color="#eeeeee", size="11pt")
168-
self.plot.setLabel("bottom", "Optimizer step", color="#d7d7d7")
169-
self.plot.setLabel("left", y_label, color="#d7d7d7")
184+
self.apply_theme(current_theme())
170185
self.plot.showGrid(x=True, y=True, alpha=0.28)
171186
self.plot.setMenuEnabled(False)
172187
self.plot.setMouseEnabled(x=True, y=True)
173-
self.plot.getAxis("bottom").setPen(pg.mkPen("#6a6a6a"))
174-
self.plot.getAxis("left").setPen(pg.mkPen("#6a6a6a"))
175-
self.plot.getAxis("bottom").setTextPen(pg.mkPen("#cfcfcf"))
176-
self.plot.getAxis("left").setTextPen(pg.mkPen("#cfcfcf"))
177188
self.plot.getPlotItem().setContentsMargins(10, 8, 10, 8)
178189
self.legend = self.plot.addLegend(offset=(12, 8), brush=pg.mkBrush(20, 20, 20, 180), pen=pg.mkPen("#444444"))
179190
self.primary_curve = self.plot.plot([], [], pen=pg.mkPen("#f5b041", width=2), name=primary_label)
@@ -188,6 +199,10 @@ def __init__(
188199
layout.addWidget(self.plot)
189200
self._refresh_plot()
190201

202+
def apply_theme(self, theme: str) -> None:
203+
"""Refresh this custom plot for the selected application theme."""
204+
_apply_plot_theme(self.plot, self.title, "Optimizer step", self.y_label, theme)
205+
191206
def clear(self) -> None:
192207
"""Remove all plotted loss values."""
193208

@@ -397,4 +412,3 @@ def _nearest_point(self, x_value: float, y_value: float) -> Optional[tuple[str,
397412
)
398413
distance = ((nearest[1] - x_value) / x_span) ** 2 + ((nearest[2] - y_value) / y_span) ** 2
399414
return nearest if distance < 0.01 else None
400-

interface/core/project_state.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@ def _default_project_state(self) -> dict[str, Any]:
2424
"schema": "drunkenbot_ide_project",
2525
"version": 1,
2626
"theme": "dark",
27+
"theme_preference_version": 1,
2728
"project_name": "",
2829
"project_dir": "",
2930
"paths": {
@@ -256,6 +257,7 @@ def _project_state_dict(self, project_name: str, project_dir: Path) -> dict[str,
256257
"schema": "drunkenbot_ide_project",
257258
"version": 1,
258259
"theme": self.theme_name,
260+
"theme_preference_version": 1,
259261
"project_name": project_name,
260262
"project_dir": str(project_dir),
261263
"created_at": created_at,

interface/core/project_state_apply.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
# ProjectStateApplyMixin mixin. Shared runtime names are provided by interface.app.
44
from typing import Any, Optional, Union # noqa: F401
55
from interface import app as _app
6+
from interface.theme import project_theme
67

78
globals().update({name: value for name, value in vars(_app).items() if not name.startswith("__")})
89

@@ -15,7 +16,7 @@ def _apply_project_state(self, data: dict[str, Any]) -> None:
1516
data: Project state loaded from JSON.
1617
"""
1718

18-
self.set_theme(data.get("theme", "dark"), persist=False)
19+
self.set_theme(project_theme(data), persist=False)
1920
self.search_box.setText(str(data.get("project_name", "")))
2021
paths = data.get("paths", {})
2122
dataset = data.get("dataset", {})

interface/core/window_core.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22

33
# WindowCoreMixin mixin. Shared runtime names are provided by interface.app.
44
from typing import Any, Optional, Union # noqa: F401
5-
from interface.widgets.app_shell import build_main_shell
5+
from interface.widgets.app_shell import build_main_shell, update_navigation_icons
66
from interface.theme import apply_theme, current_theme, normalize_theme
77
from interface import app as _app
88

@@ -80,6 +80,7 @@ def __init__(self) -> None:
8080

8181
shell = self._build_shell()
8282
self.setCentralWidget(shell)
83+
update_navigation_icons(self)
8384
self._install_ui_event_logging(shell)
8485
self._install_wheel_guard(shell)
8586
self._refresh_notification_manager()
@@ -227,6 +228,8 @@ def set_theme(self, theme: object, persist: bool = True) -> None:
227228
"""
228229
self.theme_name = normalize_theme(theme)
229230
self._apply_theme()
231+
if hasattr(self, "side_rail"):
232+
update_navigation_icons(self)
230233
self.update_theme_actions()
231234
if persist and self.current_project_file is not None:
232235
self.save_project()

interface/live_widgets.py

Lines changed: 55 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -6,9 +6,27 @@
66
import numpy as np
77
import pyqtgraph as pg
88
from 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
1010
from 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

1331
class 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

Comments
 (0)