Работа телеграм-бота с нейросетью с несколькими людьми

Я пишу своего первого бота (на aiogram), в которого я встроил style-transfer нейросеть на pytorch. Проблема у меня такая: нейросеть какое-то время (секунд 40-50) работает, а затем отправляет ответ. Если в это время боту напишет пользователь, он, конечно, проигнорирует его. Посоветуйте, как лучше оформить прием фото от пользователей (может, их id в очередь записывать и доставать оттуда, или как-то еще). Спасибо.

Запускаемый файл:

import config
from io import BytesIO
from aiogram import Bot, types
from aiogram.utils import executor
import torchvision.models as models
from aiogram.dispatcher import Dispatcher
from aiogram.types.message import ParseMode
from users import create_user_checker, users
from net import run_style_transfer, unloader, download_cnn
from aiogram.types import InlineKeyboardMarkup, InlineKeyboardButton


print('Bot is starting..')
cnn = download_cnn()
bot = Bot(token=config.TOKEN)
dp = Dispatcher(bot)
print('Bot has been started')


@dp.message_handler(commands=['start', 'help'])
async def welcome(message):
    try:
        create_user_checker(message.from_user.id)

        inline_keyboard = types.InlineKeyboardMarkup(row_width=2)
        item1 = InlineKeyboardButton("Давай!", callback_data='yes')
        item2 = InlineKeyboardButton("Чуть позже", callback_data='no')

        inline_keyboard.add(item1, item2)
        me = await bot.get_me()

        await bot.send_message(message.chat.id, f'Приветствую тебя, *{message.from_user.first_name}*! '
                                                f'Я *{me.first_name}* — бот, созданный, чтобы '
                                                f'переносить стиль одних фотографий на другие. '
                                                f'Начнем?',
                                                parse_mode=ParseMode.MARKDOWN,
                                                reply_markup=inline_keyboard)
    except Exception as e:
        await message.reply(message, "Ошибка: "+repr(e))


@dp.message_handler(commands=['transfer_style'])
async def transfer(message):
    try:

        create_user_checker(message.from_user.id)

        inline_keyboard = types.InlineKeyboardMarkup(row_width=2)
        item1 = types.InlineKeyboardButton("Давай!", callback_data='yes')
        item2 = types.InlineKeyboardButton("Чуть позже.", callback_data='no')

        inline_keyboard.add(item1, item2)

        await bot.send_message(message.chat.id, f'Итак, начнем?',
                                                parse_mode='Markdown',
                                                reply_markup=inline_keyboard)

    except Exception as e:
        await message.reply(message, "Ошибка: "+repr(e))


@dp.message_handler(content_types=['photo'])
async def get_photo(message):
    # try:
    create_user_checker(message.from_user.id)

    if users[message.from_user.id].is_getting_photos:

        print(users[message.from_user.id].counter)

        if message.media_group_id:
            users[message.from_user.id].group_counter += 1

            if users[message.from_user.id].group_counter == 1:
                await bot.send_message(message.chat.id, 'Пожалуйста, отправляйте только по одному фото в сообщении')

            return

        try:
            photo = message.photo[-1]
            print(type(photo))
            photo_id = message.photo[-1].file_id
            photo_width = message.photo[-1].width
            photo_height = message.photo[-1].height

            file = await bot.get_file(photo_id)

        except IndexError:
            await bot.send_message(message.chat.id, 'Ошибка. Попробуйте отправить то же фото еще раз')
            return

        if users[message.from_user.id].counter == 0:

            users[message.from_user.id].style_photo_size = (photo_width, photo_height)
            await photo.download(f'images/{message.from_user.id}' + '_style_photo.pickle')

            # with open(f'images/{message.from_user.id}' + '_style_photo.pickle', 'wb') as file:
            #     file.write(users[message.from_user.id].style_photo)

            await bot.send_message(message.chat.id, 'Отлично, теперь отправьте фото контента')

        if users[message.from_user.id].counter == 1:
            users[message.from_user.id].content_photo_size = (photo_width, photo_height)
            await photo.download(f'images/{message.from_user.id}' + '_content_photo.pickle')

            # with open(f'images/{message.from_user.id}' + '_content_photo.pickle', 'wb') as file:
            #     file.write(users[message.from_user.id].content_photo)

            users[message.from_user.id].is_getting_photos = False

            await bot.send_message(message.chat.id, 'Фото получил, начинаю работу!')
            await transfer_style(message)

        users[message.from_user.id].counter += 1

    elif message.media_group_id and users[message.from_user.id].not_instructed_counter < 1:
        users[message.from_user.id].not_instructed_counter += 1
        await welcome(message)

    # except Exception as e:
    #     bot.reply_to(message, "Ошибка: "+repr(e))


