X commited on
Update app.py
Browse files
app.py
CHANGED
|
@@ -8,31 +8,28 @@ import imageio
|
|
| 8 |
import os
|
| 9 |
import tempfile
|
| 10 |
|
| 11 |
-
# ============ ИСПРАВЛЕННАЯ
|
| 12 |
class FullNeuralAnimator(nn.Module):
|
| 13 |
-
"""
|
| 14 |
-
Одна нейросеть делает ВСЁ
|
| 15 |
-
"""
|
| 16 |
def __init__(self):
|
| 17 |
super().__init__()
|
| 18 |
|
| 19 |
-
#
|
| 20 |
self.enc1 = self._block(3, 32)
|
| 21 |
self.enc2 = self._block(32, 64)
|
| 22 |
self.enc3 = self._block(64, 128)
|
| 23 |
self.enc4 = self._block(128, 256)
|
| 24 |
self.pool = nn.MaxPool2d(2)
|
| 25 |
|
| 26 |
-
#
|
| 27 |
self.lstm = nn.LSTM(
|
| 28 |
-
input_size=256 * 8 * 8,
|
| 29 |
hidden_size=512,
|
| 30 |
num_layers=2,
|
| 31 |
batch_first=True,
|
| 32 |
dropout=0.2
|
| 33 |
)
|
| 34 |
|
| 35 |
-
#
|
| 36 |
self.dec4 = self._block(512, 256)
|
| 37 |
self.dec3 = self._block(256, 128)
|
| 38 |
self.dec2 = self._block(128, 64)
|
|
@@ -43,14 +40,13 @@ class FullNeuralAnimator(nn.Module):
|
|
| 43 |
self.up2 = nn.ConvTranspose2d(128, 64, 2, stride=2)
|
| 44 |
self.up1 = nn.ConvTranspose2d(64, 32, 2, stride=2)
|
| 45 |
|
| 46 |
-
# === Выход для каждого кадра ===
|
| 47 |
self.frame_generator = nn.Sequential(
|
| 48 |
nn.Conv2d(32, 16, 3, padding=1),
|
| 49 |
nn.ReLU(),
|
| 50 |
nn.Conv2d(16, 3, 3, padding=1),
|
| 51 |
nn.Tanh()
|
| 52 |
)
|
| 53 |
-
|
| 54 |
def _block(self, in_ch, out_ch):
|
| 55 |
return nn.Sequential(
|
| 56 |
nn.Conv2d(in_ch, out_ch, 3, padding=1),
|
|
@@ -64,59 +60,52 @@ class FullNeuralAnimator(nn.Module):
|
|
| 64 |
def forward(self, x, num_frames=20):
|
| 65 |
batch_size = x.size(0)
|
| 66 |
|
| 67 |
-
#
|
| 68 |
e1 = self.enc1(x)
|
| 69 |
e2 = self.enc2(self.pool(e1))
|
| 70 |
e3 = self.enc3(self.pool(e2))
|
| 71 |
e4 = self.enc4(self.pool(e3))
|
| 72 |
|
| 73 |
-
# Сохраняем skip connections
|
| 74 |
skips = [e1, e2, e3, e4]
|
| 75 |
|
| 76 |
-
#
|
| 77 |
-
|
| 78 |
-
bottleneck = e4.view(batch_size, -1) # [B, 256*8*8]
|
| 79 |
|
| 80 |
-
# === 3. Генерируем последовательность ===
|
| 81 |
frames = []
|
| 82 |
hidden = None
|
| 83 |
-
|
| 84 |
-
lstm_input = bottleneck.unsqueeze(1) # [B, 1, features]
|
| 85 |
|
| 86 |
for t in range(num_frames):
|
| 87 |
-
# LSTM предсказывает следующее состояние
|
| 88 |
lstm_out, hidden = self.lstm(lstm_input, hidden)
|
| 89 |
|
| 90 |
-
# === 4. Декодируем в кадр ===
|
| 91 |
h = lstm_out.squeeze(1).view(batch_size, 256, 8, 8)
|
| 92 |
|
| 93 |
-
#
|
| 94 |
d4 = self.up4(h)
|
| 95 |
-
#
|
| 96 |
-
if d4.size(
|
| 97 |
-
skips[3] = F.interpolate(skips[3], size=d4.size(
|
| 98 |
d4 = torch.cat([d4, skips[3]], dim=1)
|
| 99 |
d4 = self.dec4(d4)
|
| 100 |
|
| 101 |
d3 = self.up3(d4)
|
| 102 |
-
if d3.size(
|
| 103 |
-
skips[2] = F.interpolate(skips[2], size=d3.size(
|
| 104 |
d3 = torch.cat([d3, skips[2]], dim=1)
|
| 105 |
d3 = self.dec3(d3)
|
| 106 |
|
| 107 |
d2 = self.up2(d3)
|
| 108 |
-
if d2.size(
|
| 109 |
-
skips[1] = F.interpolate(skips[1], size=d2.size(
|
| 110 |
d2 = torch.cat([d2, skips[1]], dim=1)
|
| 111 |
d2 = self.dec2(d2)
|
| 112 |
|
| 113 |
d1 = self.up1(d2)
|
| 114 |
-
if d1.size(
|
| 115 |
-
skips[0] = F.interpolate(skips[0], size=d1.size(
|
| 116 |
d1 = torch.cat([d1, skips[0]], dim=1)
|
| 117 |
d1 = self.dec1(d1)
|
| 118 |
|
| 119 |
-
# Генерируем кадр
|
| 120 |
frame = self.frame_generator(d1)
|
| 121 |
frames.append(frame)
|
| 122 |
|
|
@@ -132,7 +121,6 @@ class StyleTransferAnimator(nn.Module):
|
|
| 132 |
def __init__(self):
|
| 133 |
super().__init__()
|
| 134 |
|
| 135 |
-
# Стили
|
| 136 |
self.style_embeddings = nn.ParameterDict({
|
| 137 |
'wave': nn.Parameter(torch.randn(64)),
|
| 138 |
'pulse': nn.Parameter(torch.randn(64)),
|
|
@@ -141,7 +129,6 @@ class StyleTransferAnimator(nn.Module):
|
|
| 141 |
'twist': nn.Parameter(torch.randn(64)),
|
| 142 |
})
|
| 143 |
|
| 144 |
-
# Encoder
|
| 145 |
self.encoder = nn.Sequential(
|
| 146 |
nn.Conv2d(3, 32, 4, stride=2, padding=1),
|
| 147 |
nn.ReLU(),
|
|
@@ -153,7 +140,6 @@ class StyleTransferAnimator(nn.Module):
|
|
| 153 |
nn.ReLU(),
|
| 154 |
)
|
| 155 |
|
| 156 |
-
# LSTM
|
| 157 |
self.lstm = nn.LSTM(
|
| 158 |
input_size=256 * 8 * 8,
|
| 159 |
hidden_size=512,
|
|
@@ -162,7 +148,6 @@ class StyleTransferAnimator(nn.Module):
|
|
| 162 |
dropout=0.2
|
| 163 |
)
|
| 164 |
|
| 165 |
-
# Decoder
|
| 166 |
self.decoder = nn.Sequential(
|
| 167 |
nn.ConvTranspose2d(512 + 64, 256, 4, stride=2, padding=1),
|
| 168 |
nn.ReLU(),
|
|
@@ -177,11 +162,9 @@ class StyleTransferAnimator(nn.Module):
|
|
| 177 |
def forward(self, x, style='wave', num_frames=20):
|
| 178 |
batch_size = x.size(0)
|
| 179 |
|
| 180 |
-
|
| 181 |
-
features = self.encoder(x) # [B, 256, 8, 8]
|
| 182 |
features_flat = features.view(batch_size, -1)
|
| 183 |
|
| 184 |
-
# Стиль
|
| 185 |
style_vector = self.style_embeddings[style]
|
| 186 |
style_vector = style_vector.unsqueeze(0).repeat(batch_size, 1)
|
| 187 |
|
|
@@ -190,10 +173,8 @@ class StyleTransferAnimator(nn.Module):
|
|
| 190 |
lstm_input = features_flat.unsqueeze(1)
|
| 191 |
|
| 192 |
for t in range(num_frames):
|
| 193 |
-
# LSTM
|
| 194 |
lstm_out, hidden = self.lstm(lstm_input, hidden)
|
| 195 |
|
| 196 |
-
# Декодируем со стилем
|
| 197 |
h = lstm_out.squeeze(1).view(batch_size, 256, 8, 8)
|
| 198 |
style_expanded = style_vector.view(batch_size, 64, 1, 1).repeat(1, 1, 8, 8)
|
| 199 |
decoder_input = torch.cat([h, style_expanded], dim=1)
|
|
@@ -201,7 +182,6 @@ class StyleTransferAnimator(nn.Module):
|
|
| 201 |
frame = self.decoder(decoder_input)
|
| 202 |
frames.append(frame)
|
| 203 |
|
| 204 |
-
# Обновляем вход
|
| 205 |
next_features = self.encoder(frame)
|
| 206 |
next_features = next_features.view(batch_size, -1)
|
| 207 |
lstm_input = next_features.unsqueeze(1)
|
|
@@ -214,11 +194,9 @@ class NeuralAnimator:
|
|
| 214 |
self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
| 215 |
print(f"🔥 Устройство: {self.device}")
|
| 216 |
|
| 217 |
-
# Создаём модели
|
| 218 |
self.animator = FullNeuralAnimator().to(self.device)
|
| 219 |
self.styler = StyleTransferAnimator().to(self.device)
|
| 220 |
|
| 221 |
-
# Пробуем загрузить
|
| 222 |
self.load_models()
|
| 223 |
|
| 224 |
self.animator.eval()
|
|
@@ -239,11 +217,9 @@ class NeuralAnimator:
|
|
| 239 |
print("✅ Стилизатор загружен")
|
| 240 |
|
| 241 |
def generate_animation(self, image, style='wave', num_frames=20, size=128):
|
| 242 |
-
"""Генерирует анимацию"""
|
| 243 |
if image is None:
|
| 244 |
return None
|
| 245 |
|
| 246 |
-
# Подготовка
|
| 247 |
if isinstance(image, np.ndarray):
|
| 248 |
img = Image.fromarray(image)
|
| 249 |
else:
|
|
@@ -253,21 +229,18 @@ class NeuralAnimator:
|
|
| 253 |
img_tensor = torch.from_numpy(np.array(img)).float() / 127.5 - 1
|
| 254 |
img_tensor = img_tensor.permute(2, 0, 1).unsqueeze(0).to(self.device)
|
| 255 |
|
| 256 |
-
# ВСЁ ДЕЛАЕТ НЕЙРОСЕТЬ
|
| 257 |
with torch.no_grad():
|
| 258 |
if style in ['wave', 'pulse', 'glitch', 'melt', 'twist']:
|
| 259 |
frames_tensor = self.styler(img_tensor, style=style, num_frames=num_frames)
|
| 260 |
else:
|
| 261 |
frames_tensor = self.animator(img_tensor, num_frames=num_frames)
|
| 262 |
|
| 263 |
-
# Конвертируем в кадры
|
| 264 |
frames = []
|
| 265 |
for t in range(num_frames):
|
| 266 |
frame = frames_tensor[0, t].cpu().numpy().transpose(1, 2, 0)
|
| 267 |
frame = np.clip((frame + 1) / 2, 0, 1)
|
| 268 |
frames.append((frame * 255).astype(np.uint8))
|
| 269 |
|
| 270 |
-
# Сохраняем
|
| 271 |
temp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.gif')
|
| 272 |
imageio.mimsave(temp_file.name, frames, duration=0.05, loop=0)
|
| 273 |
|
|
@@ -275,7 +248,6 @@ class NeuralAnimator:
|
|
| 275 |
|
| 276 |
# ============ ОБУЧЕНИЕ ============
|
| 277 |
def train_neural_animator():
|
| 278 |
-
"""Обучение нейросети"""
|
| 279 |
print("🧠 Обучаем нейросеть...")
|
| 280 |
|
| 281 |
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
|
@@ -290,12 +262,12 @@ def train_neural_animator():
|
|
| 290 |
|
| 291 |
print("🚀 Начинаем обучение...")
|
| 292 |
|
| 293 |
-
for epoch in range(
|
| 294 |
batch_size = 2
|
| 295 |
fake_images = torch.randn(batch_size, 3, 128, 128, device=device)
|
| 296 |
|
| 297 |
# Обучаем аниматор
|
| 298 |
-
frames = animator(fake_images, num_frames=
|
| 299 |
loss_smooth = mse(frames[:, 1:], frames[:, :-1])
|
| 300 |
loss_consistency = mse(frames.mean(dim=1), fake_images)
|
| 301 |
loss_anim = loss_smooth + 0.5 * loss_consistency
|
|
@@ -305,8 +277,7 @@ def train_neural_animator():
|
|
| 305 |
opt_anim.step()
|
| 306 |
|
| 307 |
# Обучаем стилизатор
|
| 308 |
-
|
| 309 |
-
style_frames = styler(fake_images, style=style, num_frames=15)
|
| 310 |
loss_style_smooth = mse(style_frames[:, 1:], style_frames[:, :-1])
|
| 311 |
loss_style_consistency = mse(style_frames.mean(dim=1), fake_images)
|
| 312 |
loss_style = loss_style_smooth + 0.5 * loss_style_consistency
|
|
@@ -315,7 +286,7 @@ def train_neural_animator():
|
|
| 315 |
loss_style.backward()
|
| 316 |
opt_style.step()
|
| 317 |
|
| 318 |
-
print(f"Epoch {epoch+1}/
|
| 319 |
|
| 320 |
os.makedirs('neural_models', exist_ok=True)
|
| 321 |
torch.save(animator.state_dict(), 'neural_models/animator.pth')
|
|
@@ -399,7 +370,6 @@ with gr.Blocks(title="🧠 Нейросетевая анимация") as demo:
|
|
| 399 |
variant="primary"
|
| 400 |
)
|
| 401 |
|
| 402 |
-
# Логика
|
| 403 |
generate_btn.click(
|
| 404 |
fn=generate_wrapper,
|
| 405 |
inputs=[input_image, style, frames, size],
|
|
|
|
| 8 |
import os
|
| 9 |
import tempfile
|
| 10 |
|
| 11 |
+
# ============ ИСПРАВЛЕННАЯ ВЕРСИЯ ============
|
| 12 |
class FullNeuralAnimator(nn.Module):
|
|
|
|
|
|
|
|
|
|
| 13 |
def __init__(self):
|
| 14 |
super().__init__()
|
| 15 |
|
| 16 |
+
# Encoder
|
| 17 |
self.enc1 = self._block(3, 32)
|
| 18 |
self.enc2 = self._block(32, 64)
|
| 19 |
self.enc3 = self._block(64, 128)
|
| 20 |
self.enc4 = self._block(128, 256)
|
| 21 |
self.pool = nn.MaxPool2d(2)
|
| 22 |
|
| 23 |
+
# LSTM
|
| 24 |
self.lstm = nn.LSTM(
|
| 25 |
+
input_size=256 * 8 * 8,
|
| 26 |
hidden_size=512,
|
| 27 |
num_layers=2,
|
| 28 |
batch_first=True,
|
| 29 |
dropout=0.2
|
| 30 |
)
|
| 31 |
|
| 32 |
+
# Decoder
|
| 33 |
self.dec4 = self._block(512, 256)
|
| 34 |
self.dec3 = self._block(256, 128)
|
| 35 |
self.dec2 = self._block(128, 64)
|
|
|
|
| 40 |
self.up2 = nn.ConvTranspose2d(128, 64, 2, stride=2)
|
| 41 |
self.up1 = nn.ConvTranspose2d(64, 32, 2, stride=2)
|
| 42 |
|
|
|
|
| 43 |
self.frame_generator = nn.Sequential(
|
| 44 |
nn.Conv2d(32, 16, 3, padding=1),
|
| 45 |
nn.ReLU(),
|
| 46 |
nn.Conv2d(16, 3, 3, padding=1),
|
| 47 |
nn.Tanh()
|
| 48 |
)
|
| 49 |
+
|
| 50 |
def _block(self, in_ch, out_ch):
|
| 51 |
return nn.Sequential(
|
| 52 |
nn.Conv2d(in_ch, out_ch, 3, padding=1),
|
|
|
|
| 60 |
def forward(self, x, num_frames=20):
|
| 61 |
batch_size = x.size(0)
|
| 62 |
|
| 63 |
+
# Encoder
|
| 64 |
e1 = self.enc1(x)
|
| 65 |
e2 = self.enc2(self.pool(e1))
|
| 66 |
e3 = self.enc3(self.pool(e2))
|
| 67 |
e4 = self.enc4(self.pool(e3))
|
| 68 |
|
|
|
|
| 69 |
skips = [e1, e2, e3, e4]
|
| 70 |
|
| 71 |
+
# Bottleneck
|
| 72 |
+
bottleneck = e4.view(batch_size, -1)
|
|
|
|
| 73 |
|
|
|
|
| 74 |
frames = []
|
| 75 |
hidden = None
|
| 76 |
+
lstm_input = bottleneck.unsqueeze(1)
|
|
|
|
| 77 |
|
| 78 |
for t in range(num_frames):
|
|
|
|
| 79 |
lstm_out, hidden = self.lstm(lstm_input, hidden)
|
| 80 |
|
|
|
|
| 81 |
h = lstm_out.squeeze(1).view(batch_size, 256, 8, 8)
|
| 82 |
|
| 83 |
+
# Decoder с проверкой размеров
|
| 84 |
d4 = self.up4(h)
|
| 85 |
+
# Проверяем размеры и при необходимости ресайзим
|
| 86 |
+
if d4.size(2) != skips[3].size(2) or d4.size(3) != skips[3].size(3):
|
| 87 |
+
skips[3] = F.interpolate(skips[3], size=(d4.size(2), d4.size(3)), mode='bilinear')
|
| 88 |
d4 = torch.cat([d4, skips[3]], dim=1)
|
| 89 |
d4 = self.dec4(d4)
|
| 90 |
|
| 91 |
d3 = self.up3(d4)
|
| 92 |
+
if d3.size(2) != skips[2].size(2) or d3.size(3) != skips[2].size(3):
|
| 93 |
+
skips[2] = F.interpolate(skips[2], size=(d3.size(2), d3.size(3)), mode='bilinear')
|
| 94 |
d3 = torch.cat([d3, skips[2]], dim=1)
|
| 95 |
d3 = self.dec3(d3)
|
| 96 |
|
| 97 |
d2 = self.up2(d3)
|
| 98 |
+
if d2.size(2) != skips[1].size(2) or d2.size(3) != skips[1].size(3):
|
| 99 |
+
skips[1] = F.interpolate(skips[1], size=(d2.size(2), d2.size(3)), mode='bilinear')
|
| 100 |
d2 = torch.cat([d2, skips[1]], dim=1)
|
| 101 |
d2 = self.dec2(d2)
|
| 102 |
|
| 103 |
d1 = self.up1(d2)
|
| 104 |
+
if d1.size(2) != skips[0].size(2) or d1.size(3) != skips[0].size(3):
|
| 105 |
+
skips[0] = F.interpolate(skips[0], size=(d1.size(2), d1.size(3)), mode='bilinear')
|
| 106 |
d1 = torch.cat([d1, skips[0]], dim=1)
|
| 107 |
d1 = self.dec1(d1)
|
| 108 |
|
|
|
|
| 109 |
frame = self.frame_generator(d1)
|
| 110 |
frames.append(frame)
|
| 111 |
|
|
|
|
| 121 |
def __init__(self):
|
| 122 |
super().__init__()
|
| 123 |
|
|
|
|
| 124 |
self.style_embeddings = nn.ParameterDict({
|
| 125 |
'wave': nn.Parameter(torch.randn(64)),
|
| 126 |
'pulse': nn.Parameter(torch.randn(64)),
|
|
|
|
| 129 |
'twist': nn.Parameter(torch.randn(64)),
|
| 130 |
})
|
| 131 |
|
|
|
|
| 132 |
self.encoder = nn.Sequential(
|
| 133 |
nn.Conv2d(3, 32, 4, stride=2, padding=1),
|
| 134 |
nn.ReLU(),
|
|
|
|
| 140 |
nn.ReLU(),
|
| 141 |
)
|
| 142 |
|
|
|
|
| 143 |
self.lstm = nn.LSTM(
|
| 144 |
input_size=256 * 8 * 8,
|
| 145 |
hidden_size=512,
|
|
|
|
| 148 |
dropout=0.2
|
| 149 |
)
|
| 150 |
|
|
|
|
| 151 |
self.decoder = nn.Sequential(
|
| 152 |
nn.ConvTranspose2d(512 + 64, 256, 4, stride=2, padding=1),
|
| 153 |
nn.ReLU(),
|
|
|
|
| 162 |
def forward(self, x, style='wave', num_frames=20):
|
| 163 |
batch_size = x.size(0)
|
| 164 |
|
| 165 |
+
features = self.encoder(x)
|
|
|
|
| 166 |
features_flat = features.view(batch_size, -1)
|
| 167 |
|
|
|
|
| 168 |
style_vector = self.style_embeddings[style]
|
| 169 |
style_vector = style_vector.unsqueeze(0).repeat(batch_size, 1)
|
| 170 |
|
|
|
|
| 173 |
lstm_input = features_flat.unsqueeze(1)
|
| 174 |
|
| 175 |
for t in range(num_frames):
|
|
|
|
| 176 |
lstm_out, hidden = self.lstm(lstm_input, hidden)
|
| 177 |
|
|
|
|
| 178 |
h = lstm_out.squeeze(1).view(batch_size, 256, 8, 8)
|
| 179 |
style_expanded = style_vector.view(batch_size, 64, 1, 1).repeat(1, 1, 8, 8)
|
| 180 |
decoder_input = torch.cat([h, style_expanded], dim=1)
|
|
|
|
| 182 |
frame = self.decoder(decoder_input)
|
| 183 |
frames.append(frame)
|
| 184 |
|
|
|
|
| 185 |
next_features = self.encoder(frame)
|
| 186 |
next_features = next_features.view(batch_size, -1)
|
| 187 |
lstm_input = next_features.unsqueeze(1)
|
|
|
|
| 194 |
self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
| 195 |
print(f"🔥 Устройство: {self.device}")
|
| 196 |
|
|
|
|
| 197 |
self.animator = FullNeuralAnimator().to(self.device)
|
| 198 |
self.styler = StyleTransferAnimator().to(self.device)
|
| 199 |
|
|
|
|
| 200 |
self.load_models()
|
| 201 |
|
| 202 |
self.animator.eval()
|
|
|
|
| 217 |
print("✅ Стилизатор загружен")
|
| 218 |
|
| 219 |
def generate_animation(self, image, style='wave', num_frames=20, size=128):
|
|
|
|
| 220 |
if image is None:
|
| 221 |
return None
|
| 222 |
|
|
|
|
| 223 |
if isinstance(image, np.ndarray):
|
| 224 |
img = Image.fromarray(image)
|
| 225 |
else:
|
|
|
|
| 229 |
img_tensor = torch.from_numpy(np.array(img)).float() / 127.5 - 1
|
| 230 |
img_tensor = img_tensor.permute(2, 0, 1).unsqueeze(0).to(self.device)
|
| 231 |
|
|
|
|
| 232 |
with torch.no_grad():
|
| 233 |
if style in ['wave', 'pulse', 'glitch', 'melt', 'twist']:
|
| 234 |
frames_tensor = self.styler(img_tensor, style=style, num_frames=num_frames)
|
| 235 |
else:
|
| 236 |
frames_tensor = self.animator(img_tensor, num_frames=num_frames)
|
| 237 |
|
|
|
|
| 238 |
frames = []
|
| 239 |
for t in range(num_frames):
|
| 240 |
frame = frames_tensor[0, t].cpu().numpy().transpose(1, 2, 0)
|
| 241 |
frame = np.clip((frame + 1) / 2, 0, 1)
|
| 242 |
frames.append((frame * 255).astype(np.uint8))
|
| 243 |
|
|
|
|
| 244 |
temp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.gif')
|
| 245 |
imageio.mimsave(temp_file.name, frames, duration=0.05, loop=0)
|
| 246 |
|
|
|
|
| 248 |
|
| 249 |
# ============ ОБУЧЕНИЕ ============
|
| 250 |
def train_neural_animator():
|
|
|
|
| 251 |
print("🧠 Обучаем нейросеть...")
|
| 252 |
|
| 253 |
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
|
|
|
| 262 |
|
| 263 |
print("🚀 Начинаем обучение...")
|
| 264 |
|
| 265 |
+
for epoch in range(5): # 5 эпох для быстро��о теста
|
| 266 |
batch_size = 2
|
| 267 |
fake_images = torch.randn(batch_size, 3, 128, 128, device=device)
|
| 268 |
|
| 269 |
# Обучаем аниматор
|
| 270 |
+
frames = animator(fake_images, num_frames=10)
|
| 271 |
loss_smooth = mse(frames[:, 1:], frames[:, :-1])
|
| 272 |
loss_consistency = mse(frames.mean(dim=1), fake_images)
|
| 273 |
loss_anim = loss_smooth + 0.5 * loss_consistency
|
|
|
|
| 277 |
opt_anim.step()
|
| 278 |
|
| 279 |
# Обучаем стилизатор
|
| 280 |
+
style_frames = styler(fake_images, style='wave', num_frames=10)
|
|
|
|
| 281 |
loss_style_smooth = mse(style_frames[:, 1:], style_frames[:, :-1])
|
| 282 |
loss_style_consistency = mse(style_frames.mean(dim=1), fake_images)
|
| 283 |
loss_style = loss_style_smooth + 0.5 * loss_style_consistency
|
|
|
|
| 286 |
loss_style.backward()
|
| 287 |
opt_style.step()
|
| 288 |
|
| 289 |
+
print(f"Epoch {epoch+1}/5 | Loss: {loss_anim.item():.4f} | Style: {loss_style.item():.4f}")
|
| 290 |
|
| 291 |
os.makedirs('neural_models', exist_ok=True)
|
| 292 |
torch.save(animator.state_dict(), 'neural_models/animator.pth')
|
|
|
|
| 370 |
variant="primary"
|
| 371 |
)
|
| 372 |
|
|
|
|
| 373 |
generate_btn.click(
|
| 374 |
fn=generate_wrapper,
|
| 375 |
inputs=[input_image, style, frames, size],
|