"""Configure the restoration ("расцензурить") engine.
Two kinds of engine:
• **DeepMosaics** (картинка/видео) — vendored, in-process; *reconstructs* mosaic and
locates it itself. Pick a clean model from the bundled ``models/deepmosaics`` weights
(or browse). ``mosaic_position.pth`` must sit beside it. CUDA GPU strongly recommended.
• **Diffusion-inpaint (SwarmUI)** — *regenerates* the detected (masked) regions via an
external SwarmUI server over HTTP. Needs YOLO detections for the mask and a running
SwarmUI; the model runs in SwarmUI's process (no torch here). Per-frame (best for
stills — a video sequence will flicker).
The per-engine fields live in two group widgets that are shown/hidden by the engine combo.
"""
from __future__ import annotations
from pathlib import Path
from PySide6.QtCore import Qt
from PySide6.QtWidgets import (
QApplication,
QCheckBox,
QComboBox,
QDialog,
QDialogButtonBox,
QDoubleSpinBox,
QFileDialog,
QFormLayout,
QHBoxLayout,
QLabel,
QLineEdit,
QMessageBox,
QPushButton,
QSpinBox,
QVBoxLayout,
QWidget,
)
from ..config import AppConfig
from ..core.restore.deepmosaics import discover_models
class RestoreDialog(QDialog):
def __init__(self, config: AppConfig, parent: QWidget | None = None) -> None:
super().__init__(parent)
self._cfg = config
self.setWindowTitle("Движок восстановления")
self.setMinimumWidth(580)
self.engine = QComboBox()
self.engine.addItem("DeepMosaics — картинка (покадрово, нужна модель+GPU)", "deepmosaics")
self.engine.addItem("DeepMosaics — видео (соседние кадры, лучше для роликов)", "deepmosaics_video")
self.engine.addItem("Diffusion-inpaint (SwarmUI) — перерисовка по маске YOLO", "diffusion")
self.engine.setCurrentIndex(max(0, self.engine.findData(config.restorer)))
self.engine.currentIndexChanged.connect(self._sync)
root = QVBoxLayout(self)
top = QFormLayout()
top.addRow("Движок:", self.engine)
root.addLayout(top)
root.addWidget(self._build_deepmosaics_group(config))
root.addWidget(self._build_diffusion_group(config))
buttons = QDialogButtonBox(QDialogButtonBox.Ok | QDialogButtonBox.Cancel)
buttons.accepted.connect(self.accept)
buttons.rejected.connect(self.reject)
root.addWidget(buttons)
self._sync()
# ----------------------------------------------------------- DeepMosaics UI
def _build_deepmosaics_group(self, config: AppConfig) -> QWidget:
self.dm_group = QWidget()
form = QFormLayout(self.dm_group)
form.setContentsMargins(0, 0, 0, 0)
self.model_combo = QComboBox()
self._populate_models(config.dm_model)
self.dm_gpu = QLineEdit(config.dm_gpu or "0")
self.dm_gpu.setPlaceholderText("0 = первая CUDA-карта, -1 = CPU (медленно)")
self.feed_restored = QCheckBox(
"Подавать уже расцензуренные прошлые кадры в окно (эксперим.)"
)
self.feed_restored.setChecked(bool(getattr(config, "dm_feed_restored", True)))
self.feed_restored.setToolTip(
"Только для видеодвижка: прошлые соседние кадры в окне берутся из уже\n"
"восстановленных результатов, а не из оригинала с мозаикой — больше\n"
"временной связности. Сеть обучалась на мозаичных окнах, так что эффект\n"
"не гарантирован; выключите для точной реализации DeepMosaics."
)
form.addRow("Модель:", self._with_browse(self.model_combo, self._browse_model))
form.addRow("GPU id:", self.dm_gpu)
form.addRow("", self.feed_restored)
hint = QLabel(
"Модели берутся из models/deepmosaics (рядом нужен mosaic_position.pth).\n"
"• Картинка: clean_youknow_resnet_9blocks.pth — покадрово.\n"
"• Видео: clean_youknow_video.pth — использует соседние кадры (когерентнее на "
"роликах), требует прогона по диапазону («Расцензурить все»).\n"
"На CPU медленно — лучше GPU. См. README."
)
hint.setWordWrap(True)
form.addRow(hint)
return self.dm_group
# ------------------------------------------------------------ Diffusion UI
def _build_diffusion_group(self, config: AppConfig) -> QWidget:
self.diff_group = QWidget()
form = QFormLayout(self.diff_group)
form.setContentsMargins(0, 0, 0, 0)
self.diff_url = QLineEdit(config.diff_url or "http://localhost:7801")
self.diff_url.setPlaceholderText("http://localhost:7801")
self.diff_test_btn = QPushButton("Проверить соединение")
self.diff_test_btn.clicked.connect(self._test_connection)
self.diff_test_status = QLabel("")
self.diff_test_status.setWordWrap(True)
self.diff_model = QLineEdit(config.diff_model or "")
self.diff_model.setPlaceholderText("имя чекпойнта в SwarmUI (пусто = текущий)")
self.diff_prompt = QLineEdit(config.diff_prompt or "")
self.diff_prompt.setPlaceholderText("что нарисовать в области под цензурой")
self.diff_negative = QLineEdit(config.diff_negative or "")
self.diff_negative.setPlaceholderText("чего избегать (negative prompt)")
self.diff_steps = QSpinBox()
self.diff_steps.setRange(1, 150)
self.diff_steps.setValue(int(config.diff_steps))
self.diff_cfg = QDoubleSpinBox()
self.diff_cfg.setRange(0.0, 30.0)
self.diff_cfg.setSingleStep(0.5)
self.diff_cfg.setValue(float(config.diff_cfg))
self.diff_denoise = QDoubleSpinBox()
self.diff_denoise.setRange(0.0, 1.0)
self.diff_denoise.setSingleStep(0.05)
self.diff_denoise.setValue(float(config.diff_denoise))
self.diff_denoise.setToolTip("0..1 — насколько перерисовать область (1 = полностью)")
self.diff_seed = QSpinBox()
self.diff_seed.setRange(-1, 2_147_483_647)
self.diff_seed.setValue(int(config.diff_seed))
self.diff_seed.setSpecialValueText("случайный") # at -1
self.diff_dilate = QSpinBox()
self.diff_dilate.setRange(0, 200)
self.diff_dilate.setValue(int(config.diff_mask_dilate))
self.diff_dilate.setToolTip("Расширить маску на N px (закрыть край цензуры)")
self.diff_blur = QSpinBox()
self.diff_blur.setRange(0, 200)
self.diff_blur.setValue(int(config.diff_mask_blur))
self.diff_blur.setToolTip("Размытие края маски, px (мягкий стык)")
form.addRow("SwarmUI URL:", self.diff_url)
form.addRow("", self.diff_test_btn)
form.addRow("", self.diff_test_status)
form.addRow("Чекпойнт:", self.diff_model)
form.addRow("Промпт:", self.diff_prompt)
form.addRow("Negative:", self.diff_negative)
steps_row = QWidget()
h = QHBoxLayout(steps_row)
h.setContentsMargins(0, 0, 0, 0)
h.addWidget(QLabel("Шаги:"))
h.addWidget(self.diff_steps)
h.addWidget(QLabel("CFG:"))
h.addWidget(self.diff_cfg)
h.addWidget(QLabel("Denoise:"))
h.addWidget(self.diff_denoise)
h.addWidget(QLabel("Seed:"))
h.addWidget(self.diff_seed)
form.addRow("", steps_row)
mask_row = QWidget()
hm = QHBoxLayout(mask_row)
hm.setContentsMargins(0, 0, 0, 0)
hm.addWidget(QLabel("Маска: расширить, px:"))
hm.addWidget(self.diff_dilate)
hm.addWidget(QLabel("размытие, px:"))
hm.addWidget(self.diff_blur)
form.addRow("", mask_row)
hint = QLabel(
"Diffusion перерисовывает область ЦЕНЗУРЫ заново (не восстанавливает оригинал),\n"
"опираясь на маску из детекций YOLO и промпт. Нужен запущенный сервер SwarmUI\n"
"и посчитанная детекция. Покадрово — на роликах будет мерцание. Лучше для\n"
"чёрных плашек/заливки, где DeepMosaics бессилен."
)
hint.setWordWrap(True)
form.addRow(hint)
return self.diff_group
# --------------------------------------------------------------- helpers
def _populate_models(self, current: str | None) -> None:
self.model_combo.clear()
is_video = self.engine.currentData() == "deepmosaics_video"
if is_video: # temporal engine: only the video weights (clean_*_video.pth)
models = [(n, p) for n, p in discover_models(include_video=True) if "video" in n.lower()]
else:
models = discover_models() # per-frame engine: image clean models only
for name, path in models:
self.model_combo.addItem(name, path)
if current and self.model_combo.findData(current) < 0:
self.model_combo.addItem(Path(current).stem + " (внешняя)", current)
if self.model_combo.count() == 0:
self.model_combo.addItem("(модели не найдены — положите в models/deepmosaics)", None)
idx = self.model_combo.findData(current) if current else 0
self.model_combo.setCurrentIndex(max(0, idx))
def _with_browse(self, widget: QWidget, slot) -> QWidget:
w = QWidget()
h = QHBoxLayout(w)
h.setContentsMargins(0, 0, 0, 0)
h.addWidget(widget, 1)
btn = QPushButton("Обзор…")
btn.clicked.connect(slot)
h.addWidget(btn)
return w
def _sync(self) -> None:
engine = self.engine.currentData()
is_dm = engine in ("deepmosaics", "deepmosaics_video")
is_video = engine == "deepmosaics_video"
is_diff = engine == "diffusion"
if is_dm: # repopulate model list (image vs video weights differ)
self._populate_models(self._cfg.dm_model)
self.feed_restored.setEnabled(is_video)
self.dm_group.setVisible(is_dm)
self.diff_group.setVisible(is_diff)
self.adjustSize()
def _browse_model(self) -> None:
p, _ = QFileDialog.getOpenFileName(self, "Веса DeepMosaics", "", "Веса (*.pth);;Все файлы (*.*)")
if p:
if self.model_combo.findData(p) < 0:
self.model_combo.addItem(Path(p).stem, p)
self.model_combo.setCurrentIndex(self.model_combo.findData(p))
def _test_connection(self) -> None:
"""Ping SwarmUI (GetNewSession) with the current URL and report OK / the error."""
from ..core.restore.swarmui import SwarmUIBackend
url = self.diff_url.text().strip() or "http://localhost:7801"
self.diff_test_status.setText("Проверка…")
self.diff_test_btn.setEnabled(False)
QApplication.setOverrideCursor(Qt.WaitCursor)
QApplication.processEvents()
try:
session = SwarmUIBackend(url, timeout=15.0).ping() # short timeout for the probe
except Exception as e: # noqa: BLE001 — show the server/connection error verbatim
self.diff_test_status.setText(f"✗ {e}")
QMessageBox.warning(self, "SwarmUI: соединение", str(e))
else:
self.diff_test_status.setText(
f"✓ Соединение OK (session: {session})"
)
finally:
QApplication.restoreOverrideCursor()
self.diff_test_btn.setEnabled(True)
def apply_to_config(self) -> None:
self._cfg.restorer = self.engine.currentData()
# DeepMosaics
self._cfg.dm_model = self.model_combo.currentData()
self._cfg.dm_gpu = self.dm_gpu.text().strip() or "0"
self._cfg.dm_feed_restored = self.feed_restored.isChecked()
# Diffusion (SwarmUI)
self._cfg.diff_backend = "swarmui"
self._cfg.diff_url = self.diff_url.text().strip() or "http://localhost:7801"
self._cfg.diff_model = self.diff_model.text().strip() or None
self._cfg.diff_prompt = self.diff_prompt.text()
self._cfg.diff_negative = self.diff_negative.text()
self._cfg.diff_steps = self.diff_steps.value()
self._cfg.diff_cfg = self.diff_cfg.value()
self._cfg.diff_denoise = self.diff_denoise.value()
self._cfg.diff_seed = self.diff_seed.value()
self._cfg.diff_mask_dilate = self.diff_dilate.value()
self._cfg.diff_mask_blur = self.diff_blur.value()