ML2-hw5 / app.py
merray0's picture
Upload 3 files
6137a5b verified
Raw
History Blame Contribute Delete
4.78 kB
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
""")