معاينة مختبر آمنة
Auto Encoders Py Torch
هذي معاينة منقّحة للقراءة فقط؛ ما فيه أي شيء يشتغل داخل الصفحة.
قراءة فقط
معاينة الدفتر
Auto Encoders Py Torch
> **ملاحظة بيئة التشغيل المدمجة:** هالمعاينة تستخدم عيّنة صغيرة وثابتة وآمنة من ناحية الحقوق عشان تكون النتايج قابلة للتكرار. النتايج بالحجم الكامل تحتاج مجموعة البيانات أو النموذج الموثّق بالدرس داخل بيئة خارجية معتمدة.
# [المشفّرات التلقائية (Autoencoders)](https://arxiv.org/abs/2201.03898)
إذا جينا ندرّب الشبكات العصبية الالتفافية (CNNs)، بنواجه مشكلة: نحتاج كمية كبيرة من البيانات المعلَّمة. وفي مهمة **تصنيف الصور (Image classification)**، وهي من مسائل **التصنيف (Classification)**، لازم نفرز الصور يدويًا على فئات مختلفة.
لكن نقدر نستخدم بيانات خام ما عليها علامات عشان ندرّب مستخرجات السمات في CNN؛ وهالأسلوب يسمّى **التعلّم ذاتي الإشراف (Self-supervised learning)**. بدال العلامات، نستخدم صور التدريب نفسها مدخلات ومخرجات مطلوبة من الشبكة. فكرة **المشفّر التلقائي (Autoencoder)** إن عندنا **شبكة ترميز (Encoder)** تحوّل الصورة إلى **فضاء كامن (Latent space)**، وغالبًا يكون متجهًا أصغر، وبعدها **شبكة فك الترميز (Decoder)** تحاول تعيد بناء الصورة الأصلية.
ولأننا ندرّب المشفّر التلقائي على الاحتفاظ بأكبر قدر ممكن من معلومات الصورة عشان يعيد بناءها بدقة، فالشبكة تحاول تلقى أفضل **تضمين (Embedding)** يمثّل معنى الصورة.
> **وصف الشكل:** مخطط يوضّح المشفّر التلقائي
> الصورة مأخوذة من [مدونة Keras](https://blog.keras.io/building-autoencoders-in-keras.html)
خلّونا نبني أبسط نموذج من **المشفّر التلقائي (Autoencoder)** لـ MNIST!
import torch
import torchvision
import matplotlib.pyplot as plt
from torchvision import transforms
from torch import nn
from torch import optim
from tqdm import tqdm
import numpy as np
import torch.nn.functional as F
torch.manual_seed(42)
np.random.seed(42)حدّدوا معلمات التدريب، وتأكدوا هل وحدة معالجة الرسومات متوفرة:
device = 'cuda:0' if torch.cuda.is_available() else 'cpu'
train_size = 0.9
lr = 1e-3
eps = 1e-8
batch_size = 256
epochs = 2الدالة الجاية تحمّل **مجموعة البيانات (Dataset)** MNIST وتطبّق عليها التحويلات المحددة، وبعد تقسّمها إلى مجموعتي التدريب والاختبار.
# course-edition bundled digits fixture v1
from sklearn.datasets import load_digits as _course_load_digits
def mnist(train_part, transform=None):
_course_digits = _course_load_digits()
_course_images = torch.as_tensor(_course_digits.images, dtype=torch.float32).unsqueeze(1) / 16.0
_course_images = torch.nn.functional.interpolate(
_course_images,
size=(28, 28),
mode="bilinear",
align_corners=False,
)
_course_labels = torch.as_tensor(_course_digits.target, dtype=torch.long)
dataset = torch.utils.data.TensorDataset(_course_images, _course_labels)
train_count = max(1, min(len(dataset) - 1, int(train_part * len(dataset))))
return torch.utils.data.random_split(
dataset,
[train_count, len(dataset) - train_count],
generator=torch.Generator().manual_seed(2026),
)الحين نحمّل مجموعة البيانات ونعرّف محمّلات بيانات التدريب والاختبار:
transform = transforms.Compose([transforms.ToTensor()])
train_dataset, test_dataset = mnist(train_size, transform)
train_dataloader = torch.utils.data.DataLoader(train_dataset, drop_last=True, batch_size=batch_size, shuffle=True)
test_dataloader = torch.utils.data.DataLoader(test_dataset, batch_size=1, shuffle=False)
dataloaders = (train_dataloader, test_dataloader)def plotn(n, data, noisy=False, super_res=None):
fig, ax = plt.subplots(1, n)
for i, z in enumerate(data):
if i == n:
break
preprocess = z[0].reshape(1, 28, 28) if z[0].shape[1] == 28 else z[0].reshape(1, 14, 14) if z[0].shape[1] == 14 else z[0]
if super_res is not None:
_transform = transforms.Resize((int(preprocess.shape[1] / super_res), int(preprocess.shape[2] / super_res)))
preprocess = _transform(preprocess)
if noisy:
shapes = list(preprocess.shape)
preprocess += noisify(shapes)
ax[i].imshow(preprocess[0])
plt.show()def noisify(shapes):
return np.random.normal(loc=0.5, scale=0.3, size=shapes)plotn(5, train_dataset)class Encoder(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(1, 16, kernel_size=(3, 3), padding='same')
self.maxpool1 = nn.MaxPool2d(kernel_size=(2, 2))
self.conv2 = nn.Conv2d(16, 8, kernel_size=(3, 3), padding='same')
self.maxpool2 = nn.MaxPool2d(kernel_size=(2, 2))
self.conv3 = nn.Conv2d(8, 8, kernel_size=(3, 3), padding='same')
self.maxpool3 = nn.MaxPool2d(kernel_size=(2, 2), padding=(1, 1))
self.relu = nn.ReLU()
def forward(self, input):
hidden1 = self.maxpool1(self.relu(self.conv1(input)))
hidden2 = self.maxpool2(self.relu(self.conv2(hidden1)))
encoded = self.maxpool3(self.relu(self.conv3(hidden2)))
return encodedclass Decoder(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(8, 8, kernel_size=(3, 3), padding='same')
self.upsample1 = nn.Upsample(scale_factor=(2, 2))
self.conv2 = nn.Conv2d(8, 8, kernel_size=(3, 3), padding='same')
self.upsample2 = nn.Upsample(scale_factor=(2, 2))
self.conv3 = nn.Conv2d(8, 16, kernel_size=(3, 3))
self.upsample3 = nn.Upsample(scale_factor=(2, 2))
self.conv4 = nn.Conv2d(16, 1, kernel_size=(3, 3), padding='same')
self.relu = nn.ReLU()
self.sigmoid = nn.Sigmoid()
def forward(self, input):
hidden1 = self.upsample1(self.relu(self.conv1(input)))
hidden2 = self.upsample2(self.relu(self.conv2(hidden1)))
hidden3 = self.upsample3(self.relu(self.conv3(hidden2)))
decoded = self.sigmoid(self.conv4(hidden3))
return decodedclass AutoEncoder(nn.Module):
def __init__(self, super_resolution=False):
super().__init__()
if not super_resolution:
self.encoder = Encoder()
else:
self.encoder = SuperResolutionEncoder()
self.decoder = Decoder()
def forward(self, input):
encoded = self.encoder(input)
decoded = self.decoder(encoded)
return decodedmodel = AutoEncoder().to(device)
optimizer = optim.Adam(model.parameters(), lr=lr, eps=eps)
loss_fn = nn.BCELoss()def train(dataloaders, model, loss_fn, optimizer, epochs, device, noisy=None, super_res=None):
tqdm_iter = tqdm(range(epochs))
train_dataloader, test_dataloader = dataloaders[0], dataloaders[1]
for epoch in tqdm_iter:
model.train()
train_loss = 0.0
test_loss = 0.0
for batch in train_dataloader:
imgs, labels = batch
shapes = list(imgs.shape)
if super_res is not None:
shapes[2], shapes[3] = int(shapes[2] / super_res), int(shapes[3] / super_res)
_transform = transforms.Resize((shapes[2], shapes[3]))
imgs_transformed = _transform(imgs)
imgs_transformed = imgs_transformed.to(device)
imgs = imgs.to(device)
labels = labels.to(device)
if noisy is not None:
noisy_tensor = noisy[0]
else:
noisy_tensor = torch.zeros(tuple(shapes)).to(device)
if super_res is None:
imgs_noisy = imgs + noisy_tensor
else:
imgs_noisy = imgs_transformed + noisy_tensor
imgs_noisy = torch.clamp(imgs_noisy, 0., 1.)
preds = model(imgs_noisy)
loss = loss_fn(preds, imgs)
optimizer.zero_grad()
loss.backward()
optimizer.step()
train_loss += loss.item()
model.eval()
with torch.no_grad():
for batch in test_dataloader:
imgs, labels = batch
shapes = list(imgs.shape)
if super_res is not None:
shapes[2], shapes[3] = int(shapes[2] / super_res), int(shapes[3] / super_res)
_transform = transforms.Resize((shapes[2], shapes[3]))
imgs_transformed = _transform(imgs)
imgs_transformed = imgs_transformed.to(device)
imgs = imgs.to(device)
labels = labels.to(device)
if noisy is not None:
test_noisy_tensor = noisy[1]
else:
test_noisy_tensor = torch.zeros(tuple(shapes)).to(device)
if super_res is None:
imgs_noisy = imgs + test_noisy_tensor
else:
imgs_noisy = imgs_transformed + test_noisy_tensor
imgs_noisy = torch.clamp(imgs_noisy, 0., 1.)
preds = model(imgs_noisy)
loss = loss_fn(preds, imgs)
test_loss += loss.item()
train_loss /= len(train_dataloader)
test_loss /= len(test_dataloader)
tqdm_dct = {'train loss:': train_loss, 'test loss:': test_loss}
tqdm_iter.set_postfix(tqdm_dct, refresh=True)
tqdm_iter.refresh()train(dataloaders, model, loss_fn, optimizer, epochs, device)model.eval()
predictions = []
plots = 5
for i, data in enumerate(test_dataset):
if i == plots:
break
predictions.append(model(data[0].to(device).unsqueeze(0)).detach().cpu())
plotn(plots, test_dataset)
plotn(plots, predictions)> **المهمة 1**: جرّبوا تدرّبون المشفّر التلقائي بمتجه كامن صغير جدًا، مثل 2، وارسموا النقاط اللي تمثّل الأرقام المختلفة. *تلميح: استخدموا طبقة كثيفة متصلة بالكامل بعد الجزء الالتفافي عشان تقلّلون حجم المتجه للقيمة المطلوبة.*
> **المهمة 2**: ابدأوا بأرقام مختلفة، واستخرجوا تمثيلاتها في الفضاء الكامن، ثم شوفوا وش يصير للأرقام الناتجة إذا أضفنا تشويشًا للفضاء الكامن.
## إزالة التشويش
نقدر نستخدم المشفّرات التلقائية بفعالية لإزالة التشويش من الصور. نبدأ بصور صافية ونضيف عليها تشويشًا صناعيًا، ثم ندخل النسخ المشوّشة للشبكة ونخلي الصور الصافية هي المخرجات المطلوبة.
خلّونا نشوف كيف تشتغل الفكرة مع MNIST:
plotn(5, train_dataset, noisy=True)model = AutoEncoder().to(device)
optimizer = optim.Adam(model.parameters(), lr=lr, eps=eps)
loss_fn = nn.BCELoss()noisy_tensor = torch.FloatTensor(noisify([256, 1, 28, 28])).to(device)
test_noisy_tensor = torch.FloatTensor(noisify([1, 1, 28, 28])).to(device)
noisy_tensors = (noisy_tensor, test_noisy_tensor)train(dataloaders, model, loss_fn, optimizer, 100, device, noisy=noisy_tensors)model.eval()
predictions = []
noise = []
plots = 5
for i, data in enumerate(test_dataset):
if i == plots:
break
shapes = data[0].shape
noisy_data = data[0] + test_noisy_tensor[0].detach().cpu()
noise.append(noisy_data)
predictions.append(model(noisy_data.to(device).unsqueeze(0)).detach().cpu())
plotn(plots, noise)
plotn(plots, predictions)> **تمرين:** جرّبوا مزيل التشويش المدرَّب على أرقام MNIST مع صور مختلفة. تقدرون تستخدمون مجموعة [Fashion MNIST](https://pytorch.org/vision/stable/generated/torchvision.datasets.FashionMNIST.html#torchvision.datasets.FashionMNIST) لأن صورها بالحجم نفسه. لاحظوا إن النموذج يشتغل زين بس على نوع الصور اللي تدرّب عليه؛ يعني على توزيع احتمالي مشابه لبيانات الإدخال.
## رفع الدقة
مثل إزالة التشويش، نقدر ندرّب المشفّرات التلقائية على رفع دقة الصورة. نبدأ بصور عالية الدقة ونخفّض دقتها آليًا عشان نصنع المدخلات، ثم نعطي الشبكة الصور الصغيرة مدخلات والصور عالية الدقة مخرجات مطلوبة.
خلّونا نخفّض دقة الصورة إلى 14x14 وقت التدريب.
super_res_koeff = 2.0
plotn(5, train_dataset, super_res=super_res_koeff)class SuperResolutionEncoder(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(1, 16, kernel_size=(3, 3), padding='same')
self.maxpool1 = nn.MaxPool2d(kernel_size=(2, 2))
self.conv2 = nn.Conv2d(16, 8, kernel_size=(3, 3), padding='same')
self.maxpool2 = nn.MaxPool2d(kernel_size=(2, 2), padding=(1, 1))
self.relu = nn.ReLU()
def forward(self, input):
hidden1 = self.maxpool1(self.relu(self.conv1(input)))
encoded = self.maxpool2(self.relu(self.conv2(hidden1)))
return encodedmodel = AutoEncoder(super_resolution=True).to(device)
optimizer = optim.Adam(model.parameters(), lr=lr, eps=eps)
loss_fn = nn.BCELoss()train(dataloaders, model, loss_fn, optimizer, epochs, device, super_res=2.0)model.eval()
predictions = []
plots = 5
shapes = test_dataset[0][0].shape
for i, data in enumerate(test_dataset):
if i == plots:
break
_transform = transforms.Resize((int(shapes[1] / super_res_koeff), int(shapes[2] / super_res_koeff)))
predictions.append(model(_transform(data[0]).to(device).unsqueeze(0)).detach().cpu())
plotn(plots, test_dataset, super_res=super_res_koeff)
plotn(plots, predictions)> **تمرين**: جرّبوا تدرّبون شبكة رفع الدقة على [CIFAR-10](https://pytorch.org/vision/stable/generated/torchvision.datasets.CIFAR10.html) للتكبير بمقدار 2x و4x. استخدموا التشويش مدخلًا لنموذج 4x وراقبوا النتيجة.
# [المشفّرات التلقائية التغايرية (Variational Autoencoders — VAE)](https://arxiv.org/abs/1906.02691)
المشفّرات التلقائية التقليدية تقلّل أبعاد بيانات الإدخال وتكتشف أهم السمات في الصور. لكن المتجهات الكامنة الناتجة غالبًا تكون مب واضحة المعنى. خذوا **مجموعة البيانات (Dataset)** MNIST مثالًا: مو سهل نعرف أي رقم يقابله كل متجه كامن، والمتجهات المتقاربة مو لازم تمثّل الرقم نفسه.
أما إذا كنا بندرّب نماذج *توليدية*، فنحتاج نفهم الفضاء الكامن وترتيبه بشكل أفضل. ومن هنا تجي فكرة **المشفّر التلقائي التغايري (VAE)**.
يتعلّم VAE يتنبأ بـ *توزيع إحصائي* للمعلمات الكامنة، وهذا اللي نسمّيه **التوزيع الكامن**. مثلًا، نفترض أن المتجهات الكامنة موزعة كـ $N(\mathrm{z\_mean},e^{\mathrm{z\_log}})$، حيث $\mathrm{z\_mean}, \mathrm{z\_log} \in\mathbb{R}^d$. تتنبأ شبكة الترميز في VAE بهالمعلمات، ثم يأخذ فاكّ الترميز متجهًا عشوائيًا من التوزيع ويحاول يعيد بناء العنصر الأصلي.
خلّونا نلخّصها:
* من متجه الإدخال نتنبأ بـ `z_mean` و`z_log`؛ يعني نتنبأ بلوغاريتم الانحراف المعياري بدال الانحراف نفسه.
* نسحب المتجه `sample(z_val in code)` من التوزيع $N(\mathrm{z\_mean},e^{\mathrm{z\_log\_sigma}})$.
* يحاول فاكّ الترميز يعيد بناء الصورة الأصلية باستخدام `sample` متجهَ إدخال.
> **وصف الشكل:** حُذف الأصل لأن حقوق إعادة استخدامه غير موثّقة.
> الصورة مأخوذة من [هالمقالة](https://ijdykeman.github.io/ml/2016/12/21/cvae.html) لـ Isaak Dykeman
class VAEEncoder(nn.Module):
def __init__(self, device):
super().__init__()
self.intermediate_dim = 512
self.latent_dim = 2
self.linear = nn.Linear(784, self.intermediate_dim)
self.z_mean = nn.Linear(self.intermediate_dim, self.latent_dim)
self.z_log = nn.Linear(self.intermediate_dim, self.latent_dim)
self.relu = nn.ReLU()
self.device = device
def forward(self, input):
bs = input.shape[0]
hidden = self.relu(self.linear(input))
z_mean = self.z_mean(hidden)
z_log = self.z_log(hidden)
eps = torch.FloatTensor(np.random.normal(size=(bs, self.latent_dim))).to(device)
z_val = z_mean + torch.exp(z_log) * eps
return z_mean, z_log, z_valclass VAEDecoder(nn.Module):
def __init__(self):
super().__init__()
self.intermediate_dim = 512
self.latent_dim = 2
self.linear = nn.Linear(self.latent_dim, self.intermediate_dim)
self.output = nn.Linear(self.intermediate_dim, 784)
self.relu = nn.ReLU()
self.sigmoid = nn.Sigmoid()
def forward(self, input):
hidden = self.relu(self.linear(input))
decoded = self.sigmoid(self.output(hidden))
return decodedclass VAEAutoEncoder(nn.Module):
def __init__(self, device):
super().__init__()
self.encoder = VAEEncoder(device)
self.decoder = VAEDecoder()
self.z_vals = None
def forward(self, input):
bs, c, h, w = input.shape[0], input.shape[1], input.shape[2], input.shape[3]
input = input.view(bs, -1)
encoded = self.encoder(input)
self.z_vals = encoded
decoded = self.decoder(encoded[2])
return decoded
def get_zvals(self):
return self.z_valsتستخدم المشفّرات التلقائية التغايرية دالة خسارة مركّبة من جزأين:
* **خسارة إعادة البناء (Reconstruction loss)** تقيس قرب الصورة المعاد بناؤها من الهدف، وممكن تكون MSE. وهي نفس دالة الخسارة في المشفّر التلقائي العادي.
* **خسارة KL** تخلي توزيعات المتغيرات الكامنة قريبة من التوزيع الطبيعي. وتعتمد على [تباعد كولباك–ليبلر](https://www.countbayesie.com/blog/2017/5/9/kullback-leibler-divergence-explained)، وهو مقياس يقدّر مدى التشابه بين توزيعين إحصائيين.
def vae_loss(preds, targets, z_vals):
mse = nn.MSELoss()
reconstruction_loss = mse(preds, targets.view(targets.shape[0], -1)) * 784.0
temp = 1.0 + z_vals[1] - torch.square(z_vals[0]) - torch.exp(z_vals[1])
kl_loss = -0.5 * torch.sum(temp, axis=-1)
return torch.mean(reconstruction_loss + kl_loss)model = VAEAutoEncoder(device).to(device)
optimizer = optim.RMSprop(model.parameters(), lr=lr, eps=eps)def train_vae(dataloaders, model, optimizer, epochs, device):
tqdm_iter = tqdm(range(epochs))
train_dataloader, test_dataloader = dataloaders[0], dataloaders[1]
for epoch in tqdm_iter:
model.train()
train_loss = 0.0
test_loss = 0.0
for batch in train_dataloader:
imgs, labels = batch
imgs = imgs.to(device)
labels = labels.to(device)
preds = model(imgs)
z_vals = model.get_zvals()
loss = vae_loss(preds, imgs, z_vals)
optimizer.zero_grad()
loss.backward()
optimizer.step()
train_loss += loss.item()
model.eval()
with torch.no_grad():
for batch in test_dataloader:
imgs, labels = batch
imgs = imgs.to(device)
labels = labels.to(device)
preds = model(imgs)
z_vals = model.get_zvals()
loss = vae_loss(preds, imgs, z_vals)
test_loss += loss.item()
train_loss /= len(train_dataloader)
test_loss /= len(test_dataloader)
tqdm_dct = {'train loss:': train_loss, 'test loss:': test_loss}
tqdm_iter.set_postfix(tqdm_dct, refresh=True)
tqdm_iter.refresh()train_vae(dataloaders, model, optimizer, epochs, device)model.eval()
predictions = []
plots = 5
for i, data in enumerate(test_dataset):
if i == plots:
break
predictions.append(model(data[0].to(device).unsqueeze(0)).view(1, 28, 28).detach().cpu())
plotn(plots, test_dataset)
plotn(plots, predictions)> **المهمة**: في مثالنا درّبنا VAE متصلًا بالكامل. خذوا CNN من المشفّر التلقائي التقليدي اللي فوق، وابنوا VAE قائمًا على CNN.
# [المشفّرات التلقائية الخصامية (Adversarial Autoencoders — AAE)](https://arxiv.org/abs/1511.05644)
المشفّرات التلقائية الخصامية هي **مزيج** من الشبكات التوليدية الخصامية (GANs) والمشفّرات التلقائية التغايرية (VAEs).
هنا تكون شبكة الترميز هي المولّد، ويتعلّم المميّز يفرّق بين العينات الحقيقية والعينات الناتجة من توزيع المشفّر. ثم يحاول فاكّ الترميز يعيد بناء الصورة من هالتوزيع.
بهالنهج عندنا **ثلاث دوال خسارة**: خسارة المولّد، وخسارة المميّز من GAN، وخسارة إعادة البناء من VAE.
> **وصف الشكل:** حُذف الأصل لأن حقوق إعادة استخدامه غير موثّقة.
> الصورة مأخوذة من [هالمقالة](https://blog.paperspace.com/adversarial-autoencoders-with-pytorch/) لـ Felipe Ducau
class AAEEncoder(nn.Module):
def __init__(self, input_dim, inter_dim, latent_dim):
super().__init__()
self.linear1 = nn.Linear(input_dim, inter_dim)
self.linear2 = nn.Linear(inter_dim, inter_dim)
self.linear3 = nn.Linear(inter_dim, inter_dim)
self.linear4 = nn.Linear(inter_dim, latent_dim)
self.relu = nn.ReLU()
def forward(self, input):
hidden1 = self.relu(self.linear1(input))
hidden2 = self.relu(self.linear2(hidden1))
hidden3 = self.relu(self.linear3(hidden2))
encoded = self.linear4(hidden3)
return encodedclass AAEDecoder(nn.Module):
def __init__(self, latent_dim, inter_dim, output_dim):
super().__init__()
self.linear1 = nn.Linear(latent_dim, inter_dim)
self.linear2 = nn.Linear(inter_dim, inter_dim)
self.linear3 = nn.Linear(inter_dim, inter_dim)
self.linear4 = nn.Linear(inter_dim, output_dim)
self.relu = nn.ReLU()
self.sigmoid = nn.Sigmoid()
def forward(self, input):
hidden1 = self.relu(self.linear1(input))
hidden2 = self.relu(self.linear2(hidden1))
hidden3 = self.relu(self.linear3(hidden2))
decoded = self.sigmoid(self.linear4(hidden3))
return decodedclass AAEDiscriminator(nn.Module):
def __init__(self, latent_dim, inter_dim):
super().__init__()
self.latent_dim = latent_dim
self.inter_dim = inter_dim
self.linear1 = nn.Linear(latent_dim, inter_dim)
self.linear2 = nn.Linear(inter_dim, inter_dim)
self.linear3 = nn.Linear(inter_dim, inter_dim)
self.linear4 = nn.Linear(inter_dim, inter_dim)
self.linear5 = nn.Linear(inter_dim, 1)
self.relu = nn.ReLU()
self.sigmoid = nn.Sigmoid()
def forward(self, input):
hidden1 = self.relu(self.linear1(input))
hidden2 = self.relu(self.linear2(hidden1))
hidden3 = self.relu(self.linear3(hidden2))
hidden4 = self.relu(self.linear4(hidden3))
decoded = self.sigmoid(self.linear5(hidden4))
return decoded
def get_dims(self):
return self.latent_dim, self.inter_diminput_dims = 784
inter_dims = 1000
latent_dims = 150aae_encoder = AAEEncoder(input_dims, inter_dims, latent_dims).to(device)
aae_decoder = AAEDecoder(latent_dims, inter_dims, input_dims).to(device)
aae_discriminator = AAEDiscriminator(latent_dims, int(inter_dims / 2)).to(device)lr = 1e-4
regularization_lr = 5e-5optim_encoder = optim.Adam(aae_encoder.parameters(), lr=lr)
optim_encoder_regularization = optim.Adam(aae_encoder.parameters(), lr=regularization_lr)
optim_decoder = optim.Adam(aae_decoder.parameters(), lr=lr)
optim_discriminator = optim.Adam(aae_discriminator.parameters(), lr=regularization_lr)def train_aae(dataloaders, models, optimizers, epochs, device):
tqdm_iter = tqdm(range(epochs))
train_dataloader, test_dataloader = dataloaders[0], dataloaders[1]
enc, dec, disc = models[0], models[1], models[2]
optim_enc, optim_enc_reg, optim_dec, optim_disc = optimizers[0], optimizers[1], optimizers[2], optimizers[3]
eps = 1e-9
for epoch in tqdm_iter:
enc.train()
dec.train()
disc.train()
train_reconst_loss = 0.0
train_disc_loss = 0.0
train_enc_loss = 0.0
test_reconst_loss = 0.0
test_disc_loss = 0.0
test_enc_loss = 0.0
for batch in train_dataloader:
imgs, labels = batch
imgs = imgs.view(imgs.shape[0], -1).to(device)
labels = labels.to(device)
enc.zero_grad()
dec.zero_grad()
disc.zero_grad()
encoded = enc(imgs)
decoded = dec(encoded)
reconstruction_loss = F.binary_cross_entropy(decoded, imgs)
reconstruction_loss.backward()
optim_enc.step()
optim_dec.step()
enc.eval()
latent_dim, disc_inter_dim = disc.get_dims()
real = torch.randn(imgs.shape[0], latent_dim).to(device)
disc_real = disc(real)
disc_fake = disc(enc(imgs))
disc_loss = -torch.mean(torch.log(disc_real + eps) + torch.log(1.0 - disc_fake + eps))
disc_loss.backward()
optim_disc.step()
enc.train()
disc_fake = disc(enc(imgs))
enc_loss = -torch.mean(torch.log(disc_fake + eps))
enc_loss.backward()
optim_enc_reg.step()
train_reconst_loss += reconstruction_loss.item()
train_disc_loss += disc_loss.item()
train_enc_loss += enc_loss.item()
enc.eval()
dec.eval()
disc.eval()
with torch.no_grad():
for batch in test_dataloader:
imgs, labels = batch
imgs = imgs.view(imgs.shape[0], -1).to(device)
labels = labels.to(device)
encoded = enc(imgs)
decoded = dec(encoded)
reconstruction_loss = F.binary_cross_entropy(decoded, imgs)
latent_dim, disc_inter_dim = disc.get_dims()
real = torch.randn(imgs.shape[0], latent_dim).to(device)
disc_real = disc(real)
disc_fake = disc(enc(imgs))
disc_loss = -torch.mean(torch.log(disc_real + eps) + torch.log(1.0 - disc_fake + eps))
disc_fake = disc(enc(imgs))
enc_loss = -torch.mean(torch.log(disc_fake + eps))
test_reconst_loss += reconstruction_loss.item()
test_disc_loss += disc_loss.item()
test_enc_loss += enc_loss.item()
train_reconst_loss /= len(train_dataloader)
train_disc_loss /= len(train_dataloader)
train_enc_loss /= len(train_dataloader)
test_reconst_loss /= len(test_dataloader)
test_disc_loss /= len(test_dataloader)
test_enc_loss /= len(test_dataloader)
tqdm_dct = {'train reconst loss:': train_reconst_loss, 'train disc loss:': train_disc_loss, 'train enc loss': train_enc_loss, \
'test reconst loss:': test_reconst_loss, 'test disc loss:': test_disc_loss, 'test enc loss': test_enc_loss}
tqdm_iter.set_postfix(tqdm_dct, refresh=True)
tqdm_iter.refresh()models = (aae_encoder, aae_decoder, aae_discriminator)
optimizers = (optim_encoder, optim_encoder_regularization, optim_decoder, optim_discriminator)train_aae(dataloaders, models, optimizers, epochs, device)aae_encoder.eval()
aae_decoder.eval()
predictions = []
plots = 10
for i, data in enumerate(test_dataset):
if i == plots:
break
pred = aae_decoder(aae_encoder(data[0].to(device).unsqueeze(0).view(1, 784)))
predictions.append(pred.view(1, 28, 28).detach().cpu())
plotn(plots, test_dataset)
plotn(plots, predictions)## مواد إضافية
* [تدوينة NeuroHive](https://neurohive.io/ru/osnovy-data-science/variacionnyj-avtojenkoder-vae/)
* [شرح المشفّر التلقائي التغايري](https://kvfrans.com/variational-autoencoders-explained/)
حذفنا المخرجات وعدّادات التشغيل والودجات والمحتوى النشط وقت الاستيراد. شغّل الدفاتر بس في بيئة خارجية تثق فيها.
سجّل تطبيقك
التسجيل اختياري، يفيدك تتذكر وش طبّقت، ولا يمنع إكمال الدورة.