| import streamlit as st |
| import torch |
| import torch.nn as nn |
| from PIL import Image |
| import torchvision.transforms as T |
|
|
|
|
| |
|
|
| class ResidualBlock(nn.Module): |
| def __init__(self, channels): |
| super().__init__() |
| self.block = nn.Sequential( |
| nn.ReflectionPad2d(1), |
| nn.Conv2d(channels, channels, 3), |
| nn.InstanceNorm2d(channels), |
| nn.ReLU(True), |
| nn.ReflectionPad2d(1), |
| nn.Conv2d(channels, channels, 3), |
| nn.InstanceNorm2d(channels), |
| ) |
|
|
| def forward(self, x): |
| return x + self.block(x) |
|
|
|
|
| class ResnetGenerator(nn.Module): |
| def __init__(self, in_channels=3, out_channels=3, base_channels=64, n_residual=6): |
| super().__init__() |
| layers = [ |
| nn.ReflectionPad2d(3), |
| nn.Conv2d(in_channels, base_channels, 7), |
| nn.InstanceNorm2d(base_channels), |
| nn.ReLU(True), |
| ] |
| ch = base_channels |
| for _ in range(2): |
| layers += [ |
| nn.Conv2d(ch, ch * 2, 3, stride=2, padding=1), |
| nn.InstanceNorm2d(ch * 2), |
| nn.ReLU(True), |
| ] |
| ch *= 2 |
| for _ in range(n_residual): |
| layers += [ResidualBlock(ch)] |
| for _ in range(2): |
| layers += [ |
| nn.Upsample(scale_factor=2, mode='nearest'), |
| nn.Conv2d(ch, ch // 2, 3, stride=1, padding=1), |
| nn.InstanceNorm2d(ch // 2), |
| nn.ReLU(True), |
| ] |
| ch //= 2 |
| layers += [ |
| nn.ReflectionPad2d(3), |
| nn.Conv2d(ch, out_channels, 7), |
| nn.Tanh(), |
| ] |
| self.model = nn.Sequential(*layers) |
|
|
| def forward(self, x): |
| return self.model(x) |
|
|
|
|
| |
|
|
| @st.cache_resource |
| def load_models(): |
| ckpt = torch.load('generators.pt', map_location='cpu') |
|
|
| G_A2B = ResnetGenerator(base_channels=64, n_residual=6) |
| G_A2B.load_state_dict(ckpt['G_A2B']) |
| G_A2B.eval() |
|
|
| G_B2A = ResnetGenerator(base_channels=64, n_residual=6) |
| G_B2A.load_state_dict(ckpt['G_B2A']) |
| G_B2A.eval() |
|
|
| return G_A2B, G_B2A |
|
|
|
|
| |
|
|
| IMG_SIZE = 128 |
|
|
| preprocess = T.Compose([ |
| T.Resize((IMG_SIZE, IMG_SIZE)), |
| T.ToTensor(), |
| T.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), |
| ]) |
|
|
|
|
| def tensor_to_pil(t): |
| t = t.clamp(-1, 1) * 0.5 + 0.5 |
| return T.ToPILImage()(t) |
|
|
|
|
| |
|
|
| st.set_page_config( |
| page_title="Messy ↔ Tidy Rooms", |
| page_icon="🧹", |
| layout="wide", |
| ) |
|
|
| st.title("🧹 Messy ↔ Tidy Rooms") |
|
|
| st.markdown(""" |
| CycleGAN для преобразования между неубранными и убранными комнатами. |
| |
| - **A → B**: превращает захламлённую комнату в убранную |
| - **B → A**: добавляет беспорядок в аккуратный интерьер |
| |
| Модель обучена на датасете из ~500 картинок (Kaggle + Bing Images). |
| Загрузите картинку и попробуйте! |
| """) |
|
|
| G_A2B, G_B2A = load_models() |
|
|
| col_left, col_right = st.columns([1, 2]) |
|
|
| with col_left: |
| direction = st.radio( |
| "Направление преобразования:", |
| ["Messy → Tidy (A→B)", "Tidy → Messy (B→A)"], |
| ) |
| uploaded = st.file_uploader( |
| "Загрузите фото комнаты", |
| type=["jpg", "jpeg", "png"], |
| ) |
|
|
| with col_right: |
| if uploaded is not None: |
| img = Image.open(uploaded).convert("RGB") |
|
|
| x = preprocess(img).unsqueeze(0) |
| with torch.no_grad(): |
| if "A→B" in direction: |
| y = G_A2B(x)[0] |
| else: |
| y = G_B2A(x)[0] |
|
|
| result = tensor_to_pil(y) |
|
|
| c1, c2 = st.columns(2) |
| with c1: |
| st.subheader("Оригинал") |
| st.image(img, use_container_width=True) |
| with c2: |
| st.subheader("Результат") |
| st.image(result, use_container_width=True) |
| else: |
| st.info("👆 Загрузите картинку, чтобы увидеть результат") |
|
|
| st.markdown("---") |
| st.markdown("### О модели") |
| st.markdown(""" |
| - **Архитектура**: ResNet-генератор (6 residual блоков) + PatchGAN-дискриминатор |
| - **Лоссы**: LSGAN + Cycle Consistency + Identity Loss |
| - **Датасет**: 500+ картинок интерьеров (Kaggle + Bing Images) |
| - **Обучение**: Adam, lr=2e-4, ~50 эпох на NVIDIA A100 |
| """) |