Работа телеграм-бота с нейросетью с несколькими людьми
Я пишу своего первого бота (на 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()