@dp.callback_query_handler(lambda call: True)
async def callback_inline(call):
    create_user_checker(call.from_user.id)

    try:
        if call.message:
            if call.data == 'yes':
                users[call.from_user.id].is_getting_photos = True

                await bot.send_message(call.message.chat.id, 'Отлично, тогда отправь мне сначала фото стиля, '
                                                             'а затем фото контента. Жду!')
            elif call.data == 'no':
                await bot.send_message(call.message.chat.id, 'Хорошо, пиши, как понадоблюсь!')

            await bot.edit_message_reply_markup(chat_id=call.message.chat.id,
                                                message_id=call.message.message_id,
                                                reply_markup=None)

    except Exception as e:
        await message.reply(call.message, "Ошибка: "+repr(e))


async def transfer_style(message):
    output = run_style_transfer(cnn,
                                f'images/{message.from_user.id}',
                                users[message.from_user.id].style_photo_size,
                                users[message.from_user.id].content_photo_size,
                                )

    output = unloader(output)

    bio = BytesIO()
    bio.name = f'images/{message.from_user.id}+_result.png'
    output.save(bio, 'PNG')
    bio.seek(0)

    await bot.send_photo(message.chat.id, bio, 'Вот, что у меня получилось')

    inline_keyboard = types.InlineKeyboardMarkup(row_width=2)
    item1 = InlineKeyboardButton("Давай!", callback_data='yes')
    item2 = InlineKeyboardButton("Чуть позже", callback_data='no')
    inline_keyboard.add(item1, item2)

    await bot.send_message(message.chat.id, f'Хочешь перенести стиль еще на одну фотографию?',
                           reply_markup=inline_keyboard)

if __name__ == "__main__":
    executor.start_polling(dp)

Файл с нейросетью:

import os
import copy
import torch
import config
import warnings
import torch.nn as nn
from PIL import Image
import torch.optim as optim
import torch.nn.functional as F
import torchvision.models as models
import torchvision.transforms as transforms


warnings.filterwarnings("ignore")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
imsize = config.imsize

cnn_normalization_mean = torch.tensor([0.485, 0.456, 0.406]).to(device)
cnn_normalization_std = torch.tensor([0.229, 0.224, 0.225]).to(device)

unloader = transforms.ToPILImage()
loader = transforms.Compose([
        transforms.Resize(imsize),
        transforms.ToTensor()])


def download_cnn():
    cnn = models.vgg19(pretrained=True).features.to(device).eval()
    return cnn


def image_loader(name, size):
    # image = Image.fromstring('RGB', size, byte_image)
    image = Image.open(name)
    image = loader(image).unsqueeze(0)
    print('photo readed ok')
    os.remove(name)
    return image.to(device, torch.float)


def gram_matrix(input):
    a, b, c, d = input.size()
    features = input.view(a * b, c * d)
    G = torch.mm(features, features.t())

    return G.div(a * b * c * d)


class ContentLoss(nn.Module):
    def __init__(self, target):
        super(ContentLoss, self).__init__()
        self.target = target.detach()
        self.loss = 0

    def forward(self, input):
        self.loss = F.mse_loss(input, self.target)
        return input


