X commited on
Commit
440027f
·
verified ·
1 Parent(s): 0bb2363

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +25 -55
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
- # === Encoder ===
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
- # === LSTM ===
27
  self.lstm = nn.LSTM(
28
- input_size=256 * 8 * 8, # 256 каналов * 8x8 (после 3х пулингов)
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
- # === 1. Кодируем изображение ===
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
- # === 2. Подготовка для LSTM ===
77
- # После пулингов: 256 -> 8x8
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
- # Декодер с skip connections
94
  d4 = self.up4(h)
95
- # Resize skip connection если нужно
96
- if d4.size(-1) != skips[3].size(-1):
97
- skips[3] = F.interpolate(skips[3], size=d4.size(-2:), mode='bilinear')
98
  d4 = torch.cat([d4, skips[3]], dim=1)
99
  d4 = self.dec4(d4)
100
 
101
  d3 = self.up3(d4)
102
- if d3.size(-1) != skips[2].size(-1):
103
- skips[2] = F.interpolate(skips[2], size=d3.size(-2:), mode='bilinear')
104
  d3 = torch.cat([d3, skips[2]], dim=1)
105
  d3 = self.dec3(d3)
106
 
107
  d2 = self.up2(d3)
108
- if d2.size(-1) != skips[1].size(-1):
109
- skips[1] = F.interpolate(skips[1], size=d2.size(-2:), mode='bilinear')
110
  d2 = torch.cat([d2, skips[1]], dim=1)
111
  d2 = self.dec2(d2)
112
 
113
  d1 = self.up1(d2)
114
- if d1.size(-1) != skips[0].size(-1):
115
- skips[0] = F.interpolate(skips[0], size=d1.size(-2:), mode='bilinear')
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(10):
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=15)
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
- style = 'wave'
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}/10 | Loss: {loss_anim.item():.4f} | Style: {loss_style.item():.4f}")
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],