Floral6-DiT Base Template

Это моя первая базовая ИИ-модель для генерации коротких видеороликов с машинками! Она работает на архитектуре MLP (полносвязных слоёв) прямо на CPU и собирает видео со скоростью 3 кадра в секунду из случайного шума.(Кадры для обучения сырые, но хоть есть).

Как это работает:

  • Разрешение: 64x64 (увеличивается до 512x512)
  • Скорость: 3 кадра в секунду (можнно больше)
  • Длина видео: 2 секунды (всего 6 уникальных кадров, можно больше)
  • Особенность: Модель обучается с нуля на 6 картинках датасета, но за счёт случайного шума каждый раз собирает машинцу по-новому.

Код для запуска в Google Colab

Шаг 1: Загрузка датасета (6 картинок)

import os
from google.colab import files

folder_name = "car_dataset"
if not os.path.exists(folder_name): 
    os.makedirs(folder_name)

uploaded = files.upload()

valid_extensions = ('.jpg', '.jpeg', '.png', '.bmp', '.webp')
for filename in uploaded.keys():
    if filename.lower().endswith(valid_extensions): 
        os.rename(filename, os.path.join(folder_name, filename))

total_cars = len([f for f in os.listdir(folder_name) if f.lower().endswith(valid_extensions)])
print(f"Dataset ready. Images: {total_cars}")

Шаг 2: Обучение Моста и генерация видео

import os
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset
from torchvision import transforms
from PIL import Image
import numpy as np
import cv2

folder_name = "car_dataset"
valid_extensions = ('.jpg', '.jpeg', '.png', '.bmp', '.webp')

# Настройки
fps = 3
duration = 2
noise_scale = 1.0
total_frames = fps * duration

class CarDataset(Dataset):
    def __init__(self, folder, transform=None):
        self.folder = folder
        self.transform = transform
        self.images = [f for f in os.listdir(folder) if f.lower().endswith(valid_extensions)][:6]
    def __len__(self): 
        return len(self.images)
    def __getitem__(self, idx):
        img_path = os.path.join(self.folder, self.images[idx])
        return self.transform(Image.open(img_path).convert('RGB'))

transform = transforms.Compose([
    transforms.Resize((64, 64)),
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])

dataset = CarDataset(folder=folder_name, transform=transform)

class CarBridgeTransformer(nn.Module):
    def __init__(self):
        super().__init__()
        self.bridge = nn.Sequential(
            nn.Linear(100, 256),
            nn.LeakyReLU(0.2, inplace=True),
            nn.Linear(256, 512),
            nn.LeakyReLU(0.2, inplace=True),
            nn.Linear(512, 1024),
            nn.LeakyReLU(0.2, inplace=True),
            nn.Linear(1024, 64 * 64 * 3),
            nn.Tanh()
        )
    def forward(self, noise):
        return self.bridge(noise).view(-1, 3, 64, 64)

netG = CarBridgeTransformer() # Работает прямо на CPU!
criterion = nn.MSELoss()
optimizer = optim.Adam(netG.parameters(), lr=0.002)

generated_frames = []

for i in range(total_frames):
    img_idx = i % len(dataset)
    real_car = dataset[img_idx].unsqueeze(0)

    for layer in netG.bridge:
        if hasattr(layer, 'reset_parameters'): 
            layer.reset_parameters()
            
    for epoch in range(30):
        netG.zero_grad()
        noise = torch.randn(1, 100) * noise_scale
        fake_car = netG(noise)
        loss = criterion(fake_car, real_car)
        loss.backward()
        optimizer.step()
        
    with torch.no_grad():
        random_noise = torch.randn(1, 100) * noise_scale
        generated_tensor = netG(random_noise).squeeze(0)
        
        frame = generated_tensor.permute(1, 2, 0).numpy()
        frame = (frame + 1) / 2.0 * 255.0
        frame = frame.astype(np.uint8)
        frame = cv2.cvtColor(frame, cv2.COLOR_RGB2BGR)
        
        frame_resized = cv2.resize(frame, (512, 512), interpolation=cv2.INTER_NEAREST)
        generated_frames.append(frame_resized)

video_name = 'bridge_generated_car.mp4'
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
video = cv2.VideoWriter(video_name, fourcc, fps, (512, 512))

for frame in generated_frames:
    video.write(frame)
video.release()

print(f"Video saved successfully: {video_name}")
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support