class StyleLoss(nn.Module):
    def __init__(self, target_feature):
        super(StyleLoss, self).__init__()
        self.target = gram_matrix(target_feature).detach()

    def forward(self, input):
        G = gram_matrix(input)
        self.loss = F.mse_loss(G, self.target)
        return input


class Normalization(nn.Module):
    def __init__(self, mean, std):
        super(Normalization, self).__init__()
        self.mean = torch.tensor(mean).view(-1, 1, 1)
        self.std = torch.tensor(std).view(-1, 1, 1)

    def forward(self, img):
        return (img - self.mean) / self.std


def get_style_model_and_losses(cnn,
                               style_img, content_img,
                               normalization_mean=cnn_normalization_mean,
                               normalization_std=cnn_normalization_std,
                               content_layers=config.content_layers_default,
                               style_layers=config.style_layers_default):

    cnn = copy.deepcopy(cnn)

    # нормализуем картинки
    normalization = Normalization(normalization_mean, normalization_std).to(device)

    content_losses = []
    style_losses = []

    # Начнем собирать нашу модель
    model = nn.Sequential(normalization)

    i = 0
    for layer in cnn.children():

        if isinstance(layer, nn.Conv2d):
            i += 1
            name = 'conv_{}'.format(i)

        elif isinstance(layer, nn.ReLU):
            name = 'relu_{}'.format(i)
            layer = nn.ReLU(inplace=False)

        elif isinstance(layer, nn.MaxPool2d):
            name = 'pool_{}'.format(i)

        elif isinstance(layer, nn.BatchNorm2d):
            name = 'bn_{}'.format(i)

        else:
            raise RuntimeError('Unrecognized layer: {}'.format(layer.__class__.__name__))

        model.add_module(name, layer)

        if name in content_layers:
            # добавляем в модель слой, считающий ошибку контента
            target = model(content_img).detach()
            content_loss = ContentLoss(target)
            model.add_module("content_loss_{}".format(i), content_loss)
            content_losses.append(content_loss)

        if name in style_layers:
            # добавляем в модель слой, считающий ошибку стиля
            target_feature = model(style_img).detach()
            style_loss = StyleLoss(target_feature)
            model.add_module("style_loss_{}".format(i), style_loss)
            style_losses.append(style_loss)

    # убираем лишние слои модели
    for i in range(len(model) - 1, -1, -1):
        if isinstance(model[i], ContentLoss) or isinstance(model[i], StyleLoss):
            break

    model = model[:(i + 1)]

    return model, style_losses, content_losses


def get_input_optimizer(input_img):
    optimizer = optim.LBFGS([input_img.requires_grad_()])
    return optimizer


def run_style_transfer(cnn, name,
                       content_img_size, style_img_size,
                       num_steps=75,
                       style_weight=2e4, content_weight=1):

    style_img = image_loader(name+'_style_photo.pickle', style_img_size)
    content_img = image_loader(name+'_content_photo.pickle', content_img_size)
    input_img = content_img

    model, style_losses, content_losses = get_style_model_and_losses(cnn, style_img, content_img)

    optimizer = get_input_optimizer(input_img)
    run = 0

    print('start')
    while run <= num_steps:

        def closure():
            nonlocal run

            torch.cuda.empty_cache()
            input_img.data.clamp_(0, 1)

            optimizer.zero_grad()
            model(input_img)

            style_score = 0
            content_score = 0

            for sl in style_losses:
                style_score += sl.loss

            for cl in content_losses:
                content_score += cl.loss

            style_score *= style_weight
            content_score *= content_weight

            loss = style_score + content_score
            loss.backward()

            run += 1

            if run % 25 == 0:
                print("run {}:".format(run))
                print('Style Loss: {:4f} Content Loss: {:4f}'.format(
                    style_score.item(), content_score.item()))

                print()

            return style_score + content_score

        optimizer.step(closure)

    input_img.data.clamp_(0, 1)

    return input_img.squeeze(0).cpu()


Ответы (0 шт):