/
Nemo_499
/
DecodingX-rayImages
Обзор
Документация
Войти
/
Nemo_499
/
DecodingX-rayImages
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
Interface.py
921 строка
40 KB
Вдовин Денис
vgg16
19 май 2026, 18:42
19 май 2026, 18:42
73ddac7
Код
Авторство
О чём код?
from PyQt6.QtWidgets import ( QMainWindow, QWidget, QTabWidget, QGridLayout, QVBoxLayout, QLabel, QLineEdit, QTextEdit, QPushButton, QFileDialog, QListWidget, QScrollArea, QSizePolicy, QMessageBox, QRadioButton, QButtonGroup, QHBoxLayout, QStackedWidget, QComboBox, QCheckBox ) from PyQt6.QtCore import Qt, QObject, QThread, pyqtSignal from PyQt6.QtGui import QPixmap, QIntValidator import pyqtgraph as pg from CImageLabel import ImageLabel from AIEffDet import train_efficientdet, inference_single_image import os import traceback from matplotlib.backends.backend_qt5agg import FigureCanvasQTAgg as FigureCanvas from matplotlib.figure import Figure import numpy as np import cv2 from Config_set import names as NamesClass from Config_set import set_objects class TrainingWorker(QObject): progress = pyqtSignal(int, float, float, float, float, str, bool, object, object) log = pyqtSignal(str) finished = pyqtSignal(str) error = pyqtSignal(str) save_stats_signal = pyqtSignal(str,str) def __init__(self, training_parameters): super().__init__() self.training_parameters = training_parameters def run(self): try: dataset_path = self.training_parameters["dataset_path"] epochs = self.training_parameters["epochs"] model_save_name = self.training_parameters["model_save_name"] self.log.emit(f"=== Начало обучения ===") self.log.emit(f"Датасет: {dataset_path}") self.log.emit(f"Количество эпох: {epochs}") self.log.emit(f"Имя модели: {model_save_name}") # Обучение модели с вычислением точности trained_model_path = train_efficientdet( training_parameters=self.training_parameters, callback=self.progress.emit, # Передаёт callback для прогресса save_stats_signal=self.save_stats_signal.emit ) print(trained_model_path) if trained_model_path and os.path.exists(trained_model_path): self.log.emit(f"Обучение завершено!") self.log.emit(f"Модель сохранена в: {trained_model_path}") self.finished.emit(trained_model_path) else: self.log.emit("Обучение завершено, но модель не была сохранена") self.error.emit("Модель не была сохранена") except Exception as e: error_msg = f"Критическая ошибка при обучении: {str(e)}" self.log.emit(error_msg) self.log.emit(traceback.format_exc()) self.error.emit(error_msg) class MainWindow(QMainWindow): def __init__(self): super().__init__() self.dataset_path = None self.current_image_path = None self.model_path = None self.results = [] # Графики self.epochs_x = [] self.train_loss_y = [] self.val_loss_y = [] self.train_acc_y = [] self.val_acc_y = [] self.count_epochs = 0 self.training_mode = True self.show_results_mode = True self.index_model_box = 0 self.index_objeсt_type_box = 0 self.index_objeсt_GradCAM_box = -1 self.NUM_CLASSES = 13 # Инициализация NUM_CLASSES для всех объектов по умолчанию self.setWindowTitle("Decoding X-ray Images") self.resize(1200, 700) self.tabs = QTabWidget() self.setCentralWidget(self.tabs) self.tabs.addTab(self.create_page_processing_image(), "Анализ изображений") self.tabs.addTab(self.create_page_training(), "Обучение модели") self.setStyleSheet(""" QPushButton { background-color: #4CAF50; color: white; border: none; padding: 10px; font-size: 14px; border-radius: 5px; } QPushButton:hover { background-color: #45a049; } QPushButton:pressed { background-color: #3d8b40; } QLineEdit, QTextEdit { border: 1px solid #ccc; border-radius: 3px; padding: 5px; } QLabel { font-size: 12px; } QPushButton:disabled { background-color: #b0b0b0; color: #666666; border: 1px solid #999999; } """) # ========================= # PAGE: IMAGE # ========================= def create_page_processing_image(self): page = QWidget() grid = QGridLayout(page) grid.setColumnStretch(0, 1) grid.setColumnStretch(1, 4) # Кнопки loading_image_button = QPushButton("Загрузить изображение") loading_image_button.setFixedHeight(45) self.decrypt_image_button = QPushButton("Расшифровать изображение") self.decrypt_image_button.setFixedHeight(45) self.show_results_button = QPushButton("Отобразить результат") self.show_results_button.setFixedHeight(45) select_model_button = QPushButton("Выбрать модель") select_model_button.setFixedHeight(45) self.btn_select_roi = QPushButton("Начать выбор зоны шва") self.btn_select_roi.setFixedHeight(45) self.btn_clear_roi = QPushButton("Сбросить выбор зоны") self.btn_select_roi.setFixedHeight(45) self.btn_clear_roi.setDisabled(True) # Индикатор текущей модели self.model_label = QLabel("Модель: не выбрана") self.model_label.setAlignment(Qt.AlignmentFlag.AlignCenter) self.model_label.setStyleSheet(""" QLabel { background-color: #f0f0f0; padding: 8px; border: 1px solid #ccc; border-radius: 3px; font-weight: bold; } """) # Список обнаруженных объектов self.ListWidget = QListWidget() self.ListWidget.setMouseTracking(True) self.ListWidget.entered.connect(lambda index: self.image_label.highlight_bbox(index.row())) self.ListWidget.setStyleSheet(""" QListWidget { border: 1px solid #ccc; border-radius: 3px; font-size: 12px; } QListWidget::item:hover { background-color: #e0e0e0; } """) self.object_type_box_GradCAM = QComboBox() self.object_type_box_GradCAM.setFixedHeight(40) self.object_type_box_GradCAM.addItems(["Никакие"]) self.object_type_box_GradCAM.addItems(["Поры"]) self.object_type_box_GradCAM.addItems(["Включения"]) self.object_type_box_GradCAM.addItems(["Подрезы"]) self.object_type_box_GradCAM.addItems(["Прожоги"]) self.object_type_box_GradCAM.addItems(["Утяжина"]) self.object_type_box_GradCAM.addItems(["Несплавление"]) self.object_type_box_GradCAM.addItems(["Непровар корня"]) # Вёрстка левой панели left = QWidget() left_layout = QVBoxLayout(left) left_layout.addWidget(loading_image_button) left_layout.addWidget(select_model_button) left_layout.addWidget(self.model_label) left_layout.addWidget(QLabel("Инструменты ручного выбора:")) #left_layout.addWidget(self.btn_select_roi) #left_layout.addWidget(self.btn_clear_roi) left_layout.addWidget(self.decrypt_image_button) left_layout.addWidget(self.show_results_button) left_layout.addWidget(QLabel("Обнаруженные объекты:")) left_layout.addWidget(self.ListWidget) # Вёрстка правой панели self.scroll_area = QScrollArea() self.scroll_area.setWidgetResizable(True) self.scroll_area.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Expanding) self.scroll_area.setStyleSheet(""" QScrollArea { border: 1px solid #ccc; border-radius: 3px; } """) self.image_label = ImageLabel() self.scroll_area.setWidget(self.image_label) # Вёрстка страницы grid.addWidget(left, 0, 0) grid.addWidget(self.scroll_area, 0, 1) # События loading_image_button.clicked.connect(self.load_image) select_model_button.clicked.connect(self.select_model_file_to_use) self.decrypt_image_button.clicked.connect(self.run_inference) self.show_results_button.clicked.connect(self.show_results) self.ListWidget.itemEntered.connect(self.on_item_hover) return page # ========================= # PAGE: TRAINING # ========================= def create_page_training(self): page = QWidget() grid = QGridLayout(page) grid.setColumnStretch(0, 1) grid.setColumnStretch(1, 4) # Вёрстка левой панели self.left_stack = QStackedWidget() self.left_panel_button_general = self.create_panel_button_general() self.left_panel_button_interpretation = self.create_panel_button_interpretation() self.left_stack.addWidget(self.left_panel_button_general) self.left_stack.addWidget(self.left_panel_button_interpretation) self.stats_output = QTextEdit() self.stats_output.setReadOnly(True) self.stats_output.setMaximumHeight(150) self.stats_output.setStyleSheet(""" QTextEdit { background-color: #f8f8f8; font-family: monospace; font-size: 11px; } """) # Вкладки данных обучения self.create_tab_charts() self.create_tab_matrix() #self.create_tab_interpretation() self.tabs_training = QTabWidget() self.tabs_training.addTab(self.tab_charts,"Графики обучения") self.tabs_training.addTab(self.tab_matrix,"Матрица ошибок") #self.tabs_training.addTab(self.tab_interpretation,"Интерпретация") # Вёрстка правой панели right = QWidget() right_grid = QGridLayout(right) right_grid.addWidget(self.tabs_training, 0, 0) right_grid.addWidget(QLabel("Лог обучения:"), 1, 0, 1, 2) right_grid.addWidget(self.stats_output, 2, 0, 1, 2) # Вёрстка страницы grid.addWidget(self.left_stack, 0, 0) grid.addWidget(right, 0, 1) # События self.tabs_training.currentChanged.connect(self.on_training_tab_changed) return page def create_tab_charts(self): self.tab_charts = QWidget() # Графики error_graph = pg.PlotWidget(title="График ошибки") error_graph.setLabel('bottom', 'Эпоха', units=None) error_graph.setLabel('left', 'Потери', units=None) error_graph.addLegend() error_graph.setBackground('w') error_graph.showGrid(x=True, y=True) accuracy_graph = pg.PlotWidget(title="График точности") accuracy_graph.setLabel('bottom', 'Эпоха', units=None) accuracy_graph.setLabel('left', 'Точность', units=None) accuracy_graph.addLegend() accuracy_graph.setBackground('w') accuracy_graph.showGrid(x=True, y=True) self.loss_train_curve = error_graph.plot(pen=pg.mkPen(color='b', width=2), name='Ошибка обучения') self.loss_val_curve = error_graph.plot(pen=pg.mkPen(color='orange', width=2), name='Ошибка валидации') self.acc_train_curve = accuracy_graph.plot(pen=pg.mkPen(color='b', width=2), name='Точность обучения') self.acc_val_curve = accuracy_graph.plot(pen=pg.mkPen(color='orange', width=2), name='Точность валидации') right_grid = QGridLayout(self.tab_charts) right_grid.addWidget(error_graph, 0, 0) right_grid.addWidget(accuracy_graph, 0, 1) def create_tab_matrix(self): self.tab_matrix = QWidget() layout = QVBoxLayout(self.tab_matrix) # layout='constrained' автоматически управляет отступами, чтобы подписи не накладывались self.fig = Figure(figsize=(8, 8), dpi=100, layout='constrained') self.canvas = FigureCanvas(self.fig) # Позволяет холсту расширяться во все стороны self.canvas.setSizePolicy( QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Expanding ) # Добавляем без alignment, чтобы виджет заполнял весь layout layout.addWidget(self.canvas) cm = np.zeros((self.NUM_CLASSES + 1, self.NUM_CLASSES + 1)) self.draw_confusion_matrix(cm, NamesClass[self.index_objeсt_type_box], 0, flag_val=True, flag_save = False) def create_tab_interpretation(self): self.tab_interpretation = QWidget() def create_panel_button_general(self): panel_button_general = QWidget() # Радио кнопки режима обучения radio_train_new = QRadioButton("С нуля") radio_train_further_education = QRadioButton("Дообучение") radio_train_new.setChecked(True) # левая = True по умолчанию # Группа radio_group = QButtonGroup(self) radio_group.addButton(radio_train_new, 1) radio_group.addButton(radio_train_further_education, 0) # Вёрстка радиокнопок radio_layout = QHBoxLayout() radio_layout.addWidget(radio_train_new) radio_layout.addWidget(radio_train_further_education) radio_layout.addStretch() # Кнопки loading_dataset_button = QPushButton("Загрузить датасет") loading_dataset_button.setFixedHeight(45) self.launching_training_button = QPushButton("Запуск обучения") self.launching_training_button.setFixedHeight(45) self.loading_weights_AI = QPushButton("Загрузить веса ИИ") self.loading_weights_AI.setFixedHeight(45) self.loading_weights_AI.setDisabled(True) # Поля ввода числа эпох self.number_epochs = QLineEdit() self.number_epochs.setFixedHeight(45) self.number_epochs.setPlaceholderText("Количество эпох") self.number_epochs.setValidator(QIntValidator(1, 10000)) # Поле для пути сохранения модели self.model_save_name = QLineEdit() self.model_save_name.setFixedHeight(40) self.model_save_name.setPlaceholderText("Путь для сохранения модели (pth)") self.model_save_name.setText("efficientdet_xray") # Комбобоксы для выбора модели и определяемых объектов self.confid_model_box = QComboBox() self.confid_model_box.setFixedHeight(40) self.confid_model_box.addItems(["EfficientDet_lite0"]) self.confid_model_box.addItems(["EfficientDet_lite1"]) self.confid_model_box.addItems(["EfficientDet_d0"]) self.confid_model_box.addItems(["EfficientDet_d0_ap"]) self.confid_model_box.addItems(["EfficientDet_d1"]) self.confid_model_box.addItems(["EfficientDet_d1_ap"]) self.confid_model_box.addItems(["EfficientDet_d2"]) self.confid_model_box.addItems(["EfficientDet_d2_ap"]) self.object_type_box = QComboBox() self.object_type_box.setFixedHeight(40) self.object_type_box.addItems(["Все объекты"]) self.object_type_box.addItems(["Пора, Включение, Подрез"]) self.object_type_box.addItems(["Прожог, Трещина"]) self.object_type_box.addItems(["Эталоны"]) self.object_type_box.addItems(["Утяжина, Несплавление, Непровар корня"]) self.clahe_flag = QCheckBox("Использование CLAHE") panel_button_general_layout = QVBoxLayout(panel_button_general) panel_button_general_layout.addWidget(QLabel("Выберите модель")) panel_button_general_layout.addWidget(self.confid_model_box) panel_button_general_layout.addWidget(QLabel("Выберите набор объектов")) panel_button_general_layout.addWidget(self.object_type_box) panel_button_general_layout.addWidget(QLabel("Выберите режим обучения ИИ")) panel_button_general_layout.addLayout(radio_layout) panel_button_general_layout.addWidget(self.clahe_flag) panel_button_general_layout.addWidget(QLabel("Количество эпох:")) panel_button_general_layout.addWidget(self.number_epochs) panel_button_general_layout.addWidget(self.loading_weights_AI) panel_button_general_layout.addWidget(QLabel("Имя модели:")) panel_button_general_layout.addWidget(self.model_save_name) panel_button_general_layout.addWidget(loading_dataset_button) panel_button_general_layout.addWidget(self.launching_training_button) panel_button_general_layout.addStretch() # События loading_dataset_button.clicked.connect(self.select_dataset_folder) self.launching_training_button.clicked.connect( lambda: self.start_training() ) radio_group.idClicked.connect(self.on_training_mode_changed) self.loading_weights_AI.clicked.connect(self.select_model_file_to_train) self.object_type_box.currentIndexChanged.connect(self.on_selected_object_type) self.confid_model_box.currentIndexChanged.connect(self.on_selected_config_model) return panel_button_general def create_panel_button_interpretation(self): panel_button_interpretation = QWidget() # Кнопки button_forward = QPushButton("Следующий") button_forward.setFixedHeight(45) button_back = QPushButton("Предыдущий") button_back.setFixedHeight(45) button_switch_view = QPushButton("Переключить вид") button_switch_view.setFixedHeight(45) # Вёрстка """, Qt.AlignmentFlag.AlignTop""" grid = QGridLayout(panel_button_interpretation) grid.addWidget(button_forward, 0, 1, Qt.AlignmentFlag.AlignTop) grid.addWidget(button_back, 0, 0, Qt.AlignmentFlag.AlignTop) grid.addWidget(button_switch_view, 1, 0, 1, 2, Qt.AlignmentFlag.AlignTop) grid.setRowStretch(0, 0) # Первая строка не растягивается (0) grid.setRowStretch(1, 1) # Вторая строка растягивается (если есть) return panel_button_interpretation # ========================= # Логика # ========================= def load_image(self): path, _ = QFileDialog.getOpenFileName( self, "Выбрать изображение", "", "Images (*.png *.jpg *.jpeg *.bmp *.tiff)" ) if path: self.current_image_path = path pixmap = QPixmap(path) if pixmap.isNull(): QMessageBox.critical(self, "Ошибка", "Не удалось загрузить изображение") else: self.image_label.set_image(pixmap) self.stats_output.append(f"[Изображение] Загружено: {os.path.basename(path)}") def select_dataset_folder(self): self.dataset_path = QFileDialog.getExistingDirectory( self, "Выберите папку с датасетом" ) if self.dataset_path: # Проверка структуры датасета required_folders = ["train/images", "train/labels", "val/images", "val/labels"] missing = [] for folder in required_folders: if not os.path.exists(os.path.join(self.dataset_path, folder)): missing.append(folder) if missing: QMessageBox.warning(self, "Предупреждение", f"В датасете отсутствуют некоторые папки:\n" + "\n".join(missing) + f"\n\nТребуемая структура:\n{self.dataset_path}/\n train/\n images/\n labels/\n val/\n images/\n labels/") else: self.stats_output.append(f"[Датасет] Выбран: {self.dataset_path}") QMessageBox.information(self, "Успех", "Датасет загружен успешно!") def select_model_file_to_use(self): if(self.select_model_file()): buf = self.model_label.text().replace("Модель: ","").replace(".pth","") if("_A_" in buf): self.index_objeсt_type_box = 1 elif("_E_" in buf): self.index_objeсt_type_box = 2 elif("_D_" in buf): self.index_objeсt_type_box = 4 else: self.index_objeсt_type_box = 0 def select_model_file_to_train(self): if(self.select_model_file()): # Сброс графиков перед новым обучением self.epochs_x = [] self.train_loss_y = [] self.val_loss_y = [] self.train_acc_y = [] self.val_acc_y = [] self.count_epochs = 0 # Очистка кривых на графике self.loss_train_curve.setData([], []) self.loss_val_curve.setData([], []) self.acc_train_curve.setData([], []) self.acc_val_curve.setData([], []) dir_path = os.path.dirname(self.model_path) data_epochs_path = os.path.join(dir_path, self.model_save_name.text() + ".txt") with open(data_epochs_path, 'r', encoding='utf-8') as file: for line in file: epoch, train_loss, val_loss, train_acc, val_acc = line.strip().split('\t') self.epochs_x.append(int(epoch)) self.train_loss_y.append(float(train_loss)) self.val_loss_y.append(float(val_loss)) self.train_acc_y.append(float(train_acc)) self.val_acc_y.append(float(val_acc)) self.count_epochs+=1 self.loss_train_curve.setData(self.epochs_x, self.train_loss_y) self.loss_val_curve.setData(self.epochs_x, self.val_loss_y) self.acc_train_curve.setData(self.epochs_x, self.train_acc_y) self.acc_val_curve.setData(self.epochs_x, self.val_acc_y) def select_model_file(self): """Выбор файла модели""" path, _ = QFileDialog.getOpenFileName( self, "Выбрать модель", "", "Models (*.pth)" ) if path: buf_name = (os.path.basename(path)).replace(".pth","") self.model_path = path self.model_label.setText(f"Модель: {buf_name}") self.model_save_name.setText(buf_name) self.stats_output.append(f"[Модель] Выбрана: {os.path.basename(path)}") return True return False def start_training(self): if not self.training_mode and not self.model_path: QMessageBox.warning(self, "Ошибка", "Выберите файл с весами ИИ") return if not self.number_epochs.text(): QMessageBox.warning(self, "Ошибка", "Введите количество эпох") return if self.model_save_name.text() == "": QMessageBox.warning(self, "Ошибка", "Введите имя модели") return if not self.dataset_path: QMessageBox.warning(self, "Ошибка", "Сначала выберите датасет") return if(self.training_mode): # Сброс графиков перед новым обучением self.epochs_x = [] self.train_loss_y = [] self.val_loss_y = [] self.train_acc_y = [] self.val_acc_y = [] self.count_epochs = 0 # Очистка кривых на графике self.loss_train_curve.setData([], []) self.loss_val_curve.setData([], []) self.acc_train_curve.setData([], []) self.acc_val_curve.setData([], []) self.epochs = int(self.number_epochs.text()) training_parameters = { "model_type": self.index_model_box, "object_type": self.index_objeсt_type_box, "dataset_path": self.dataset_path, "epochs": self.epochs, "model_save_name": self.model_save_name.text(), "training_mode": self.training_mode, "count_epochs": self.count_epochs, "model_path": self.model_path, "clahe":self.clahe_flag.isChecked() } self.thread = QThread() self.worker = TrainingWorker(training_parameters) self.worker.moveToThread(self.thread) self.thread.started.connect(self.worker.run) self.worker.progress.connect(self.handle_progress) self.worker.log.connect(self.stats_output.append) self.worker.finished.connect(self.on_training_finished) self.worker.error.connect(self.on_training_error) self.worker.save_stats_signal.connect(self.write_data) # Очистка при завершении self.worker.finished.connect(self.thread.quit) self.worker.finished.connect(self.worker.deleteLater) self.thread.finished.connect(self.thread.deleteLater) self.thread.start() # Блокировка кнопки на время обучения if self.launching_training_button: self.launching_training_button.setEnabled(False) self.launching_training_button.setText("Обучение...") self.confid_model_box.setEnabled(False) self.object_type_box.setEnabled(False) self.thread.finished.connect(lambda: self.launching_training_button.setEnabled(True)) self.thread.finished.connect(lambda: self.launching_training_button.setText("Запуск обучения")) self.thread.finished.connect(lambda: self.confid_model_box.setEnabled(True)) self.thread.finished.connect(lambda: self.object_type_box.setEnabled(True)) def on_training_error(self, error_msg): QMessageBox.critical(self, "Ошибка обучения", error_msg) self.launching_training_button.setEnabled(True) self.launching_training_button.setText("Запуск обучения") self.confid_model_box.setEnabled(True) self.object_type_box.setEnabled(True) def on_training_finished(self, model_path): """Обработка завершения обучения""" if model_path and os.path.exists(model_path): self.model_path = model_path self.model_label.setText(f"Модель: {os.path.basename(self.model_path)}") QMessageBox.information( self, "Обучение завершено", f"Модель успешно обучена и сохранена:\n{model_path}\n\n" ) self.count_epochs = 0 self.count_epochs += len(self.epochs_x) self.write_data(self.model_path, self.model_save_name.text()) else: QMessageBox.warning( self, "Обучение завершено", "Обучение завершено, но модель не была сохранена." ) def run_inference(self): if not hasattr(self, "current_image_path") or not self.current_image_path: QMessageBox.warning(self, "Ошибка", "Сначала загрузите изображение") return if not os.path.exists(self.model_path): reply = QMessageBox.question( self, "Модель не найдена", f"Модель не найдена по пути:\n{self.model_path}\n\n" f"Хотите выбрать другую модель?", QMessageBox.StandardButton.Yes | QMessageBox.StandardButton.No ) if reply == QMessageBox.StandardButton.Yes: self.select_model_file() return try: self.results = [] self.ListWidget.clear() # Отображение индикатора загрузки self.stats_output.append("[Анализ] Начинаю обработку изображения...") self.results = inference_single_image( model_path=self.model_path, image_path=self.current_image_path, confidence_threshold=0.3, index_objeсt_type_box = self.index_objeсt_type_box ) if self.results: self.stats_output.append(f"[Анализ] Найдено объектов: {len(self.results)}") QMessageBox.information( self, "Анализ завершен", f"Найдено объектов: {len(self.results)}\n" f"Нажмите 'Отобразить результат' для визуализации." ) else: self.stats_output.append("[Анализ] Объекты не обнаружены") QMessageBox.information( self, "Анализ завершен", "Объекты не обнаружены на изображении." ) except Exception as e: self.stats_output.append(f"[Ошибка] Не удалось выполнить анализ: {str(e)}") QMessageBox.critical( self, "Ошибка анализа", f"Не удалось выполнить анализ:\n{str(e)}" ) def show_results(self): if not hasattr(self, "results") or not self.results: QMessageBox.warning(self, "Ошибка", "Нет результатов для отображения") return if (self.show_results_mode): # Установка bounding boxes для отображения self.image_label.set_bboxes(self.results) # Очистка и заполнение списка self.ListWidget.clear() for i, obj in enumerate(self.results): self.ListWidget.addItem( f"Объект {i + 1} | Класс: {obj['label']} | Уверенность={obj['score']:.2f}" ) self.stats_output.append(f"[Визуализация] Отображено {len(self.results)} объектов") self.show_results_mode = False self.show_results_button.setText("Скрыть результаты") else: self.image_label.clear_boxes() self.show_results_mode = True self.show_results_button.setText("Отобразить результаты") def on_item_hover(self, item): # Обработка наведения на элемент списка if item is not None: index = self.ListWidget.row(item) self.image_label.highlight_bbox(index) def handle_progress(self, epoch, train_loss, val_loss, train_acc, val_acc, stats_text, flag, matrix_val, matrix_train): self.stats_output.append(stats_text) if(flag): self.epochs_x.append(epoch + 1 + self.count_epochs) self.train_loss_y.append(train_loss) self.val_loss_y.append(val_loss) self.train_acc_y.append(train_acc) self.val_acc_y.append(val_acc) self.loss_train_curve.setData(self.epochs_x, self.train_loss_y) self.loss_val_curve.setData(self.epochs_x, self.val_loss_y) self.acc_train_curve.setData(self.epochs_x, self.train_acc_y) self.acc_val_curve.setData(self.epochs_x, self.val_acc_y) if matrix_train is not None: self.draw_confusion_matrix(matrix_train, NamesClass[self.index_objeсt_type_box], epoch + 1 + self.count_epochs, flag_val=False, flag_save=True) if matrix_val is not None: self.draw_confusion_matrix(matrix_val, NamesClass[self.index_objeсt_type_box], epoch + 1 + self.count_epochs, flag_val=True, flag_save=True) def on_training_mode_changed(self, id_: int): self.training_mode = bool(id_) self.loading_weights_AI.setDisabled(self.training_mode) print("training_mode =", self.training_mode) def on_training_tab_changed(self, index): """Обработчик смены вкладки в tabs_training""" if index == 2: # Индекс вкладки "Интерпретация" self.left_stack.setCurrentIndex(1) # Показываем панель интерпретации else: # Для остальных вкладок (Графики и Матрица ошибок) self.left_stack.setCurrentIndex(0) # Показываем стандартную панель обучения def draw_confusion_matrix(self, matrix, class_names, epoch, flag_val, flag_save = True): self.fig.clf() ax = self.fig.add_subplot(111) # Отрисовка тепловой карты if matrix.max() == 0: # Если матрица нулевая, используем фиксированный диапазон cax = ax.matshow(matrix, cmap='Blues', vmin=0, vmax=1) else: cax = ax.matshow(matrix, cmap='Blues') self.fig.colorbar(cax, ax=ax) # Полный список имен (классы + Background) all_names = class_names + ["Background"] indices = np.arange(len(all_names)) # 1. Создание точек на осях ax.set_xticks(indices) ax.set_yticks(indices) # 2. Установка подписей с поворотом # rotation=45 позволит уместить длинные названия "Непровар корня" и т.д. # ha='left' выравнивает текст по левому краю для верхней оси matshow ax.set_xticklabels(all_names, rotation=45, ha='left') ax.set_yticklabels(all_names) thresh = matrix.max() / 2. for i in range(matrix.shape[0]): # Строки (Реальные) for j in range(matrix.shape[1]): # Столбцы (Предсказанные) ax.text(j, i, f"{int(matrix[i, j])}", ha="center", va="center", color="white" if matrix[i, j] > thresh else "black", fontsize=8) # Предсказанные сверху в matshow, но xlabel ставит подпись снизу ax.set_xlabel('ПРЕДСКАЗАНЫЕ', fontsize=10, fontweight='bold', labelpad=15) ax.set_ylabel('РЕАЛЬНЫЕ', fontsize=10, fontweight='bold', labelpad=15) buf_pref="val" if flag_val else "train" if (flag_save): # 1. Формирование пути к папке folder_path = os.path.join("Interpretations", self.model_save_name.text(), "matrix") # 2. Создание всех папок, если их нет os.makedirs(folder_path, exist_ok=True) # 3. Формирование полного пути к файлу с расширением file_path = os.path.join(folder_path, f"epoch_{buf_pref}_{epoch}.png") # 4. Сохранение self.fig.savefig(file_path, dpi=300, bbox_inches='tight') # 4. Обновление холст а в Qt if flag_val: self.canvas.draw_idle() def write_data(self, model_path, model_save_name): dir_path = os.path.dirname(model_path) data_epochs_path = os.path.join(dir_path, model_save_name + ".txt") with open(data_epochs_path, 'w', encoding='utf-8') as file: for epoch, train_loss, val_loss, train_acc, val_acc in zip(self.epochs_x, self.train_loss_y, self.val_loss_y, self.train_acc_y, self.val_acc_y): file.write(f"{epoch}\t{train_loss}\t{val_loss}\t{train_acc}\t{val_acc}\n") def on_selected_config_model(self, index): """Обработчик выбора элемента""" # Получить индекс выбранного элемента self.index_model_box = index def on_selected_object_type(self, index): """Обработчик выбора элемента""" # Получить индекс выбранного элемента self.index_objeсt_type_box = index self.NUM_CLASSES = set_objects[self.index_objeсt_type_box][0] # Обновляем матрицу ошибок на вкладке "Матрица ошибок" cm = np.zeros((self.NUM_CLASSES + 1, self.NUM_CLASSES + 1)) self.draw_confusion_matrix(cm, NamesClass[self.index_objeсt_type_box], 0, flag_val=True, flag_save = False) def filling_object_type_box_GradCAM(self, index): if(index == 0): self.object_type_box_GradCAM.clear() self.object_type_box_GradCAM.addItems(["Никакие"]) self.object_type_box_GradCAM.addItems(["Поры"]) self.object_type_box_GradCAM.addItems(["Включения"]) self.object_type_box_GradCAM.addItems(["Подрезы"]) self.object_type_box_GradCAM.addItems(["Прожоги"]) self.object_type_box_GradCAM.addItems(["Эталон 1"]) self.object_type_box_GradCAM.addItems(["Эталон 2"]) self.object_type_box_GradCAM.addItems(["Эталон 3"]) self.object_type_box_GradCAM.addItems(["Утяжина"]) self.object_type_box_GradCAM.addItems(["Несплавление"]) self.object_type_box_GradCAM.addItems(["Непровар корня"]) elif(index == 1): self.object_type_box_GradCAM.clear() self.object_type_box_GradCAM.addItems(["Поры"]) self.object_type_box_GradCAM.addItems(["Включения"]) self.object_type_box_GradCAM.addItems(["Подрезы"]) self.object_type_box_GradCAM.addItems(["Прожоги"]) elif(index == 3): self.object_type_box_GradCAM.clear() self.object_type_box_GradCAM.addItems(["Утяжина"]) self.object_type_box_GradCAM.addItems(["Несплавление"]) self.object_type_box_GradCAM.addItems(["Непровар корня"])