| from __future__ import annotations |
|
|
| from dataclasses import asdict |
| from datetime import datetime |
| import json |
| from pathlib import Path |
| from uuid import uuid4 |
|
|
| from PySide6.QtCore import QSize, Qt, QThread, QTimer, QUrl, Signal |
| from PySide6.QtGui import QBrush, QDesktopServices, QIcon, QImage, QImageReader, QPixmap |
| from PySide6.QtWidgets import ( |
| QAbstractItemView, |
| QComboBox, |
| QDialog, |
| QDialogButtonBox, |
| QFileDialog, |
| QFrame, |
| QGridLayout, |
| QHBoxLayout, |
| QHeaderView, |
| QLabel, |
| QLineEdit, |
| QListWidget, |
| QListWidgetItem, |
| QMessageBox, |
| QPlainTextEdit, |
| QProgressBar, |
| QPushButton, |
| QSpinBox, |
| QTabWidget, |
| QTableWidget, |
| QTableWidgetItem, |
| QVBoxLayout, |
| QWidget, |
| ) |
|
|
| from adam.assets import Asset, AssetRegistry |
| from adam.config import ConfigManager |
| from adam.job_manager import JobManager |
| from adam.models import Job |
| from adam.eve import EveResult, EveVisionModel, save_eve_results |
| from adam.studio import ( |
| PreviewEvaluation, |
| StudioStore, |
| TrainingRecipe, |
| caption_path, |
| checkpoint_files, |
| exact_duplicate_groups, |
| ) |
| from adam.ui.theme import COLORS |
|
|
|
|
| def _card() -> QFrame: |
| frame = QFrame() |
| frame.setProperty("card", True) |
| return frame |
|
|
|
|
| def _title(text: str) -> QLabel: |
| label = QLabel(text) |
| label.setObjectName("CardTitle") |
| return label |
|
|
|
|
| def _header(title: str, subtitle: str) -> QWidget: |
| widget = QWidget() |
| layout = QVBoxLayout(widget) |
| layout.setContentsMargins(0, 0, 0, 10) |
| heading = QLabel(title) |
| heading.setObjectName("PageTitle") |
| detail = QLabel(subtitle) |
| detail.setProperty("muted", True) |
| detail.setWordWrap(True) |
| layout.addWidget(heading) |
| layout.addWidget(detail) |
| return widget |
|
|
|
|
| def _thumbnail(path: Path, width: int, height: int) -> QPixmap: |
| """Decode close to display size instead of loading a full-resolution image.""" |
| reader = QImageReader(str(path)) |
| reader.setAutoTransform(True) |
| source_size = reader.size() |
| if source_size.isValid(): |
| source_size.scale(QSize(width, height), Qt.KeepAspectRatio) |
| reader.setScaledSize(source_size) |
| image = reader.read() |
| return QPixmap.fromImage(image) if not image.isNull() else QPixmap() |
|
|
|
|
| class ImageScanWorker(QThread): |
| scanned = Signal(object, int) |
|
|
| def __init__(self, folder: str, token: int, limit: int = 2500) -> None: |
| super().__init__() |
| self.folder = folder |
| self.token = token |
| self.limit = limit |
|
|
| def run(self) -> None: |
| root = Path(self.folder).expanduser() |
| paths: list[Path] = [] |
| if root.is_dir(): |
| for path in root.rglob("*"): |
| if self.isInterruptionRequested(): |
| return |
| if path.is_file() and path.suffix.casefold() in { |
| ".png", ".jpg", ".jpeg", ".webp", ".bmp", ".gif" |
| }: |
| paths.append(path) |
| if len(paths) >= self.limit: |
| break |
| if not self.isInterruptionRequested(): |
| self.scanned.emit(sorted(paths), self.token) |
|
|
|
|
| class EveReviewWorker(QThread): |
| progress = Signal(int, str) |
| completed = Signal(object) |
| failed = Signal(str) |
|
|
| def __init__( |
| self, |
| paths: list[Path], |
| good_references: list[Path], |
| bad_references: list[Path], |
| keep_threshold: float, |
| reject_threshold: float, |
| ) -> None: |
| super().__init__() |
| self.paths = paths |
| self.good_references = good_references |
| self.bad_references = bad_references |
| self.keep_threshold = keep_threshold |
| self.reject_threshold = reject_threshold |
|
|
| def run(self) -> None: |
| try: |
| self.progress.emit(1, "Loading EVE's local vision model…") |
|
|
| def on_progress(done: int, total: int) -> None: |
| if self.isInterruptionRequested(): |
| raise RuntimeError("EVE review cancelled.") |
| percent = 5 + int(done / max(1, total) * 94) |
| self.progress.emit(percent, f"EVE analyzed {done} of {total} images…") |
|
|
| results = EveVisionModel().review( |
| self.paths, |
| self.good_references, |
| self.bad_references, |
| keep_threshold=self.keep_threshold, |
| reject_threshold=self.reject_threshold, |
| progress=on_progress, |
| ) |
| self.progress.emit(100, "EVE finished sorting the dataset.") |
| self.completed.emit(results) |
| except Exception as exc: |
| self.failed.emit(str(exc)) |
|
|
|
|
| class EveReviewDialog(QDialog): |
| applied = Signal(object) |
|
|
| def __init__( |
| self, |
| root_path: Path, |
| dataset_path: str, |
| paths: list[Path], |
| store: StudioStore, |
| parent: QWidget | None = None, |
| ) -> None: |
| super().__init__(parent) |
| self.root_path = root_path |
| self.dataset_path = dataset_path |
| self.paths = paths |
| self.store = store |
| self.results: list[EveResult] = [] |
| self.worker: EveReviewWorker | None = None |
| self.setWindowTitle("EVE · AI Dataset Review") |
| self.setMinimumSize(980, 700) |
| root = QVBoxLayout(self) |
| root.addWidget(_header( |
| "EVE · AI Dataset Review", |
| "Give EVE a few good examples and optional bad examples. EVE proposes Keep, Reject, or Uncertain; nothing changes until you apply the review.", |
| )) |
|
|
| references = QHBoxLayout() |
| self.good_references = self._reference_panel( |
| references, "GOOD REFERENCES", "Add good images…", self._add_good_references |
| ) |
| self.bad_references = self._reference_panel( |
| references, "BAD REFERENCES (OPTIONAL)", "Add bad images…", self._add_bad_references |
| ) |
| root.addLayout(references) |
|
|
| controls = QHBoxLayout() |
| self.keep_threshold = QSpinBox(); self.keep_threshold.setRange(51, 99); self.keep_threshold.setValue(75); self.keep_threshold.setSuffix("%") |
| self.reject_threshold = QSpinBox(); self.reject_threshold.setRange(1, 49); self.reject_threshold.setValue(25); self.reject_threshold.setSuffix("%") |
| self.analyze_button = QPushButton("Analyze dataset with EVE") |
| self.analyze_button.setProperty("primary", True) |
| self.analyze_button.clicked.connect(self._analyze) |
| controls.addWidget(QLabel("Keep at or above")); controls.addWidget(self.keep_threshold) |
| controls.addWidget(QLabel("Reject at or below")); controls.addWidget(self.reject_threshold) |
| controls.addStretch(); controls.addWidget(self.analyze_button) |
| root.addLayout(controls) |
| self.progress = QProgressBar(); self.progress.setRange(0, 100); self.progress.setValue(0) |
| self.status = QLabel( |
| "Tip: 3–10 varied good references work best. Bad references help EVE distinguish visually similar mistakes." |
| ) |
| self.status.setProperty("muted", True); self.status.setWordWrap(True) |
| root.addWidget(self.progress); root.addWidget(self.status) |
|
|
| self.tabs = QTabWidget() |
| self.result_lists: dict[str, QListWidget] = {} |
| for decision, label in (("keep", "KEEP"), ("reject", "REJECT"), ("unreviewed", "UNCERTAIN")): |
| gallery = QListWidget() |
| gallery.setViewMode(QListWidget.IconMode) |
| gallery.setIconSize(QSize(120, 90)); gallery.setGridSize(QSize(165, 145)) |
| gallery.setResizeMode(QListWidget.Adjust) |
| gallery.setSelectionMode(QAbstractItemView.ExtendedSelection) |
| gallery.itemDoubleClicked.connect(lambda item: QDesktopServices.openUrl(QUrl.fromLocalFile(str(item.data(Qt.UserRole))))) |
| self.result_lists[decision] = gallery |
| self.tabs.addTab(gallery, label) |
| root.addWidget(self.tabs, 1) |
|
|
| moves = QHBoxLayout() |
| select_all = QPushButton("Select all in current group") |
| clear_selection = QPushButton("Clear selection") |
| to_keep = QPushButton("Move selected to Keep") |
| to_reject = QPushButton("Move selected to Reject") |
| to_uncertain = QPushButton("Move selected to Uncertain") |
| select_all.clicked.connect(self._select_all_current) |
| clear_selection.clicked.connect(self._clear_current_selection) |
| to_keep.clicked.connect(lambda: self._move_selected("keep")) |
| to_reject.clicked.connect(lambda: self._move_selected("reject")) |
| to_uncertain.clicked.connect(lambda: self._move_selected("unreviewed")) |
| moves.addWidget(select_all); moves.addWidget(clear_selection) |
| moves.addWidget(to_keep); moves.addWidget(to_reject); moves.addWidget(to_uncertain); moves.addStretch() |
| root.addLayout(moves) |
|
|
| buttons = QDialogButtonBox(QDialogButtonBox.Close) |
| self.close_button = buttons.button(QDialogButtonBox.Close) |
| self.apply_button = QPushButton("Apply EVE review") |
| self.apply_button.setProperty("primary", True); self.apply_button.setEnabled(False) |
| self.apply_button.clicked.connect(self._apply) |
| buttons.addButton(self.apply_button, QDialogButtonBox.AcceptRole) |
| buttons.rejected.connect(self.reject) |
| root.addWidget(buttons) |
|
|
| def _reference_panel(self, row: QHBoxLayout, title: str, button_text: str, callback) -> QListWidget: |
| frame = _card(); layout = QVBoxLayout(frame); layout.addWidget(_title(title)) |
| listing = QListWidget(); listing.setMaximumHeight(115) |
| button = QPushButton(button_text); button.clicked.connect(callback) |
| clear = QPushButton("Clear"); clear.clicked.connect(listing.clear) |
| actions = QHBoxLayout(); actions.addWidget(button); actions.addWidget(clear) |
| layout.addWidget(listing); layout.addLayout(actions); row.addWidget(frame, 1) |
| return listing |
|
|
| @staticmethod |
| def _reference_paths(listing: QListWidget) -> list[Path]: |
| return [Path(str(listing.item(index).data(Qt.UserRole))) for index in range(listing.count())] |
|
|
| def _add_references(self, listing: QListWidget) -> None: |
| selected, _ = QFileDialog.getOpenFileNames( |
| self, "Choose EVE reference images", self.dataset_path, |
| "Images (*.png *.jpg *.jpeg *.webp *.bmp *.gif)", |
| ) |
| existing = {str(path) for path in self._reference_paths(listing)} |
| for raw_path in selected: |
| path = str(Path(raw_path).resolve()) |
| if path in existing: |
| continue |
| item = QListWidgetItem(Path(path).name); item.setData(Qt.UserRole, path) |
| listing.addItem(item); existing.add(path) |
|
|
| def _add_good_references(self) -> None: |
| self._add_references(self.good_references) |
|
|
| def _add_bad_references(self) -> None: |
| self._add_references(self.bad_references) |
|
|
| def _analyze(self) -> None: |
| good = self._reference_paths(self.good_references) |
| if not good: |
| QMessageBox.information(self, "Good references required", "Add at least one good reference image for EVE.") |
| return |
| if self.reject_threshold.value() >= self.keep_threshold.value(): |
| QMessageBox.warning(self, "Check thresholds", "Reject confidence must be lower than Keep confidence.") |
| return |
| self.analyze_button.setEnabled(False); self.apply_button.setEnabled(False) |
| self.close_button.setEnabled(False) |
| self.progress.setValue(0); self.status.setText("EVE is starting. The first run may download its vision model once.") |
| self.worker = EveReviewWorker( |
| self.paths, good, self._reference_paths(self.bad_references), |
| self.keep_threshold.value() / 100, self.reject_threshold.value() / 100, |
| ) |
| self.worker.progress.connect(self._progress) |
| self.worker.completed.connect(self._completed) |
| self.worker.failed.connect(self._failed) |
| self.worker.finished.connect(self._worker_finished) |
| self.worker.start() |
|
|
| def _progress(self, percent: int, message: str) -> None: |
| self.progress.setValue(percent); self.status.setText(message) |
|
|
| def _completed(self, results: object) -> None: |
| if not isinstance(results, list): |
| self._failed("EVE returned an invalid review.") |
| return |
| self.results = results |
| save_eve_results(self.root_path, self.dataset_path, self.results) |
| self._rebuild_results() |
| self.apply_button.setEnabled(True) |
| counts = {key: sum(result.suggestion == key for result in self.results) for key in self.result_lists} |
| self.status.setText( |
| f"EVE proposes {counts['keep']} Keep, {counts['reject']} Reject, and {counts['unreviewed']} Uncertain. Review both sides before applying." |
| ) |
|
|
| def _failed(self, message: str) -> None: |
| self.status.setText(f"EVE could not finish: {message}") |
| QMessageBox.warning(self, "EVE review stopped", message) |
|
|
| def _worker_finished(self) -> None: |
| self.analyze_button.setEnabled(True) |
| self.close_button.setEnabled(True) |
| if self.worker: |
| self.worker.deleteLater() |
| self.worker = None |
|
|
| def _rebuild_results(self) -> None: |
| for listing in self.result_lists.values(): |
| listing.clear() |
| for result in self.results: |
| path = Path(result.path) |
| decision_score = result.match_score if result.suggestion == "keep" else 1.0 - result.match_score if result.suggestion == "reject" else result.match_score |
| text = f"{path.name}\n{decision_score * 100:.0f}% " + ("match" if result.suggestion != "reject" else "reject confidence") |
| item = QListWidgetItem(QIcon(_thumbnail(path, 120, 90)), text) |
| item.setData(Qt.UserRole, result.path) |
| self.result_lists[result.suggestion].addItem(item) |
| self._update_tab_labels() |
|
|
| def _update_tab_labels(self) -> None: |
| labels = {"keep": "KEEP", "reject": "REJECT", "unreviewed": "UNCERTAIN"} |
| for index, key in enumerate(("keep", "reject", "unreviewed")): |
| self.tabs.setTabText(index, f"{labels[key]} ({self.result_lists[key].count()})") |
|
|
| def _current_result_list(self) -> QListWidget: |
| return self.tabs.currentWidget() |
|
|
| def _select_all_current(self) -> None: |
| self._current_result_list().selectAll() |
|
|
| def _clear_current_selection(self) -> None: |
| self._current_result_list().clearSelection() |
|
|
| def _move_selected(self, decision: str) -> None: |
| source = self._current_result_list() |
| selected_items = source.selectedItems() |
| if not selected_items: |
| return |
| destination = self.result_lists[decision] |
| if source is destination: |
| return |
| selected_paths = {str(item.data(Qt.UserRole)) for item in selected_items} |
| for result in self.results: |
| if result.path in selected_paths: |
| result.suggestion = decision |
| result.decision_confidence = 1.0 |
| |
| |
| selected_rows = sorted((source.row(item) for item in selected_items), reverse=True) |
| moved_items = [source.takeItem(row) for row in selected_rows] |
| for item in reversed(moved_items): |
| destination.addItem(item) |
| item.setSelected(True) |
| self._update_tab_labels() |
|
|
| def _apply(self) -> None: |
| counts = {key: sum(result.suggestion == key for result in self.results) for key in self.result_lists} |
| answer = QMessageBox.question( |
| self, "Apply EVE review", |
| f"Apply {counts['keep']} Keep and {counts['reject']} Reject decisions?\n\n" |
| f"The {counts['unreviewed']} uncertain images will remain unreviewed. Rejected files are not moved until you choose Exclude rejected.", |
| ) |
| if answer != QMessageBox.Yes: |
| return |
| decisions = {result.path: result.suggestion for result in self.results} |
| changed = self.store.apply_decisions(self.dataset_path, decisions) |
| save_eve_results(self.root_path, self.dataset_path, self.results) |
| self.applied.emit(decisions) |
| self.status.setText(f"Applied EVE review ({changed} decisions changed). You can continue reviewing manually.") |
|
|
| def closeEvent(self, event) -> None: |
| if self.worker and self.worker.isRunning(): |
| self.worker.requestInterruption() |
| self.status.setText("EVE is stopping after the current image. The window will be safe to close when analysis ends.") |
| event.ignore() |
| return |
| super().closeEvent(event) |
|
|
|
|
| class DatasetReviewTab(QWidget): |
| def __init__(self, assets: AssetRegistry, store: StudioStore) -> None: |
| super().__init__() |
| self.assets = assets |
| self.store = store |
| self.paths: list[Path] = [] |
| self.dataset_path = "" |
| self._load_index = 0 |
| self._load_token = 0 |
| self._requested_row = 0 |
| self._scan_workers: set[ImageScanWorker] = set() |
|
|
| root = QVBoxLayout(self) |
| root.setContentsMargins(10, 14, 10, 10) |
| top = QVBoxLayout() |
| dataset_row = QHBoxLayout() |
| review_actions = QHBoxLayout() |
| self.dataset = QComboBox() |
| self.dataset.setMinimumWidth(300) |
| browse = QPushButton("Open another dataset…") |
| browse.clicked.connect(self._browse) |
| refresh = QPushButton("Refresh") |
| refresh.clicked.connect(self.refresh) |
| duplicates = QPushButton("Check duplicates") |
| duplicates.clicked.connect(self._duplicates) |
| keep_all = QPushButton("Keep all images") |
| keep_all.setProperty("primary", True) |
| keep_all.setToolTip( |
| "Accept every image in this dataset, then reject only the individual images you do not want." |
| ) |
| keep_all.clicked.connect(self._keep_all_images) |
| eve_review = QPushButton("EVE AI Review…") |
| eve_review.setToolTip( |
| "Sort this dataset from visual reference images, then review EVE's Keep, Reject, and Uncertain groups." |
| ) |
| eve_review.clicked.connect(self._open_eve_review) |
| apply_rejected = QPushButton("Exclude rejected") |
| apply_rejected.setToolTip( |
| "Move rejected images out of the training dataset into ADAM's recoverable quarantine." |
| ) |
| apply_rejected.clicked.connect(self._apply_rejected) |
| restore_rejected = QPushButton("Restore excluded") |
| restore_rejected.clicked.connect(self._restore_rejected) |
| dataset_row.addWidget(QLabel("Dataset")) |
| dataset_row.addWidget(self.dataset, 1) |
| dataset_row.addWidget(browse) |
| dataset_row.addWidget(refresh) |
| review_actions.addWidget(duplicates) |
| review_actions.addWidget(keep_all) |
| review_actions.addWidget(eve_review) |
| review_actions.addWidget(apply_rejected) |
| review_actions.addWidget(restore_rejected) |
| review_actions.addStretch() |
| top.addLayout(dataset_row) |
| top.addLayout(review_actions) |
| root.addLayout(top) |
|
|
| self.summary = QLabel("Choose a dataset to begin reviewing it.") |
| self.summary.setProperty("muted", True) |
| root.addWidget(self.summary) |
|
|
| body = QHBoxLayout() |
| self.gallery = QListWidget() |
| self.gallery.setViewMode(QListWidget.IconMode) |
| self.gallery.setIconSize(QPixmap(150, 110).size()) |
| self.gallery.setGridSize(QPixmap(178, 158).size()) |
| self.gallery.setResizeMode(QListWidget.Adjust) |
| self.gallery.setSelectionMode(QAbstractItemView.SingleSelection) |
| self.gallery.currentRowChanged.connect(self._selected) |
| body.addWidget(self.gallery, 3) |
|
|
| detail = _card() |
| detail.setMinimumWidth(330) |
| detail_layout = QVBoxLayout(detail) |
| detail_layout.addWidget(_title("IMAGE REVIEW")) |
| self.preview = QLabel("Select an image") |
| self.preview.setAlignment(Qt.AlignCenter) |
| self.preview.setMinimumHeight(230) |
| self.preview.setStyleSheet( |
| f"background: #050d14; border: 1px solid {COLORS['border']}; border-radius: 8px;" |
| ) |
| self.file_label = QLabel() |
| self.file_label.setWordWrap(True) |
| self.file_label.setProperty("muted", True) |
| self.caption = QPlainTextEdit() |
| self.caption.setPlaceholderText("Caption text stored beside the image…") |
| self.caption.setMaximumHeight(120) |
| detail_layout.addWidget(self.preview) |
| detail_layout.addWidget(self.file_label) |
| detail_layout.addWidget(QLabel("Caption")) |
| detail_layout.addWidget(self.caption) |
| buttons = QGridLayout() |
| keep = QPushButton("Keep") |
| keep.setProperty("primary", True) |
| reject = QPushButton("Reject") |
| reject.setProperty("danger", True) |
| restore = QPushButton("Mark unreviewed") |
| save_caption = QPushButton("Save caption") |
| keep.clicked.connect(lambda: self._decide("keep")) |
| reject.clicked.connect(lambda: self._decide("reject")) |
| restore.clicked.connect(lambda: self._decide("unreviewed")) |
| save_caption.clicked.connect(self._save_caption) |
| buttons.addWidget(keep, 0, 0) |
| buttons.addWidget(reject, 0, 1) |
| buttons.addWidget(restore, 1, 0) |
| buttons.addWidget(save_caption, 1, 1) |
| detail_layout.addLayout(buttons) |
| detail_layout.addStretch() |
| body.addWidget(detail, 2) |
| root.addLayout(body, 1) |
|
|
| self.dataset.currentIndexChanged.connect(self.refresh) |
| self.reload_assets() |
|
|
| def reload_assets(self) -> None: |
| current = self.dataset.currentData() |
| self.dataset.blockSignals(True) |
| self.dataset.clear() |
| for asset in self.assets.assets: |
| if asset.kind == "dataset" and Path(asset.path).is_dir(): |
| self.dataset.addItem(asset.name, asset.path) |
| self.dataset.blockSignals(False) |
| index = self.dataset.findData(current) |
| if index >= 0: |
| self.dataset.setCurrentIndex(index) |
| self.refresh() |
|
|
| def _browse(self) -> None: |
| selected = QFileDialog.getExistingDirectory( |
| self, "Choose an image dataset", self.dataset_path or str(Path.home()) |
| ) |
| if not selected: |
| return |
| index = self.dataset.findData(selected) |
| if index < 0: |
| self.dataset.addItem(Path(selected).name, selected) |
| index = self.dataset.count() - 1 |
| self.dataset.setCurrentIndex(index) |
|
|
| def refresh(self) -> None: |
| previous_path = self.dataset_path |
| self.dataset_path = str(self.dataset.currentData() or "") |
| if self.dataset_path != previous_path: |
| self._requested_row = 0 |
| self.gallery.clear() |
| self._load_token += 1 |
| token = self._load_token |
| self._load_index = 0 |
| for worker in self._scan_workers: |
| worker.requestInterruption() |
| if not self.dataset_path: |
| self.paths = [] |
| self.summary.setText("Choose a dataset to begin reviewing it.") |
| self.preview.setText("No supported images found") |
| return |
| self.summary.setText("Scanning dataset…") |
| worker = ImageScanWorker(self.dataset_path, token) |
| self._scan_workers.add(worker) |
| worker.scanned.connect(self._scan_finished) |
| worker.finished.connect( |
| lambda worker=worker: self._scan_workers.discard(worker) |
| ) |
| worker.finished.connect(worker.deleteLater) |
| worker.start() |
|
|
| def _scan_finished(self, paths: object, token: int) -> None: |
| if token != self._load_token or not isinstance(paths, list): |
| return |
| self.paths = paths |
| self._update_review_summary(loading=True) |
| if not self.paths: |
| self.preview.setText("No supported images found") |
| return |
| QTimer.singleShot(0, lambda: self._load_next_thumbnail(token)) |
|
|
| def _update_review_summary(self, *, loading: bool = False) -> None: |
| """Update counts without rescanning files or rebuilding thumbnails.""" |
| review = self.store.review(self.dataset_path) |
| kept = rejected = 0 |
| for path in self.paths: |
| decision = review.decisions.get(str(path.resolve()), "unreviewed") if review else "unreviewed" |
| kept += decision == "keep" |
| rejected += decision == "reject" |
| reviewed = kept + rejected |
| captions = sum(caption_path(path).is_file() for path in self.paths) |
| self.summary.setText( |
| f"{len(self.paths)} images · {reviewed} reviewed · {kept} kept · " |
| f"{rejected} rejected · {captions} captions" |
| + (" · Loading thumbnails…" if loading else "") |
| ) |
|
|
| @staticmethod |
| def _style_review_item(item: QListWidgetItem, path: Path, decision: str) -> None: |
| prefix = {"keep": "✓ ", "reject": "× ", "unreviewed": ""}[decision] |
| item.setText(prefix + path.name) |
| if decision == "reject": |
| item.setForeground(Qt.red) |
| elif decision == "keep": |
| item.setForeground(Qt.green) |
| else: |
| item.setForeground(QBrush()) |
|
|
| def _load_next_thumbnail(self, token: int) -> None: |
| """Decode one image per event-loop turn so large datasets stay responsive.""" |
| if token != self._load_token or self._load_index >= len(self.paths): |
| if token == self._load_token: |
| self.summary.setText(self.summary.text().replace(" · Loading thumbnails…", "")) |
| return |
| path = self.paths[self._load_index] |
| review = self.store.review(self.dataset_path) if self.dataset_path else None |
| decision = review.decisions.get(str(path.resolve()), "unreviewed") if review else "unreviewed" |
| item = QListWidgetItem(path.name) |
| item.setData(Qt.UserRole, str(path)) |
| pixmap = _thumbnail(path, 150, 110) |
| if not pixmap.isNull(): |
| item.setIcon(QIcon(pixmap)) |
| self._style_review_item(item, path, decision) |
| self.gallery.addItem(item) |
| target_row = min(self._requested_row, len(self.paths) - 1) |
| if self._load_index == target_row: |
| self.gallery.setCurrentRow(target_row) |
| self._load_index += 1 |
| QTimer.singleShot(0, lambda: self._load_next_thumbnail(token)) |
|
|
| def _current_path(self) -> Path | None: |
| item = self.gallery.currentItem() |
| return Path(str(item.data(Qt.UserRole))) if item else None |
|
|
| def _selected(self, _row: int) -> None: |
| path = self._current_path() |
| if not path: |
| return |
| pixmap = _thumbnail(path, 310, 260) |
| self.preview.setPixmap( |
| pixmap |
| ) |
| self.file_label.setText(str(path)) |
| try: |
| text = caption_path(path).read_text(encoding="utf-8") |
| except OSError: |
| text = "" |
| self.caption.setPlainText(text) |
|
|
| def _decide(self, decision: str) -> None: |
| path = self._current_path() |
| if not path or not self.dataset_path: |
| return |
| row = self.gallery.currentRow() |
| self.store.set_decision(self.dataset_path, str(path), decision) |
| item = self.gallery.item(row) |
| if item is not None: |
| self._style_review_item(item, path, decision) |
| self._update_review_summary() |
| if decision in {"keep", "reject"} and row + 1 < self.gallery.count(): |
| self.gallery.setCurrentRow(row + 1) |
| self.gallery.scrollToItem(self.gallery.currentItem()) |
|
|
| def _save_caption(self) -> None: |
| path = self._current_path() |
| if not path: |
| return |
| try: |
| caption_path(path).write_text( |
| self.caption.toPlainText().strip() + "\n", encoding="utf-8" |
| ) |
| except OSError as exc: |
| QMessageBox.warning(self, "Caption not saved", str(exc)) |
| return |
| self.summary.setText(self.summary.text() + " · Caption saved") |
|
|
| def _apply_rejected(self) -> None: |
| if not self.dataset_path: |
| return |
| review = self.store.review(self.dataset_path) |
| count = sum(value == "reject" for value in review.decisions.values()) |
| if not count: |
| QMessageBox.information( |
| self, "No rejected images", "Mark images as rejected before excluding them." |
| ) |
| return |
| answer = QMessageBox.question( |
| self, |
| "Exclude rejected images", |
| f"Move {count} rejected image(s) and their captions out of this training dataset?\n\n" |
| "They remain recoverable with Restore excluded.", |
| ) |
| if answer != QMessageBox.Yes: |
| return |
| moved = self.store.apply_rejections(self.dataset_path) |
| self._requested_row = 0 |
| self.refresh() |
| self.summary.setText( |
| f"Excluded {moved} rejected image(s) from training. They remain recoverable." |
| ) |
|
|
| def _keep_all_images(self) -> None: |
| if not self.dataset_path or not self.paths: |
| QMessageBox.information( |
| self, "No images", "Choose a dataset and wait for its images to finish loading." |
| ) |
| return |
| answer = QMessageBox.question( |
| self, |
| "Keep all images", |
| f"Mark all {len(self.paths)} images in this dataset as kept?\n\n" |
| "You can still reject individual images afterward.", |
| ) |
| if answer != QMessageBox.Yes: |
| return |
| changed = self.store.set_all_decisions( |
| self.dataset_path, self.paths, "keep" |
| ) |
| for row in range(self.gallery.count()): |
| item = self.gallery.item(row) |
| self._style_review_item(item, Path(str(item.data(Qt.UserRole))), "keep") |
| self._update_review_summary() |
| self.summary.setText( |
| self.summary.text() |
| + (f" · All images kept ({changed} changed)" if changed else " · All images already kept") |
| ) |
|
|
| def _open_eve_review(self) -> None: |
| if not self.dataset_path or not self.paths: |
| QMessageBox.information( |
| self, "No dataset ready", "Choose a dataset and wait for its image scan to finish." |
| ) |
| return |
| root_path = self.store.path.parent.parent |
| dialog = EveReviewDialog(root_path, self.dataset_path, self.paths, self.store, self) |
| dialog.applied.connect(self._eve_decisions_applied) |
| dialog.exec() |
|
|
| def _eve_decisions_applied(self, decisions: object) -> None: |
| if not isinstance(decisions, dict): |
| return |
| for row in range(self.gallery.count()): |
| item = self.gallery.item(row) |
| path = Path(str(item.data(Qt.UserRole))) |
| decision = str(decisions.get(str(path.resolve()), "unreviewed")) |
| self._style_review_item(item, path, decision) |
| self._update_review_summary() |
| self.summary.setText(self.summary.text() + " · EVE review applied") |
|
|
| def _restore_rejected(self) -> None: |
| if not self.dataset_path: |
| return |
| restored = self.store.restore_rejections(self.dataset_path) |
| self._requested_row = 0 |
| self.refresh() |
| self.summary.setText( |
| f"Restored {restored} excluded image(s)." |
| if restored |
| else "No excluded images were available to restore." |
| ) |
|
|
| def _duplicates(self) -> None: |
| groups = exact_duplicate_groups(self.paths) |
| exact_members = {value for group in groups for value in group} |
| hashes: list[tuple[Path, int]] = [] |
| for path in self.paths[:500]: |
| image = QImage(str(path)) |
| if image.isNull(): |
| continue |
| sample = image.convertToFormat(QImage.Format_Grayscale8).scaled( |
| 8, 8, Qt.IgnoreAspectRatio, Qt.SmoothTransformation |
| ) |
| values = [ |
| sample.pixelColor(x, y).red() for y in range(8) for x in range(8) |
| ] |
| average = sum(values) / len(values) |
| bits = 0 |
| for index, value in enumerate(values): |
| if value >= average: |
| bits |= 1 << index |
| hashes.append((path, bits)) |
| near: list[tuple[str, str]] = [] |
| for index, (first_path, first_hash) in enumerate(hashes): |
| for second_path, second_hash in hashes[index + 1 :]: |
| if str(first_path) in exact_members and str(second_path) in exact_members: |
| continue |
| if (first_hash ^ second_hash).bit_count() <= 5: |
| near.append((first_path.name, second_path.name)) |
| if len(near) >= 20: |
| break |
| if len(near) >= 20: |
| break |
| if not groups and not near: |
| QMessageBox.information( |
| self, |
| "Duplicate check", |
| "No exact or visually similar duplicate candidates were found.", |
| ) |
| return |
| lines = [ |
| " = ".join(Path(value).name for value in group) for group in groups[:20] |
| ] |
| lines.extend(f"≈ {first} / {second}" for first, second in near) |
| scope = ( |
| " Visual similarity checked the first 500 images." |
| if len(self.paths) > 500 |
| else "" |
| ) |
| QMessageBox.warning( |
| self, |
| "Duplicate candidates found", |
| f"{len(groups)} exact group(s) and {len(near)} visually similar " |
| f"candidate pair(s).{scope}\n\n" + "\n".join(lines), |
| ) |
|
|
|
|
| class ExperimentsTab(QWidget): |
| def __init__( |
| self, jobs: JobManager, assets: AssetRegistry, store: StudioStore |
| ) -> None: |
| super().__init__() |
| self.jobs = jobs |
| self.assets = assets |
| self.store = store |
| root = QVBoxLayout(self) |
| root.setContentsMargins(10, 14, 10, 10) |
| hint = QLabel( |
| "Select one run for details or two runs to compare their recipes and outcomes." |
| ) |
| hint.setProperty("muted", True) |
| root.addWidget(hint) |
| self.table = QTableWidget(0, 7) |
| self.table.setHorizontalHeaderLabels( |
| ["RUN", "PROJECT", "TRAINER", "EPOCHS", "STATUS", "CREATED", "OUTPUT"] |
| ) |
| self.table.setSelectionBehavior(QAbstractItemView.SelectRows) |
| self.table.setSelectionMode(QAbstractItemView.ExtendedSelection) |
| self.table.setEditTriggers(QAbstractItemView.NoEditTriggers) |
| self.table.verticalHeader().hide() |
| header = self.table.horizontalHeader() |
| header.setSectionResizeMode(1, QHeaderView.Stretch) |
| for column in (0, 2, 3, 4, 5, 6): |
| header.setSectionResizeMode(column, QHeaderView.ResizeToContents) |
| self.table.itemSelectionChanged.connect(self._show_selection) |
| root.addWidget(self.table, 3) |
| detail = _card() |
| detail_layout = QVBoxLayout(detail) |
| detail_layout.addWidget(_title("RUN COMPARISON")) |
| self.comparison = QPlainTextEdit() |
| self.comparison.setReadOnly(True) |
| self.comparison.setMaximumHeight(180) |
| detail_layout.addWidget(self.comparison) |
| actions = QHBoxLayout() |
| self.best = QPushButton("Mark model as best") |
| self.open = QPushButton("Open output") |
| self.recipe = QPushButton("Save as reusable recipe") |
| self.best.clicked.connect(self._toggle_best) |
| self.open.clicked.connect(self._open_output) |
| self.recipe.clicked.connect(self._save_recipe) |
| actions.addWidget(self.best) |
| actions.addWidget(self.recipe) |
| actions.addWidget(self.open) |
| detail_layout.addLayout(actions) |
| root.addWidget(detail, 1) |
| jobs.job_created.connect(lambda _job: self.refresh()) |
| jobs.job_updated.connect(lambda _job: self.refresh()) |
| self.refresh() |
|
|
| @staticmethod |
| def _training(job: Job) -> tuple[str, int]: |
| for step in job.plan.steps: |
| if step.tool_id.endswith("_trainer"): |
| return step.tool_id.removesuffix("_trainer").upper(), int( |
| step.arguments.get("epochs", 0) or 0 |
| ) |
| return "—", 0 |
|
|
| def refresh(self) -> None: |
| selected = {job.id for job in self._selected_jobs()} |
| self.table.setRowCount(len(self.jobs.jobs)) |
| for row, job in enumerate(self.jobs.jobs): |
| trainer, epochs = self._training(job) |
| values = [ |
| job.id, |
| job.plan.project_name, |
| trainer, |
| str(epochs or "—"), |
| job.status.value, |
| self._date(job.created_at), |
| "Ready" if job.output_folder else "—", |
| ] |
| for column, value in enumerate(values): |
| item = QTableWidgetItem(value) |
| if column != 1: |
| item.setTextAlignment(Qt.AlignCenter) |
| self.table.setItem(row, column, item) |
| if job.id in selected: |
| self.table.selectRow(row) |
| self._show_selection() |
|
|
| def _selected_jobs(self) -> list[Job]: |
| rows = sorted({index.row() for index in self.table.selectionModel().selectedRows()}) |
| result = [] |
| for row in rows: |
| item = self.table.item(row, 0) |
| if item: |
| try: |
| result.append(self.jobs.get(item.text())) |
| except KeyError: |
| pass |
| return result |
|
|
| def _show_selection(self) -> None: |
| jobs = self._selected_jobs() |
| if not jobs: |
| self.comparison.setPlainText("Select a run to inspect it.") |
| else: |
| blocks = [] |
| for job in jobs[:2]: |
| trainer, epochs = self._training(job) |
| duration = self._duration(job) |
| blocks.append( |
| f"{job.plan.project_name} ({job.id})\n" |
| f" Trainer: {trainer} · Epochs: {epochs or '—'}\n" |
| f" Status: {job.status.value} · Progress: {job.progress}%\n" |
| f" Duration: {duration} · Output: {job.output_folder or 'Not recorded'}" |
| ) |
| self.comparison.setPlainText("\n\n".join(blocks)) |
| one = len(jobs) == 1 |
| self.open.setEnabled(one and bool(jobs[0].output_folder)) |
| self.recipe.setEnabled(one and self._training(jobs[0])[0] != "—") |
| model = self._model_for_job(jobs[0]) if one else None |
| self.best.setEnabled(model is not None) |
| if model: |
| self.best.setText( |
| "Unmark best" if model.id in self.store.best_models else "Mark model as best" |
| ) |
|
|
| def _model_for_job(self, job: Job) -> Asset | None: |
| if not job.output_folder: |
| return None |
| output = Path(job.output_folder) |
| return next( |
| ( |
| asset |
| for asset in self.assets.assets |
| if asset.kind == "model" |
| and ( |
| Path(asset.path) == output |
| or output in Path(asset.path).parents |
| or Path(asset.path) in output.parents |
| ) |
| ), |
| None, |
| ) |
|
|
| def _toggle_best(self) -> None: |
| jobs = self._selected_jobs() |
| model = self._model_for_job(jobs[0]) if len(jobs) == 1 else None |
| if model: |
| self.store.toggle_best(model.id) |
| self._show_selection() |
|
|
| def _open_output(self) -> None: |
| jobs = self._selected_jobs() |
| if len(jobs) == 1 and jobs[0].output_folder: |
| QDesktopServices.openUrl(QUrl.fromLocalFile(jobs[0].output_folder)) |
|
|
| def _save_recipe(self) -> None: |
| jobs = self._selected_jobs() |
| if len(jobs) != 1: |
| return |
| job = jobs[0] |
| for step in job.plan.steps: |
| if step.tool_id.endswith("_trainer"): |
| self.store.add_recipe( |
| TrainingRecipe( |
| name=f"{job.plan.project_name} recipe", |
| trainer=step.tool_id.removesuffix("_trainer"), |
| epochs=int(step.arguments.get("epochs", 100) or 100), |
| base_model=str(step.arguments.get("base_model", "")), |
| notes=f"Created from job {job.id}", |
| ) |
| ) |
| QMessageBox.information( |
| self, "Recipe saved", "The run settings are now reusable in Recipes." |
| ) |
| return |
|
|
| @staticmethod |
| def _date(value: str) -> str: |
| try: |
| return datetime.fromisoformat(value).astimezone().strftime("%b %d · %H:%M") |
| except ValueError: |
| return value[:16] |
|
|
| @staticmethod |
| def _duration(job: Job) -> str: |
| if not job.started_at or not job.ended_at: |
| return "In progress" if job.started_at else "Not started" |
| try: |
| seconds = int( |
| ( |
| datetime.fromisoformat(job.ended_at) |
| - datetime.fromisoformat(job.started_at) |
| ).total_seconds() |
| ) |
| return f"{seconds // 3600}h {(seconds % 3600) // 60}m" if seconds >= 3600 else f"{seconds // 60}m {seconds % 60}s" |
| except ValueError: |
| return "Unknown" |
|
|
|
|
| class PreviewLabTab(QWidget): |
| plan_requested = Signal(str) |
|
|
| def __init__(self, assets: AssetRegistry, store: StudioStore) -> None: |
| super().__init__() |
| self.assets = assets |
| self.store = store |
| self._output_paths: list[Path] = [] |
| self._output_index = 0 |
| self._output_token = 0 |
| self._output_scan_workers: set[ImageScanWorker] = set() |
| root = QHBoxLayout(self) |
| root.setContentsMargins(10, 14, 10, 10) |
| form = _card() |
| form_layout = QGridLayout(form) |
| form_layout.addWidget(_title("PROMPT & CHECKPOINT EVALUATION"), 0, 0, 1, 2) |
| self.model = QComboBox() |
| self.checkpoint = QComboBox() |
| self.prompt = QPlainTextEdit() |
| self.prompt.setPlaceholderText( |
| "A consistent evaluation prompt for comparing model checkpoints…" |
| ) |
| self.prompt.setMaximumHeight(110) |
| self.seed = QSpinBox() |
| self.seed.setRange(0, 2_147_483_647) |
| self.count = QSpinBox() |
| self.count.setRange(1, 16) |
| self.count.setValue(4) |
| self.rating = QSpinBox() |
| self.rating.setRange(0, 5) |
| self.rating.setSuffix(" / 5") |
| self.notes = QPlainTextEdit() |
| self.notes.setPlaceholderText("What worked, what drifted, and what to try next…") |
| self.notes.setMaximumHeight(100) |
| rows = ( |
| ("Model", self.model), |
| ("Checkpoint", self.checkpoint), |
| ("Evaluation prompt", self.prompt), |
| ("Seed", self.seed), |
| ("Preview count", self.count), |
| ("Rating", self.rating), |
| ("Notes", self.notes), |
| ) |
| for row, (label, widget) in enumerate(rows, 1): |
| form_layout.addWidget(QLabel(label), row, 0, Qt.AlignTop) |
| form_layout.addWidget(widget, row, 1) |
| save = QPushButton("Save evaluation") |
| save.clicked.connect(self._save) |
| request = QPushButton("Plan preview job") |
| request.setProperty("primary", True) |
| request.clicked.connect(self._request) |
| form_layout.addWidget(save, len(rows) + 1, 0) |
| form_layout.addWidget(request, len(rows) + 1, 1) |
| root.addWidget(form, 2) |
|
|
| history = _card() |
| history_layout = QVBoxLayout(history) |
| history_layout.addWidget(_title("MODEL OUTPUTS")) |
| self.outputs = QListWidget() |
| self.outputs.setViewMode(QListWidget.IconMode) |
| self.outputs.setIconSize(QPixmap(110, 80).size()) |
| self.outputs.setMaximumHeight(210) |
| self.outputs.itemDoubleClicked.connect( |
| lambda item: QDesktopServices.openUrl( |
| QUrl.fromLocalFile(str(item.data(Qt.UserRole))) |
| ) |
| ) |
| history_layout.addWidget(self.outputs) |
| history_layout.addWidget(_title("EVALUATION HISTORY")) |
| self.history = QListWidget() |
| history_layout.addWidget(self.history) |
| root.addWidget(history, 1) |
| self.model.currentIndexChanged.connect(self._reload_checkpoints) |
| self.reload_assets() |
|
|
| def reload_assets(self) -> None: |
| current = self.model.currentData() |
| self.model.blockSignals(True) |
| self.model.clear() |
| for asset in self.assets.assets: |
| if asset.kind == "model": |
| star = "★ " if asset.id in self.store.best_models else "" |
| self.model.addItem(star + asset.name, asset.id) |
| self.model.blockSignals(False) |
| index = self.model.findData(current) |
| self.model.setCurrentIndex(index if index >= 0 else 0) |
| self._reload_checkpoints() |
| self._reload_history() |
|
|
| def _asset(self) -> Asset | None: |
| model_id = str(self.model.currentData() or "") |
| return next((asset for asset in self.assets.assets if asset.id == model_id), None) |
|
|
| def _reload_checkpoints(self) -> None: |
| self.checkpoint.clear() |
| self.outputs.clear() |
| self._output_token += 1 |
| token = self._output_token |
| self._output_index = 0 |
| for worker in self._output_scan_workers: |
| worker.requestInterruption() |
| asset = self._asset() |
| if not asset: |
| return |
| found = checkpoint_files(asset.path) |
| if asset.checkpoint and Path(asset.checkpoint).exists(): |
| found = [Path(asset.checkpoint), *[path for path in found if str(path) != asset.checkpoint]] |
| if not found: |
| self.checkpoint.addItem("Latest model output", asset.path) |
| else: |
| for path in found: |
| self.checkpoint.addItem(path.name, str(path)) |
| worker = ImageScanWorker(asset.path, token, limit=80) |
| self._output_scan_workers.add(worker) |
| worker.scanned.connect(self._preview_scan_finished) |
| worker.finished.connect( |
| lambda worker=worker: self._output_scan_workers.discard(worker) |
| ) |
| worker.finished.connect(worker.deleteLater) |
| worker.start() |
|
|
| def _preview_scan_finished(self, paths: object, token: int) -> None: |
| if token != self._output_token or not isinstance(paths, list): |
| return |
| self._output_paths = paths |
| QTimer.singleShot(0, lambda: self._load_next_preview(token)) |
|
|
| def _load_next_preview(self, token: int) -> None: |
| if token != self._output_token or self._output_index >= len(self._output_paths): |
| return |
| path = self._output_paths[self._output_index] |
| item = QListWidgetItem(path.name) |
| item.setData(Qt.UserRole, str(path)) |
| pixmap = _thumbnail(path, 110, 80) |
| if not pixmap.isNull(): |
| item.setIcon(QIcon(pixmap)) |
| self.outputs.addItem(item) |
| self._output_index += 1 |
| QTimer.singleShot(0, lambda: self._load_next_preview(token)) |
|
|
| def _save(self) -> None: |
| asset = self._asset() |
| if not asset: |
| QMessageBox.warning(self, "No model", "Register or finish a model first.") |
| return |
| self.store.add_evaluation( |
| PreviewEvaluation( |
| model_id=asset.id, |
| checkpoint=str(self.checkpoint.currentData() or ""), |
| prompt=self.prompt.toPlainText().strip(), |
| seed=self.seed.value(), |
| rating=self.rating.value(), |
| notes=self.notes.toPlainText().strip(), |
| ) |
| ) |
| self._reload_history() |
|
|
| def _reload_history(self) -> None: |
| self.history.clear() |
| names = {asset.id: asset.name for asset in self.assets.assets} |
| for evaluation in reversed(self.store.evaluations[-100:]): |
| self.history.addItem( |
| f"{'★' * evaluation.rating or 'Unrated'} · " |
| f"{names.get(evaluation.model_id, 'Unknown model')}\n" |
| f"{evaluation.prompt or 'No prompt recorded'}" |
| ) |
|
|
| def _request(self) -> None: |
| asset = self._asset() |
| if not asset: |
| return |
| prompt = self.prompt.toPlainText().strip() |
| checkpoint = str(self.checkpoint.currentData() or "") |
| request = ( |
| f"Generate {self.count.value()} previews for the {asset.name} model" |
| + (f" from checkpoint {checkpoint}" if checkpoint else "") |
| + (f" using this evaluation prompt: {prompt}" if prompt else "") |
| + f". Use seed {self.seed.value()}." |
| ) |
| self.plan_requested.emit(request) |
|
|
|
|
| class RecipesTab(QWidget): |
| plan_requested = Signal(str) |
|
|
| def __init__(self, store: StudioStore) -> None: |
| super().__init__() |
| self.store = store |
| root = QHBoxLayout(self) |
| root.setContentsMargins(10, 14, 10, 10) |
| self.list = QListWidget() |
| self.list.currentRowChanged.connect(self._selected) |
| root.addWidget(self.list, 1) |
| form = _card() |
| layout = QGridLayout(form) |
| layout.addWidget(_title("REPRODUCIBLE TRAINING RECIPE"), 0, 0, 1, 2) |
| self.name = QLineEdit() |
| self.trainer = QComboBox() |
| self.trainer.addItem("LoRA", "lora") |
| self.trainer.addItem("DDPM", "ddpm") |
| self.trainer.addItem("Flow Matching", "flow") |
| self.epochs = QSpinBox() |
| self.epochs.setRange(1, 100_000) |
| self.epochs.setValue(100) |
| self.images = QSpinBox() |
| self.images.setRange(10, 100_000) |
| self.images.setValue(60) |
| self.base_model = QLineEdit() |
| self.preview_prompt = QLineEdit() |
| self.notes = QPlainTextEdit() |
| self.notes.setMaximumHeight(100) |
| for row, (label, widget) in enumerate( |
| ( |
| ("Name", self.name), |
| ("Trainer", self.trainer), |
| ("Epochs", self.epochs), |
| ("Dataset target", self.images), |
| ("Base model", self.base_model), |
| ("Preview prompt", self.preview_prompt), |
| ("Notes", self.notes), |
| ), |
| 1, |
| ): |
| layout.addWidget(QLabel(label), row, 0) |
| layout.addWidget(widget, row, 1) |
| save = QPushButton("Save recipe") |
| save.clicked.connect(self._save) |
| use = QPushButton("Create model from recipe") |
| use.setProperty("primary", True) |
| use.clicked.connect(self._use) |
| layout.addWidget(save, 8, 0) |
| layout.addWidget(use, 8, 1) |
| transfer = QHBoxLayout() |
| export = QPushButton("Export recipe…") |
| import_button = QPushButton("Import recipe…") |
| export.clicked.connect(self._export) |
| import_button.clicked.connect(self._import) |
| transfer.addWidget(import_button) |
| transfer.addWidget(export) |
| layout.addLayout(transfer, 9, 0, 1, 2) |
| root.addWidget(form, 2) |
| self.refresh() |
|
|
| def refresh(self) -> None: |
| row = self.list.currentRow() |
| self.list.clear() |
| for recipe in self.store.recipes: |
| item = QListWidgetItem( |
| f"{recipe.name}\n{recipe.trainer.upper()} · {recipe.epochs} epochs" |
| ) |
| item.setData(Qt.UserRole, recipe.id) |
| self.list.addItem(item) |
| if self.list.count(): |
| self.list.setCurrentRow(max(0, min(row, self.list.count() - 1))) |
|
|
| def _current(self) -> TrainingRecipe | None: |
| item = self.list.currentItem() |
| recipe_id = str(item.data(Qt.UserRole)) if item else "" |
| return next((recipe for recipe in self.store.recipes if recipe.id == recipe_id), None) |
|
|
| def _selected(self, _row: int) -> None: |
| recipe = self._current() |
| if not recipe: |
| return |
| self.name.setText(recipe.name) |
| self.trainer.setCurrentIndex(max(0, self.trainer.findData(recipe.trainer))) |
| self.epochs.setValue(recipe.epochs) |
| self.images.setValue(recipe.image_count) |
| self.base_model.setText(recipe.base_model) |
| self.preview_prompt.setText(recipe.preview_prompt) |
| self.notes.setPlainText(recipe.notes) |
|
|
| def _save(self) -> None: |
| current = self._current() |
| recipe = TrainingRecipe( |
| id=current.id if current else uuid4().hex[:10], |
| created_at=current.created_at if current else datetime.now().astimezone().isoformat(), |
| name=self.name.text().strip() or "Untitled recipe", |
| trainer=str(self.trainer.currentData()), |
| epochs=self.epochs.value(), |
| image_count=self.images.value(), |
| base_model=self.base_model.text().strip(), |
| preview_prompt=self.preview_prompt.text().strip(), |
| notes=self.notes.toPlainText().strip(), |
| ) |
| self.store.add_recipe(recipe) |
| self.refresh() |
|
|
| def _use(self) -> None: |
| trainer = str(self.trainer.currentData()).upper() |
| self.plan_requested.emit( |
| f"Create a {trainer} model using {self.images.value()} images and train " |
| f"for {self.epochs.value()} epochs. Use the recipe named " |
| f"{self.name.text().strip() or 'Untitled recipe'}." |
| ) |
|
|
| def _export(self) -> None: |
| recipe = self._current() |
| if not recipe: |
| QMessageBox.information(self, "No recipe", "Select or save a recipe first.") |
| return |
| selected, _filter = QFileDialog.getSaveFileName( |
| self, |
| "Export ADAM recipe", |
| f"{recipe.name}.adam-recipe.json", |
| "ADAM recipes (*.json)", |
| ) |
| if selected: |
| try: |
| Path(selected).write_text( |
| json.dumps( |
| {"format": "adam-training-recipe-v1", "recipe": asdict(recipe)}, |
| indent=2, |
| ), |
| encoding="utf-8", |
| ) |
| except OSError as exc: |
| QMessageBox.warning(self, "Recipe not exported", str(exc)) |
|
|
| def _import(self) -> None: |
| selected, _filter = QFileDialog.getOpenFileName( |
| self, "Import ADAM recipe", "", "ADAM recipes (*.json)" |
| ) |
| if not selected: |
| return |
| try: |
| payload = json.loads(Path(selected).read_text(encoding="utf-8")) |
| if payload.get("format") != "adam-training-recipe-v1": |
| raise ValueError("This is not an ADAM training recipe.") |
| recipe = TrainingRecipe.from_dict(dict(payload["recipe"])) |
| recipe.id = uuid4().hex[:10] |
| self.store.add_recipe(recipe) |
| except (OSError, ValueError, TypeError, KeyError, json.JSONDecodeError) as exc: |
| QMessageBox.warning(self, "Recipe not imported", str(exc)) |
| return |
| self.refresh() |
|
|
|
|
| class StudioPage(QWidget): |
| plan_requested = Signal(str) |
|
|
| def __init__( |
| self, |
| root_path: Path, |
| jobs: JobManager, |
| assets: AssetRegistry, |
| config: ConfigManager, |
| ) -> None: |
| super().__init__() |
| del config |
| self.store = StudioStore(root_path) |
| root = QVBoxLayout(self) |
| root.setContentsMargins(24, 20, 24, 17) |
| root.setSpacing(8) |
| root.addWidget( |
| _header( |
| "Training studio", |
| "Review datasets, compare experiments, evaluate checkpoints, and preserve reproducible recipes.", |
| ) |
| ) |
| self.tabs = QTabWidget() |
| self.datasets = DatasetReviewTab(assets, self.store) |
| self.experiments = ExperimentsTab(jobs, assets, self.store) |
| self.previews = PreviewLabTab(assets, self.store) |
| self.recipes = RecipesTab(self.store) |
| self.tabs.addTab(self.datasets, "Datasets") |
| self.tabs.addTab(self.experiments, "Experiments") |
| self.tabs.addTab(self.previews, "Checkpoint Lab") |
| self.tabs.addTab(self.recipes, "Recipes") |
| root.addWidget(self.tabs, 1) |
| self.previews.plan_requested.connect(self.plan_requested) |
| self.recipes.plan_requested.connect(self.plan_requested) |
|
|
| def refresh(self) -> None: |
| self.datasets.assets.load() |
| self.datasets.reload_assets() |
| self.experiments.assets.load() |
| self.experiments.refresh() |
| self.previews.assets.load() |
| self.previews.reload_assets() |
| self.recipes.refresh() |
|
|
| def shutdown(self) -> None: |
| workers = [ |
| *self.datasets._scan_workers, |
| *self.previews._output_scan_workers, |
| ] |
| for worker in workers: |
| worker.requestInterruption() |
| for worker in workers: |
| worker.wait(1000) |
|
|