/
Arthur159
/
chat_bot_BLS
Обзор
Документация
Войти
/
Arthur159
/
chat_bot_BLS
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
ready_model.py
80 строк
4 KB
Arthur159
upload files
02 июн 2025, 16:35
02 июн 2025, 16:35
b4c2ea2
Код
Авторство
О чём код?
import torch import torch.nn as nn from torchvision import transforms from PIL import Image import os def ready_model(): # Определение модели CNN class CNN(nn.Module): def __init__(self): super(CNN, self).__init__() self.conv1 = nn.Conv2d(1, 32, kernel_size=3, padding=1) # Один канал для градаций серого self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1) self.fc1 = nn.Linear(64 * 7 * 7, 128) self.fc2 = nn.Linear(128, 10) def forward(self, x): x = nn.functional.relu(self.conv1(x)) x = nn.functional.max_pool2d(x, 2) x = nn.functional.relu(self.conv2(x)) x = nn.functional.max_pool2d(x, 2) x = x.view(-1, 64 * 7 * 7) x = nn.functional.relu(self.fc1(x)) x = self.fc2(x) return x # Загрузка модели model = CNN() model.load_state_dict(torch.load('model.pth')) model.eval() # Переключаем модель в режим оценки # Предобработка данных для предсказания transform = transforms.Compose([ transforms.Grayscale(), # Преобразуем в градации серого transforms.Resize((28, 28)), # Изменение размера до 28x28 transforms.ToTensor(), # Преобразование в тензор transforms.Normalize((0.5,), (0.5,)) # Нормализация ]) # Функция для предсказания цифр на изображениях def predict_digits(image_paths): digits = [] for image_path in image_paths: image = Image.open(image_path) # Загрузка изображения с использованием PIL image = transform(image).unsqueeze(0) # Применение трансформации и добавление размерности батча with torch.no_grad(): # Отключаем градиенты для экономии памяти output = model(image) # Прямой проход _, predicted = torch.max(output.data, 1) # Получаем предсказанный класс digits.append(predicted.item()) # Добавляем предсказанную цифру в список return digits # Загрузка изображений из директории base_dir = 'output_images' # Укажите путь к вашей директории с поддиректориями result_numbers = [] # Список для хранения результатов # Проход по поддиректориям for subdir in sorted(os.listdir(base_dir)): subdir_path = os.path.join(base_dir, subdir) if os.path.isdir(subdir_path): # Проверяем, является ли это директорией image_files = sorted(os.listdir(subdir_path)) # Сортируем файлы для последовательной обработки image_paths = [os.path.join(subdir_path, img) for img in image_files if img.endswith('.png')] # Убедитесь, что это изображения PNG if len(image_paths) == 3: # Проверяем, есть ли три изображения predicted_digits = predict_digits(image_paths) # Предсказание цифр на изображениях result_number = ''.join(map(str, predicted_digits)) # Объединение предсказанных цифр в одно число result_numbers.append(result_number) # Добавляем в список else: result_numbers.append('') # Если не три изображения, добавляем пустое значение # Вывод результатов print(f'Предсказанные числа: {result_numbers}') return result_numbers