الوحدة 6 — نماذج الانتشار: التشويش ثمّ تعلّم إزالته
في 2020 أظهرت ورقة DDPM لـHo وJain وAbbeel أنّ فكرة قديمة نسبيًّا (منذ 2015 مع Sohl-Dickstein) يمكن أن تُنافس GAN على جودة الصور. ثمّ في 2022 أنتجت Stable Diffusion وDALL-E 2 وImagen صورًا مذهلة من جمل نصّيّة، وأعادت رسم المشهد. المفاجأة أنّ فكرة النموذج بسيطة عجيبة: نشوّش صورةً تدريجيًّا حتّى نحوّلها إلى ضجيج خالص، ثمّ نعلّم شبكة أن تعكس هذه العمليّة.
الفكرة: مسار أمامي معلوم ومسار خلفي متعلَّم
في الطرف الأمامي، عمليّة بسيطة بلا تعلّم: نأخذ صورة حقيقيّة، ونضيف إليها ضجيجًا طبيعيًّا بجرعات صغيرة في كلّ خطوة، على مدى خطوة (عادةً ). في نهاية المسار، لا يمكن تمييزها عن ضجيج طبيعيّ محض.
الخصيصة الجميلة التي تُبرّر كلّ الرياضيّات التالية: يمكن أخذ عيّنة مباشرة من بمعلومية عبر صيغة مغلقة (بفضل تركيب توزيعات طبيعيّة):
حيث . لسنا مضطرّين لتشغيل المسار خطوةً خطوة أثناء التدريب.
في الطرف الخلفي، الشبكة تُدرَّب لتُنبِّئ الضجيج المضاف بمعلومية الصورة المشوَّشة والخطوة الزمنيّة . الخسارة أبسط ما يمكن تخيّله:
هذا هو DDPM كلّه: خطأ تربيعيّ على الضجيج المتوقّع. مقارنة بلعبة GAN المعقّدة، البساطة صادمة.
حلقة التدريب في PyTorch
import torch
T = 1000
betas = torch.linspace(1e-4, 0.02, T) # جدول $\beta_t$ الخطّي
alphas = 1.0 - betas
alphas_cum = torch.cumprod(alphas, dim=0) # $\bar\alpha_t$
def entrainement_pas(unet, x_0, optimiseur):
B = x_0.size(0)
t = torch.randint(0, T, (B,)) # خطوة زمنيّة عشوائيّة لكلّ صورة
epsilon = torch.randn_like(x_0)
a_bar_t = alphas_cum[t].view(-1, 1, 1, 1)
x_t = torch.sqrt(a_bar_t) * x_0 + torch.sqrt(1 - a_bar_t) * epsilon
epsilon_pred = unet(x_t, t)
perte = torch.mean((epsilon - epsilon_pred) ** 2)
optimiseur.zero_grad()
perte.backward()
optimiseur.step()
return perte.item()
ثلاث ملاحظات جوهرية. الأولى: كلّ صورة في الدفعة تُشوَّش بمعلوم مختلف ، فتتعلّم الشبكة إزالة التشويش في كلّ المستويات. الثانية: الشبكة تعرف لأنّها تُحقَن كتضمين إلى U-Net (كما نُدخل موضع الرمز إلى المحوّل). الثالث: لا يوجد مميّز، ولا لعبة، ولا انهيار أنماط. التدريب مستقرّ بشكل مذهل.