BigUI WIP: Some Model Manager Cleanup

This commit is contained in:
firestarsdog
2026-04-11 04:54:16 -04:00
parent 161e0e10c8
commit 851666e244
5 changed files with 713 additions and 234 deletions
@@ -198,11 +198,12 @@ class AetherTile(Widget):
class HubTile(AetherTile):
def __init__(self, title: str, desc: str, icon_path: str, on_click: Callable | None = None, starpilot_icon: bool = False, bg_color: rl.Color | str | None = None):
def __init__(self, title: str, desc: str, icon_path: str, on_click: Callable | None = None, starpilot_icon: bool = False, bg_color: rl.Color | str | None = None, get_status: Callable[[], str] | None = None):
if bg_color:
super().__init__(surface_color=bg_color, on_click=on_click)
else:
super().__init__(on_click=on_click)
self.get_status = get_status
self.title = title
self.desc = desc
if icon_path:
@@ -214,6 +215,17 @@ class HubTile(AetherTile):
def _render(self, rect: rl.Rectangle):
face = self._render_layers(rect)
status_text = self.get_status() if self.get_status else ""
if status_text:
import re
m = re.search(r'(\d+)%$', status_text)
if m:
ratio = min(1.0, max(0.0, float(m.group(1)) / 100.0))
if ratio > 0.05:
fill_rect = rl.Rectangle(face.x, face.y, face.width * ratio, face.height)
rl.draw_rectangle_rounded(fill_rect, TILE_RADIUS, 10, rl.Color(255, 255, 255, 40))
content_pad = SPACING.tile_content
max_w = face.width - (content_pad * 2)
lines = self._wrap_text(self._font_title, self.title, max_w, 30)
@@ -224,8 +236,9 @@ class HubTile(AetherTile):
for i, line in enumerate(lines):
self._draw_text_fit(self._font_title, line, rl.Vector2(face.x + content_pad, ty + i * (line_h + line_spacing)), max_w, line_h, align_center=True)
if self.desc:
desc_lines = self._wrap_text(self._font_desc, self.desc, max_w, 18, max_lines=3)
desc_to_render = status_text if status_text else self.desc
if desc_to_render:
desc_lines = self._wrap_text(self._font_desc, desc_to_render, max_w, 18, max_lines=3)
desc_y = ty + len(lines) * (line_h + line_spacing) + SPACING.lg
for i, line in enumerate(desc_lines):
self._draw_text_fit(self._font_desc, line, rl.Vector2(face.x + content_pad, desc_y + i * 20), max_w, 18, align_center=True)
@@ -492,14 +505,14 @@ class AetherSliderDialog(Widget):
if self._is_pressed_ok:
self._ok_target = 0.0
if rl.check_collision_point_rec(mouse_pos, self._ok_rect):
self._user_callback(DialogResult.CONFIRM, self._current_val)
gui_app.set_modal_overlay(None)
self._user_callback(DialogResult.CONFIRM, self._current_val)
self._is_pressed_ok = False
if self._is_pressed_cancel:
self._cancel_target = 0.0
if rl.check_collision_point_rec(mouse_pos, self._cancel_rect):
self._user_callback(DialogResult.CANCEL, self._current_val)
gui_app.set_modal_overlay(None)
self._user_callback(DialogResult.CANCEL, self._current_val)
self._is_pressed_cancel = False
def _render(self, rect: rl.Rectangle):
@@ -635,26 +648,46 @@ class TileGrid(Widget):
self.tiles = []
self._uniform_width = uniform_width
@property
def gap(self) -> int:
return self._gap
def add_tile(self, tile: Widget): self.tiles.append(tile)
def clear(self): self.tiles.clear()
def get_column_count(self, tile_count: int | None = None) -> int:
count = len(self.tiles) if tile_count is None else tile_count
if count <= 0:
return self._columns or 1
if self._columns is not None:
return self._columns
if count == 1: return 1
if count == 2: return 2
if count == 3: return 3
if count == 4: return 2
if count <= 6: return 3
return 4
def get_row_count(self, tile_count: int | None = None) -> int:
count = len(self.tiles) if tile_count is None else tile_count
if count <= 0:
return 0
cols = self.get_column_count(count)
return (count + cols - 1) // cols
def get_internal_gap_height(self, tile_count: int | None = None) -> float:
rows = self.get_row_count(tile_count)
return self._gap * max(0, rows - 1)
def _render(self, rect: rl.Rectangle):
self.set_rect(rect)
if not self.tiles:
return
tiles_to_render = list(self.tiles)
count = len(tiles_to_render)
if self._columns is not None:
cols = self._columns
else:
if count == 1: cols = 1
elif count == 2: cols = 2
elif count == 3: cols = 3
elif count == 4: cols = 2
elif count <= 6: cols = 3
else: cols = 4
rows = (count + cols - 1) // cols
cols = self.get_column_count(count)
rows = self.get_row_count(count)
tile_h = (rect.height - (self._gap * (rows - 1))) / rows
uniform_tile_w = (rect.width - (self._gap * (cols - 1))) / cols if self._uniform_width else 0
tile_idx = 0
@@ -1,12 +1,14 @@
from __future__ import annotations
from dataclasses import dataclass
from collections.abc import Callable
import json
import os
import shutil
import threading
from pathlib import Path
from openpilot.selfdrive.ui.ui_state import ui_state
from openpilot.starpilot.assets.model_manager import ModelManager
from openpilot.system.hardware import HARDWARE, PC
from openpilot.starpilot.assets.model_manager import ModelManager, TINYGRAD_VERSIONS, canonical_model_key, is_builtin_model_key, model_key_aliases
from openpilot.starpilot.common.starpilot_variables import MODELS_PATH, update_starpilot_toggles
from openpilot.system.hardware import HARDWARE
from openpilot.system.ui.lib.application import gui_app
from openpilot.system.ui.lib.multilang import tr, tr_noop
from openpilot.system.ui.widgets import DialogResult
@@ -15,105 +17,158 @@ from openpilot.system.ui.widgets.selection_dialog import SelectionDialog
from openpilot.selfdrive.ui.layouts.settings.starpilot.panel import StarPilotPanel
from openpilot.selfdrive.ui.layouts.settings.starpilot.aethergrid import AetherSliderDialog
@dataclass
class ModelCatalogEntry:
key: str
name: str
series: str
version: str
released: str
builtin: bool
installed: bool
partial: bool
community_favorite: bool
user_favorite: bool
def _clean_model_name(name: str) -> str:
return str(name or "").replace("_default", "").replace("(Default)", "").strip()
class StarPilotDrivingModelLayout(StarPilotPanel):
def __init__(self):
super().__init__()
if PC:
self._model_dir = Path(os.path.expanduser("~/.comma/starpilot/data/models"))
else:
self._model_dir = Path("/data/starpilot/models")
self._model_dir = MODELS_PATH
self._model_dir.mkdir(parents=True, exist_ok=True)
self._available_models = []
self._available_model_names = []
self._available_model_series = []
self._available_model_versions = []
self._model_released_dates = {}
self._model_file_to_name = {}
self._model_series_map = {}
self._model_version_map = {}
self._current_model_name = tr("Default")
self._catalog_entries: dict[str, ModelCatalogEntry] = {}
self._model_file_to_name: dict[str, str] = {}
self._model_file_to_name_processed: dict[str, str] = {}
self._model_series_map: dict[str, str] = {}
self._model_released_dates: dict[str, str] = {}
self._model_version_map: dict[str, str] = {}
self._community_favorites: set[str] = set()
self._user_favorites: set[str] = set()
self._current_model_key = self._default_model_key()
self._current_model_name = self._default_model_name()
self.CATEGORIES = [
self.SECTIONS = [
{
"title": tr_noop("Select Model"),
"type": "value",
"icon": "toggle_icons/icon_steering.png",
"on_click": self._on_select_model_clicked,
"get_value": lambda: self._current_model_name,
"visible": lambda: not self._params.get_bool("ModelRandomizer"),
"color": "#597497"
"title": tr_noop("Model Selection"),
"columns": 1,
"uniform_width": True,
"categories": [
{
"title": tr_noop("Select Model"),
"type": "value",
"icon": "toggle_icons/icon_steering.png",
"on_click": self._on_select_model_clicked,
"get_value": lambda: self._current_model_name,
"visible": lambda: not self._params.get_bool("ModelRandomizer"),
"color": "#597497"
},
],
},
{
"title": tr_noop("Download Models"),
"type": "hub",
"icon": "toggle_icons/icon_system.png",
"on_click": self._on_download_clicked,
"color": "#597497"
"title": tr_noop("Model Actions"),
"columns": 2,
"uniform_width": True,
"categories": [
{
"title": tr_noop("Download Models"),
"type": "hub",
"icon": "toggle_icons/icon_system.png",
"on_click": self._on_download_clicked,
"color": "#597497",
"get_status": lambda: self._params_memory.get("ModelDownloadProgress", encoding="utf-8") if self._is_download_active() else ""
},
{
"title": tr_noop("Delete Models"),
"type": "hub",
"icon": "toggle_icons/icon_system.png",
"on_click": self._on_delete_clicked,
"color": "#597497"
},
],
},
{
"title": tr_noop("Delete Models"),
"type": "hub",
"icon": "toggle_icons/icon_system.png",
"on_click": self._on_delete_clicked,
"color": "#597497"
"title": tr_noop("Automation"),
"columns": 2,
"uniform_width": True,
"categories": [
{
"title": tr_noop("Model Randomizer"),
"type": "toggle",
"icon": "toggle_icons/icon_conditional.png",
"get_state": lambda: self._params.get_bool("ModelRandomizer"),
"set_state": self._on_model_randomizer_toggled,
"color": "#597497"
},
{
"title": tr_noop("Auto Download"),
"type": "toggle",
"icon": "toggle_icons/icon_system.png",
"get_state": lambda: self._params.get_bool("AutomaticallyDownloadModels"),
"set_state": lambda s: self._params.put_bool("AutomaticallyDownloadModels", s),
"color": "#597497"
},
],
},
{
"title": tr_noop("Model Randomizer"),
"type": "toggle",
"icon": "toggle_icons/icon_conditional.png",
"get_state": lambda: self._params.get_bool("ModelRandomizer"),
"set_state": self._on_model_randomizer_toggled,
"color": "#597497"
},
{
"title": tr_noop("Recovery Power"),
"type": "value",
"icon": "toggle_icons/icon_road.png",
"get_value": lambda: f"{self._params.get_float('RecoveryPower'):.1f}",
"on_click": self._on_recovery_power_clicked,
"visible": lambda: self._params.get_int("TuningLevel") == 3,
"color": "#597497"
},
{
"title": tr_noop("Stop Distance"),
"type": "value",
"icon": "toggle_icons/icon_road.png",
"get_value": lambda: f"{self._params.get_float('StopDistance'):.1f}m",
"on_click": self._on_stop_distance_clicked,
"visible": lambda: self._params.get_int("TuningLevel") == 3,
"color": "#597497"
},
{
"title": tr_noop("Auto Download"),
"type": "toggle",
"icon": "toggle_icons/icon_system.png",
"get_state": lambda: self._params.get_bool("AutomaticallyDownloadModels"),
"set_state": lambda s: self._params.put_bool("AutomaticallyDownloadModels", s),
"color": "#597497"
},
{
"title": tr_noop("Blacklist"),
"type": "hub",
"icon": "toggle_icons/icon_system.png",
"on_click": self._on_blacklist_clicked,
"title": tr_noop("Randomizer Details"),
"columns": 2,
"uniform_width": True,
"visible": lambda: self._params.get_bool("ModelRandomizer"),
"color": "#597497"
"categories": [
{
"title": tr_noop("Blacklist"),
"type": "hub",
"icon": "toggle_icons/icon_system.png",
"on_click": self._on_blacklist_clicked,
"color": "#597497"
},
{
"title": tr_noop("Ratings"),
"type": "hub",
"icon": "toggle_icons/icon_system.png",
"on_click": self._on_scores_clicked,
"color": "#597497"
},
],
},
{
"title": tr_noop("Ratings"),
"type": "hub",
"icon": "toggle_icons/icon_system.png",
"on_click": self._on_scores_clicked,
"visible": lambda: self._params.get_bool("ModelRandomizer"),
"color": "#597497"
"title": tr_noop("Advanced Tuning"),
"columns": 2,
"uniform_width": True,
"visible": lambda: self._params.get_int("TuningLevel") == 3,
"categories": [
{
"title": tr_noop("Recovery Power"),
"type": "value",
"icon": "toggle_icons/icon_road.png",
"get_value": lambda: f"{self._params.get_float('RecoveryPower'):.1f}",
"on_click": self._on_recovery_power_clicked,
"color": "#597497"
},
{
"title": tr_noop("Stop Distance"),
"type": "value",
"icon": "toggle_icons/icon_road.png",
"get_value": lambda: f"{self._params.get_float('StopDistance'):.1f}m",
"on_click": self._on_stop_distance_clicked,
"color": "#597497"
},
],
},
]
self._model_manager = ModelManager(self._params, self._params_memory)
self._download_thread = None
self._manifest_fetch_thread = None
self._manifest_fetched = False
self._fetch_manifest_async()
self._update_model_metadata()
self._rebuild_grid()
@@ -123,56 +178,187 @@ class StarPilotDrivingModelLayout(StarPilotPanel):
def show_event(self):
super().show_event()
self._fetch_manifest_async()
self._update_model_metadata()
def _is_model_installed(self, key: str) -> bool:
if not key: return False
has_thneed = (self._model_dir / f"{key}.thneed").exists()
if has_thneed: return True
return (self._model_dir / f"{key}_driving_policy_tinygrad.pkl").exists()
def _fetch_manifest_async(self):
if self._manifest_fetch_thread is not None and self._manifest_fetch_thread.is_alive():
return
def _task():
self._model_manager.update_models()
self._manifest_fetched = True
self._manifest_fetch_thread = threading.Thread(target=_task, daemon=True)
self._manifest_fetch_thread.start()
def _default_model_key(self) -> str:
default_key = self._params.get_default_value("Model") or self._params.get_default_value("DrivingModel")
if isinstance(default_key, bytes):
default_key = default_key.decode("utf-8", errors="ignore")
return canonical_model_key(str(default_key or "").strip()) or "sc2"
def _default_model_name(self) -> str:
default_name = self._params.get_default_value("DrivingModelName")
if isinstance(default_name, bytes):
default_name = default_name.decode("utf-8", errors="ignore")
return _clean_model_name(default_name or "") or "South Carolina"
def _default_model_version(self) -> str:
default_version = self._params.get_default_value("ModelVersion") or self._params.get_default_value("DrivingModelVersion")
if isinstance(default_version, bytes):
default_version = default_version.decode("utf-8", errors="ignore")
return str(default_version or "").strip() or "v11"
def _current_selected_key(self) -> str:
current_key = self._params.get("Model", encoding="utf-8") or self._params.get("DrivingModel", encoding="utf-8") or ""
return canonical_model_key(str(current_key).strip()) or self._default_model_key()
def _load_on_disk_files(self) -> set[str]:
try:
return {entry.name for entry in self._model_dir.iterdir()}
except Exception:
return set()
def _is_model_installed(self, key: str, version: str = "", on_disk_files: set[str] | None = None) -> bool:
model_key = canonical_model_key(str(key or "").strip())
if not model_key:
return False
if is_builtin_model_key(model_key):
return True
files = on_disk_files if on_disk_files is not None else self._load_on_disk_files()
if f"{model_key}.thneed" in files:
return True
if version in TINYGRAD_VERSIONS:
required_files = set(self._required_files_for_version(model_key, version))
return required_files.issubset(files)
if version == "v7":
return f"{model_key}.pkl" in files
return any(file.startswith(f"{model_key}.") or file.startswith(f"{model_key}_") for file in files)
def _required_files_for_version(self, key: str, version: str) -> list[str]:
files = [
f"{key}_driving_policy_tinygrad.pkl",
f"{key}_driving_vision_tinygrad.pkl",
f"{key}_driving_policy_metadata.pkl",
f"{key}_driving_vision_metadata.pkl",
]
if version == "v12":
files.extend([
f"{key}_driving_off_policy_tinygrad.pkl",
f"{key}_driving_off_policy_metadata.pkl",
])
return files
def _ensure_default_model_visible(self):
default_key = self._default_model_key()
default_name = self._default_model_name()
default_series = tr("Custom Series")
default_released = ""
for alias in model_key_aliases(default_key):
alias = canonical_model_key(alias)
if alias not in self._model_file_to_name:
continue
default_name = self._model_file_to_name.get(alias, default_name)
default_series = self._model_series_map.get(alias, default_series)
default_released = self._model_released_dates.get(alias, default_released)
if alias != default_key:
self._model_file_to_name.pop(alias, None)
self._model_file_to_name_processed.pop(alias, None)
self._model_series_map.pop(alias, None)
self._model_released_dates.pop(alias, None)
self._model_version_map.pop(alias, None)
self._catalog_entries.pop(alias, None)
version = self._model_version_map.get(default_key, self._default_model_version())
self._model_file_to_name[default_key] = default_name
self._model_file_to_name_processed[default_key] = _clean_model_name(default_name)
self._model_series_map[default_key] = default_series
if default_released:
self._model_released_dates[default_key] = default_released
self._model_version_map.setdefault(default_key, version)
def _build_catalog_entries(self, on_disk_files: set[str]):
self._catalog_entries.clear()
self._model_file_to_name.clear()
self._model_file_to_name_processed.clear()
self._model_series_map.clear()
self._model_released_dates.clear()
self._model_version_map.clear()
available_models = [entry.strip() for entry in (self._params.get("AvailableModels", encoding="utf-8") or "").split(",")]
available_names = [entry.strip() for entry in (self._params.get("AvailableModelNames", encoding="utf-8") or "").split(",")]
available_series = [entry.strip() for entry in (self._params.get("AvailableModelSeries", encoding="utf-8") or "").split(",")]
available_versions = [entry.strip() for entry in (self._params.get("ModelVersions", encoding="utf-8") or "").split(",")]
released_dates = [entry.strip() for entry in (self._params.get("ModelReleasedDates", encoding="utf-8") or "").split(",")]
self._community_favorites = {canonical_model_key(entry.strip()) for entry in (self._params.get("CommunityFavorites", encoding="utf-8") or "").split(",") if entry.strip()}
self._user_favorites = {canonical_model_key(entry.strip()) for entry in (self._params.get("UserFavorites", encoding="utf-8") or "").split(",") if entry.strip()}
size = min(len(available_models), len(available_names))
for i in range(size):
canonical_key = canonical_model_key(available_models[i])
name = available_names[i].strip()
if not canonical_key or not name:
continue
series = available_series[i].strip() if i < len(available_series) and available_series[i].strip() else tr("Custom Series")
version = available_versions[i].strip() if i < len(available_versions) else ""
released = released_dates[i].strip() if i < len(released_dates) else ""
self._model_file_to_name.setdefault(canonical_key, name)
self._model_file_to_name_processed.setdefault(canonical_key, _clean_model_name(name))
self._model_series_map.setdefault(canonical_key, series)
if released:
self._model_released_dates.setdefault(canonical_key, released)
if version:
self._model_version_map.setdefault(canonical_key, version)
self._ensure_default_model_visible()
for key, name in self._model_file_to_name.items():
version = self._model_version_map.get(key, self._default_model_version() if is_builtin_model_key(key) else "")
installed = self._is_model_installed(key, version, on_disk_files)
partial = (not is_builtin_model_key(key)) and (not installed) and any(file.startswith(f"{key}.") or file.startswith(f"{key}_") for file in on_disk_files)
self._catalog_entries[key] = ModelCatalogEntry(
key=key,
name=name,
series=self._model_series_map.get(key, tr("Custom Series")),
version=version,
released=self._model_released_dates.get(key, ""),
builtin=is_builtin_model_key(key),
installed=installed,
partial=partial,
community_favorite=(key in self._community_favorites),
user_favorite=(key in self._user_favorites),
)
def _update_model_metadata(self):
available_models_raw = self._params.get("AvailableModels", encoding='utf-8')
if not available_models_raw:
manager = ModelManager(self._params, self._params_memory)
manager.update_models()
available_models_raw = self._params.get("AvailableModels", encoding='utf-8')
on_disk_files = self._load_on_disk_files()
self._build_catalog_entries(on_disk_files)
self._available_models = [m.strip() for m in (available_models_raw or "").split(",") if m.strip()]
self._available_model_names = [m.strip() for m in (self._params.get("AvailableModelNames", encoding='utf-8') or "").split(",") if m.strip()]
self._available_model_series = [m.strip() for m in (self._params.get("AvailableModelSeries", encoding='utf-8') or "").split(",") if m.strip()]
self._available_model_versions = [m.strip() for m in (self._params.get("ModelVersions", encoding='utf-8') or "").split(",")]
released_dates = (self._params.get("ModelReleasedDates", encoding='utf-8') or "").split(",")
self._current_model_key = self._current_selected_key()
current_entry = self._catalog_entries.get(self._current_model_key)
if current_entry is None or not current_entry.installed:
self._current_model_key = self._default_model_key()
current_entry = self._catalog_entries.get(self._current_model_key)
self._model_file_to_name.clear()
self._model_series_map.clear()
self._model_version_map.clear()
self._model_released_dates.clear()
if current_entry is not None:
self._current_model_name = current_entry.name
else:
self._current_model_name = self._default_model_name()
size = min(len(self._available_models), len(self._available_model_names))
for i in range(size):
key = self._available_models[i].strip()
name = self._available_model_names[i].strip()
if not key or not name: continue
series = self._available_model_series[i].strip() if i < len(self._available_model_series) else tr("Custom Series")
self._model_file_to_name[key] = name
self._model_series_map[key] = series
if i < len(self._available_model_versions):
v = self._available_model_versions[i].strip()
if v: self._model_version_map[key] = v
if i < len(released_dates):
d = released_dates[i].strip()
if d: self._model_released_dates[key] = d
model_key = self._params.get("Model") or self._params.get("DrivingModel")
if model_key and isinstance(model_key, bytes): model_key = model_key.decode()
if not model_key or not self._is_model_installed(model_key):
model_key = self._params.get_default_value("Model") or self._params.get_default_value("DrivingModel") or ""
if model_key and isinstance(model_key, bytes): model_key = model_key.decode()
self._current_model_name = self._model_file_to_name.get(model_key, "Default")
def _show_selection_dialog(self, title: str, options: dict[str, str] | list[str], current_val: str, on_confirm: Callable):
def _show_selection_dialog(self, title: str, options: dict[str, str] | list[str], current_val: str, on_confirm: Callable, current_key: str = ""):
if not options:
gui_app.set_modal_overlay(alert_dialog(tr("No options available.")))
return
@@ -186,11 +372,16 @@ class StarPilotDrivingModelLayout(StarPilotPanel):
grouped = {}
name_to_key = {}
key_to_display = {}
name_counts = {}
for key, name in options.items():
series = self._model_series_map.get(key, tr("Custom Series"))
if series not in grouped: grouped[series] = []
grouped[series].append(name)
name_to_key[name] = key
name_counts[name] = name_counts.get(name, 0) + 1
display_name = name if name_counts[name] == 1 else f"{name} [{key}]"
grouped[series].append(display_name)
name_to_key[display_name] = key
key_to_display[key] = display_name
for series in grouped: grouped[series].sort()
sorted_series = sorted(grouped.keys())
@@ -214,8 +405,10 @@ class StarPilotDrivingModelLayout(StarPilotPanel):
user_favs = [f.strip() for f in (self._params.get("UserFavorites", encoding='utf-8') or "").split(",") if f.strip()]
comm_favs = [f.strip() for f in (self._params.get("CommunityFavorites", encoding='utf-8') or "").split(",") if f.strip()]
current_display = key_to_display.get(current_key, current_val)
dialog = SelectionDialog(
title, final_grouped, current_val, on_close=_on_close_grouped,
title, final_grouped, current_display, on_close=_on_close_grouped,
model_released_dates=self._model_released_dates,
model_file_to_name=self._model_file_to_name,
user_favorites=user_favs,
@@ -224,23 +417,75 @@ class StarPilotDrivingModelLayout(StarPilotPanel):
)
gui_app.set_modal_overlay(dialog)
def _is_download_active(self) -> bool:
return bool(self._params_memory.get("ModelToDownload", encoding="utf-8") or self._params_memory.get_bool("DownloadAllModels"))
def _selected_model_version(self, model_key: str) -> str:
version = self._model_version_map.get(model_key, "")
if version:
return version
try:
versions_file = self._model_dir / ".model_versions.json"
if versions_file.is_file():
payload = json.loads(versions_file.read_text())
if isinstance(payload, dict):
for alias in model_key_aliases(model_key):
resolved = str(payload.get(alias, "")).strip()
if resolved:
return resolved
except Exception:
pass
if is_builtin_model_key(model_key):
return self._default_model_version()
return ""
def _build_selectable_models(self) -> dict[str, str]:
models: dict[str, str] = {}
for key, entry in self._catalog_entries.items():
if entry.installed:
models[key] = entry.name
return models
def _build_deletable_models(self) -> dict[str, str]:
installed = self._build_selectable_models()
default_key = self._default_model_key()
current_name = _clean_model_name(self._current_model_name)
default_name = _clean_model_name(installed.get(default_key, self._default_model_name()))
deletable: dict[str, str] = {}
for key, display_name in installed.items():
processed_name = _clean_model_name(display_name)
if processed_name == current_name or processed_name == default_name:
continue
deletable[key] = display_name
return deletable
def _on_select_model_clicked(self):
installed_models = {k: v for k, v in self._model_file_to_name.items() if self._is_model_installed(k)}
if not installed_models: return
self._update_model_metadata()
installed_models = self._build_selectable_models()
if not installed_models:
gui_app.set_modal_overlay(alert_dialog(tr("No downloaded models found.")))
return
def _on_confirm(model_key):
self._params.put("Model", model_key)
self._params.put("DrivingModel", model_key)
selected_model = canonical_model_key(model_key)
self._params.put("Model", selected_model)
self._params.put("DrivingModel", selected_model)
self._params.put("DrivingModelName", installed_models[model_key])
mv = self._model_version_map.get(model_key, "")
if mv:
self._params.put("ModelVersion", mv)
self._params.put("DrivingModelVersion", mv)
resolved_version = self._selected_model_version(selected_model)
if resolved_version:
self._params.put("ModelVersion", resolved_version)
self._params.put("DrivingModelVersion", resolved_version)
update_starpilot_toggles()
self._update_model_metadata()
if ui_state.started:
gui_app.set_modal_overlay(ConfirmDialog(tr("Reboot required. Reboot now?"), tr("Reboot"), tr("Cancel"), on_close=lambda res: HARDWARE.reboot() if res == DialogResult.CONFIRM else None))
self._params.put_bool("OnroadCycleRequested", True)
gui_app.set_modal_overlay(alert_dialog(tr("Drive-cycle requested for immediate apply.")))
self._show_selection_dialog(tr("Select Driving Model"), installed_models, self._current_model_name, _on_confirm)
self._show_selection_dialog(tr("Select Driving Model"), installed_models, self._current_model_name, _on_confirm, current_key=self._current_model_key)
def _on_recovery_power_clicked(self):
def on_close(res, val):
@@ -259,29 +504,48 @@ class StarPilotDrivingModelLayout(StarPilotPanel):
gui_app.set_modal_overlay(AetherSliderDialog(tr("Stop Distance"), 4.0, 10.0, 0.1, self._params.get_float("StopDistance"), on_close, unit="m", color="#597497"))
def _on_download_clicked(self):
is_downloading = self._params_memory.get("ModelToDownload") or self._params_memory.get_bool("DownloadAllModels")
self._update_model_metadata()
if ui_state.started:
gui_app.set_modal_overlay(alert_dialog(tr("Cannot download models while driving.")))
return
is_downloading = self._is_download_active()
if is_downloading:
self._params_memory.remove("ModelToDownload")
self._params_memory.put_bool("DownloadAllModels", False)
self._params_memory.remove("ModelDownloadProgress")
self._params_memory.put_bool("CancelModelDownload", True)
return
not_installed = {key: entry.name for key, entry in self._catalog_entries.items() if not entry.installed}
if not not_installed:
gui_app.set_modal_overlay(alert_dialog(tr("All models are already installed.")))
return
not_installed = {k: v for k, v in self._model_file_to_name.items() if not self._is_model_installed(k)}
self._show_selection_dialog(tr("Select Model to Download"), not_installed, "", lambda mk: self._params_memory.put("ModelToDownload", mk))
def _on_delete_clicked(self):
installed = {k: v for k, v in self._model_file_to_name.items() if self._is_model_installed(k)}
dk = self._params.get_default_value("Model") or ""
if isinstance(dk, bytes): dk = dk.decode()
ck = self._params.get("Model", encoding='utf-8') or ""
deletable = {k: v for k, v in installed.items() if k != dk and k != ck}
self._update_model_metadata()
if ui_state.started:
gui_app.set_modal_overlay(alert_dialog(tr("Cannot delete model files while driving.")))
return
if self._is_download_active():
gui_app.set_modal_overlay(alert_dialog(tr("Cannot delete model files while a download is in progress.")))
return
deletable = self._build_deletable_models()
if not deletable:
gui_app.set_modal_overlay(alert_dialog(tr("No deletable models found.")))
return
def _on_confirm(mk):
def _execute_delete(res):
if res == DialogResult.CONFIRM:
for file in self._model_dir.iterdir():
if file.name.startswith(mk): file.unlink()
self._update_model_metadata()
if res == DialogResult.CONFIRM:
for file in self._model_dir.iterdir():
if not (file.name == f"{mk}.thneed" or file.name == f"{mk}.pkl" or file.name.startswith(f"{mk}_")):
continue
if file.is_file():
file.unlink(missing_ok=True)
self._update_model_metadata()
self._rebuild_grid()
gui_app.set_modal_overlay(ConfirmDialog(tr(f"Delete '{deletable[mk]}'?"), tr("Delete"), on_close=_execute_delete))
self._show_selection_dialog(tr("Select Model to Delete"), deletable, "", _on_confirm)
@@ -317,7 +581,7 @@ class StarPilotDrivingModelLayout(StarPilotPanel):
def _on_model_randomizer_toggled(self, state: bool):
self._params.put_bool("ModelRandomizer", state)
if state:
not_installed = [k for k in self._model_file_to_name if not self._is_model_installed(k)]
not_installed = [key for key, entry in self._catalog_entries.items() if not entry.installed]
if not_installed:
def _on_download_confirm(res):
if res == DialogResult.CONFIRM:
@@ -326,16 +590,32 @@ class StarPilotDrivingModelLayout(StarPilotPanel):
gui_app.set_modal_overlay(ConfirmDialog(tr("Download all models for Randomizer?"), tr("Download All"), on_close=_on_download_confirm))
def _update_state(self):
if getattr(self, "_manifest_fetched", False):
self._manifest_fetched = False
self._update_model_metadata()
self._rebuild_grid()
model_to_download = self._params_memory.get("ModelToDownload", encoding='utf-8') or ""
download_all = self._params_memory.get_bool("DownloadAllModels")
is_downloading = bool(model_to_download or download_all)
if is_downloading and (self._download_thread is None or not self._download_thread.is_alive()):
def _download_task():
try:
if download_all: self._model_manager.download_all_models()
else: self._model_manager.download_model(model_to_download)
except: pass
finally: self._download_thread = None
self._download_thread = threading.Thread(target=_download_task, daemon=True)
self._download_thread.start()
def _download_task():
try:
if download_all:
self._model_manager.download_all_models()
else:
self._model_manager.download_model(model_to_download)
except Exception:
pass
finally:
self._params_memory.remove("CancelModelDownload")
self._params_memory.remove("ModelToDownload")
self._params_memory.put_bool("DownloadAllModels", False)
self._params_memory.remove("ModelDownloadProgress")
self._download_thread = None
self._update_model_metadata()
self._rebuild_grid()
self._download_thread = threading.Thread(target=_download_task, daemon=True)
self._download_thread.start()
@@ -9,6 +9,7 @@ from openpilot.common.params import Params
from openpilot.system.ui.lib.multilang import tr
from openpilot.system.ui.widgets import Widget
from openpilot.selfdrive.ui.layouts.settings.starpilot.aethergrid import TileGrid, HubTile, ToggleTile, ValueTile, SPACING
from openpilot.selfdrive.ui.layouts.settings.starpilot.sectioned_panel import SectionedTileLayout, TileSection
class StarPilotPanelType(IntEnum):
@@ -46,7 +47,9 @@ class StarPilotPanel(Widget):
self._sub_panels: dict[str, Widget] = {}
self._scroller = None
self._tile_grid = None
self._sectioned_grid = None
self.CATEGORIES = []
self.SECTIONS = []
def set_navigate_callback(self, callback: Callable):
self._navigate_callback = callback
@@ -57,7 +60,74 @@ class StarPilotPanel(Widget):
def set_current_sub_panel(self, sub_panel: str):
self._current_sub_panel = sub_panel
def _is_category_visible(self, cat: dict) -> bool:
visible_fn = cat.get("visible")
return visible_fn is None or visible_fn()
def _build_tile(self, cat: dict) -> Widget | None:
tile_type = cat.get("type", "hub")
if tile_type == "hub":
on_click = cat.get("on_click")
if on_click is None:
on_click = lambda c=cat: self._navigate_to(c["panel"])
return HubTile(
title=tr(cat["title"]),
desc=tr(cat.get("desc", "")),
icon_path=cat.get("icon"),
on_click=on_click,
starpilot_icon=cat.get("starpilot_icon", True),
bg_color=cat.get("color"),
get_status=cat.get("get_status"),
)
if tile_type == "toggle":
raw_set_state = cat["set_state"]
def on_toggle(state: bool, setter=raw_set_state):
setter(state)
self._rebuild_grid()
return ToggleTile(title=tr(cat["title"]), get_state=cat["get_state"], set_state=on_toggle, icon_path=cat.get("icon"), bg_color=cat.get("color"), desc=tr(cat.get("desc", "")), is_enabled=cat.get("is_enabled"), disabled_label=cat.get("disabled_label", ""))
if tile_type == "value":
return ValueTile(title=tr(cat["title"]), get_value=cat["get_value"], on_click=cat["on_click"], icon_path=cat.get("icon"), bg_color=cat.get("color"), is_enabled=cat.get("is_enabled"), desc=tr(cat.get("desc", "")))
return None
def _build_tile_grid(self, categories: list[dict], columns: int | None = None, padding: int | None = None, uniform_width: bool = False) -> TileGrid:
grid = TileGrid(columns=columns, padding=padding, uniform_width=uniform_width)
for cat in categories:
if not self._is_category_visible(cat):
continue
tile = self._build_tile(cat)
if tile is not None:
grid.add_tile(tile)
return grid
def _rebuild_grid(self):
if self.SECTIONS:
if self._sectioned_grid is None:
self._sectioned_grid = SectionedTileLayout()
sections: list[TileSection] = []
for section in self.SECTIONS:
visible_fn = section.get("visible")
if visible_fn is not None and not visible_fn():
continue
grid = self._build_tile_grid(
section.get("categories", []),
columns=section.get("columns", 2),
padding=section.get("padding", SPACING.tile_gap),
uniform_width=section.get("uniform_width", True),
)
if grid.tiles:
sections.append(TileSection(tr(section["title"]), grid))
self._sectioned_grid.set_sections(sections)
return
if not self.CATEGORIES:
return
@@ -67,38 +137,11 @@ class StarPilotPanel(Widget):
self._tile_grid.clear()
for cat in self.CATEGORIES:
visible_fn = cat.get("visible")
if visible_fn is not None and not visible_fn():
if not self._is_category_visible(cat):
continue
tile_type = cat.get("type", "hub")
if tile_type == "hub":
on_click = cat.get("on_click")
if on_click is None:
on_click = lambda c=cat: self._navigate_to(c["panel"])
tile = HubTile(
title=tr(cat["title"]),
desc=tr(cat.get("desc", "")),
icon_path=cat.get("icon"),
on_click=on_click,
starpilot_icon=cat.get("starpilot_icon", True),
bg_color=cat.get("color"),
)
elif tile_type == "toggle":
raw_set_state = cat["set_state"]
def on_toggle(state: bool, setter=raw_set_state):
setter(state)
self._rebuild_grid()
tile = ToggleTile(title=tr(cat["title"]), get_state=cat["get_state"], set_state=on_toggle, icon_path=cat.get("icon"), bg_color=cat.get("color"), desc=tr(cat.get("desc", "")), is_enabled=cat.get("is_enabled"), disabled_label=cat.get("disabled_label", ""))
elif tile_type == "value":
tile = ValueTile(title=tr(cat["title"]), get_value=cat["get_value"], on_click=cat["on_click"], icon_path=cat.get("icon"), bg_color=cat.get("color"), is_enabled=cat.get("is_enabled"), desc=tr(cat.get("desc", "")))
else:
continue
self._tile_grid.add_tile(tile)
tile = self._build_tile(cat)
if tile is not None:
self._tile_grid.add_tile(tile)
def _navigate_to(self, sub_panel: str):
self._current_sub_panel = sub_panel
@@ -113,6 +156,8 @@ class StarPilotPanel(Widget):
def _render(self, rect: rl.Rectangle):
if self._current_sub_panel and self._current_sub_panel in self._sub_panels:
self._sub_panels[self._current_sub_panel].render(rect)
elif self.SECTIONS and self._sectioned_grid:
self._sectioned_grid.render(rect)
elif self.CATEGORIES and self._tile_grid:
self._tile_grid.render(rect)
elif self._scroller:
@@ -123,9 +168,20 @@ class StarPilotPanel(Widget):
self._rebuild_grid()
if self._current_sub_panel and self._current_sub_panel in self._sub_panels:
self._sub_panels[self._current_sub_panel].show_event()
elif self.SECTIONS and self._sectioned_grid:
self._sectioned_grid.show_event()
elif self._scroller:
self._scroller.show_event()
def hide_event(self):
super().hide_event()
if self._current_sub_panel and self._current_sub_panel in self._sub_panels:
self._sub_panels[self._current_sub_panel].hide_event()
elif self.SECTIONS and self._sectioned_grid:
self._sectioned_grid.hide_event()
elif self._scroller:
self._scroller.hide_event()
def create_tile_panel(categories: list[dict], sub_panels: dict[str, Widget] | None = None) -> StarPilotPanel:
panel = StarPilotPanel()
@@ -0,0 +1,132 @@
from __future__ import annotations
from dataclasses import dataclass
import pyray as rl
from openpilot.system.ui.lib.application import FontWeight, gui_app
from openpilot.system.ui.lib.text_measure import measure_text_cached
from openpilot.system.ui.widgets import Widget
from openpilot.selfdrive.ui.layouts.settings.starpilot.aethergrid import SPACING, TileGrid
@dataclass(frozen=True)
class TileSection:
title: str
grid: TileGrid
class SectionedTileLayout(Widget):
def __init__(self, section_gap: int = SPACING.section_gap, title_height: int = 32, title_gap: int = SPACING.sm,
min_row_height: int = 150, max_row_height: int = 280, top_padding: int = 0,
horizontal_padding: int = SPACING.xl, max_content_width: int | None = 1440):
super().__init__()
self._sections: list[TileSection] = []
self._section_gap = section_gap
self._title_height = title_height
self._title_gap = title_gap
self._min_row_height = min_row_height
self._max_row_height = max_row_height
self._top_padding = top_padding
self._horizontal_padding = horizontal_padding
self._max_content_width = max_content_width
self._title_font_size = 26
self._font_title = gui_app.font(FontWeight.BOLD)
self._is_active = False
def set_sections(self, sections: list[TileSection]):
if self._is_active:
for section in self._sections:
section.grid.hide_event()
self._sections = list(sections)
if self._is_active:
for section in self._sections:
section.grid.show_event()
def clear(self):
self._sections.clear()
def show_event(self):
self._is_active = True
super().show_event()
for section in self._sections:
section.grid.show_event()
def hide_event(self):
self._is_active = False
super().hide_event()
for section in self._sections:
section.grid.hide_event()
def _title_block_height(self, section: TileSection) -> int:
return (self._title_height + self._title_gap) if section.title else 0
def _section_band_height(self, section: TileSection, row_height: float) -> float:
rows = section.grid.get_row_count()
if rows <= 0:
return 0.0
return (rows * row_height) + (section.grid.gap * max(0, rows - 1))
def _content_rect(self, rect: rl.Rectangle) -> rl.Rectangle:
content_x = rect.x + self._horizontal_padding
content_w = max(0.0, rect.width - (self._horizontal_padding * 2))
if self._max_content_width is not None and content_w > self._max_content_width:
content_w = float(self._max_content_width)
content_x = rect.x + (rect.width - content_w) / 2
return rl.Rectangle(content_x, rect.y, content_w, rect.height)
def _compute_row_height(self, rect: rl.Rectangle, sections: list[TileSection]) -> float:
total_rows = sum(section.grid.get_row_count() for section in sections)
if total_rows <= 0:
return 0.0
total_title_height = sum(self._title_block_height(section) for section in sections)
total_section_gaps = self._section_gap * max(0, len(sections) - 1)
total_internal_gaps = sum(section.grid.get_internal_gap_height() for section in sections)
# Clamp oversized sections so tiles keep a touch-friendly shape instead of stretching to fill the full panel.
fit_row_height = max(0.0, (rect.height - self._top_padding - total_title_height - total_section_gaps - total_internal_gaps) / total_rows)
if fit_row_height >= self._min_row_height:
return min(fit_row_height, self._max_row_height)
return fit_row_height
def _draw_section_title(self, rect: rl.Rectangle, title: str):
title_text = title.upper()
spacing = round(self._title_font_size * 0.08)
size = measure_text_cached(self._font_title, title_text, self._title_font_size, spacing=spacing)
text_y = rect.y + (rect.height - size.y) / 2
text_pos = rl.Vector2(round(rect.x), round(text_y))
rl.draw_text_ex(self._font_title, title_text, rl.Vector2(text_pos.x + 1, text_pos.y + 1), self._title_font_size, spacing, rl.Color(0, 0, 0, 90))
rl.draw_text_ex(self._font_title, title_text, text_pos, self._title_font_size, spacing, rl.Color(255, 255, 255, 215))
line_x = rect.x + size.x + SPACING.lg
line_w = rect.width - (line_x - rect.x)
if line_w <= 0:
return
line_y = int(rect.y + rect.height / 2)
rl.draw_rectangle(int(line_x), line_y, int(line_w), 2, rl.Color(255, 255, 255, 36))
def _render(self, rect: rl.Rectangle):
self.set_rect(rect)
sections = [section for section in self._sections if section.grid.tiles]
if not sections:
return
content_rect = self._content_rect(rect)
row_height = self._compute_row_height(content_rect, sections)
if row_height <= 0:
return
y = content_rect.y + self._top_padding
for index, section in enumerate(sections):
if section.title:
title_rect = rl.Rectangle(content_rect.x, y, content_rect.width, self._title_height)
self._draw_section_title(title_rect, section.title)
y += self._title_height + self._title_gap
active_grid_height = (section.grid.get_row_count() * row_height) + section.grid.get_internal_gap_height()
section.grid.render(rl.Rectangle(content_rect.x, y, content_rect.width, active_grid_height))
y += self._section_band_height(section, row_height)
if index < len(sections) - 1:
y += self._section_gap
+7 -29
View File
@@ -64,17 +64,14 @@ class SelectionHeader(Widget):
self._pressed = False
class SelectionItem(Widget):
def __init__(self, text: str, is_selected: bool, is_favorite: bool, callback: Callable[[str], None], fav_callback: Callable[[str], None] = None):
def __init__(self, text: str, is_selected: bool, callback: Callable[[str], None]):
super().__init__()
self._text = text
self._is_selected = is_selected
self._is_favorite = is_favorite
self._callback = callback
self._fav_callback = fav_callback
self._font = gui_app.font(FontWeight.MEDIUM)
self._font_size = 48
self._pressed = False
self._fav_pressed = False
self.set_rect(rl.Rectangle(0, 0, 0, 110))
def set_parent_rect(self, parent_rect: rl.Rectangle) -> None:
@@ -94,14 +91,9 @@ class SelectionItem(Widget):
if self._is_selected:
rl.draw_rectangle_rounded_lines_ex(rect, 0.1, 10, 3, rl.WHITE)
# Favorite Star - Left side
star = "" if self._is_favorite else ""
star_pos = rl.Vector2(rect.x + 25, rect.y + (rect.height - self._font_size) / 2)
rl.draw_text_ex(self._font, star, star_pos, self._font_size + 10, 0, rl.WHITE)
# Text
text_size = rl.measure_text_ex(self._font, self._text, self._font_size, 0)
text_pos = rl.Vector2(rect.x + 90, rect.y + (rect.height - text_size.y) / 2)
text_pos = rl.Vector2(rect.x + 40, rect.y + (rect.height - text_size.y) / 2)
rl.draw_text_ex(self._font, self._text, text_pos, self._font_size, 0, rl.WHITE)
# Indicator (Dot for selection instead of radio)
@@ -109,25 +101,15 @@ class SelectionItem(Widget):
circle_center = rl.Vector2(rect.x + rect.width - 50, rect.y + rect.height / 2)
rl.draw_circle_v(circle_center, 12, rl.WHITE)
@property
def _fav_rect(self) -> rl.Rectangle:
return rl.Rectangle(self._rect.x, self._rect.y, 80, self._rect.height)
def _handle_mouse_press(self, mouse_pos):
if rl.check_collision_point_rec(mouse_pos, self._fav_rect):
self._fav_pressed = True
elif rl.check_collision_point_rec(mouse_pos, self._hit_rect):
if rl.check_collision_point_rec(mouse_pos, self._hit_rect):
self._pressed = True
def _handle_mouse_release(self, mouse_pos):
if self._fav_pressed and rl.check_collision_point_rec(mouse_pos, self._fav_rect):
if self._fav_callback:
self._fav_callback(self._text)
elif self._pressed and rl.check_collision_point_rec(mouse_pos, self._hit_rect):
if self._pressed and rl.check_collision_point_rec(mouse_pos, self._hit_rect):
if self._callback:
self._callback(self._text)
self._pressed = False
self._fav_pressed = False
class SelectionDialog(Widget):
def __init__(self, title: str, options, current_selection: str = "",
@@ -214,20 +196,16 @@ class SelectionDialog(Widget):
for model in sorted_models:
key = self._name_to_file.get(model, model)
is_selected = (model == self._selected_value or key == self._selected_value)
is_fav = key in self._user_favorites or key in self._community_favorites
items.append(SelectionItem(
text=model,
is_selected=is_selected,
is_favorite=is_fav,
callback=self._on_item_selected,
fav_callback=self._toggle_favorite if self._favorites_editable else None
callback=self._on_item_selected
))
else:
for option in self._options_raw:
items.append(SelectionItem(
text=option,
is_selected=(option == self._selected_value),
is_favorite=False,
callback=self._on_item_selected
))
@@ -257,14 +235,14 @@ class SelectionDialog(Widget):
item._is_selected = (item._text == val)
def _cancel_button_callback(self):
gui_app.set_modal_overlay(None)
if self._on_close:
self._on_close(DialogResult.CANCEL, "")
gui_app.set_modal_overlay(None)
def _confirm_button_callback(self):
gui_app.set_modal_overlay(None)
if self._on_close:
self._on_close(DialogResult.CONFIRM, self._selected_value)
gui_app.set_modal_overlay(None)
def show_event(self):
super().show_event()