ژرفا — خودآموز یادگیری عمیق و مهندسی هوش مصنوعی (متوسط)

فصل ۸ از ۹

پیشرفت ترم
۰٪

ترم ۱ · مکانیزمِ یادگیری: PyTorch از پایه

`loss` پایین نمی‌آید: چک‌لیستِ واقعی

فصل ۸پیش‌نمایش رایگان
۱۰ دقیقه مطالعه فصل ۸

در این فصل چه یاد می‌گیری#

این فصل احتمالاً بیشتر از هر فصلِ دیگرِ این ترم به کارت می‌آید. چون شبکه‌ات خراب خواهد شد — و برخلافِ خطای نحوی، بیشترِ این خرابی‌ها هیچ پیامی نمی‌دهند. فقط loss پایین نمی‌آید و تو نمی‌دانی از کجا شروع کنی.

اینجا چهار علتِ رایج را می‌بینی، هر کدام با کدی که واقعاً خراب می‌شود، نشانهٔ مخصوصِ خودش، و راهِ تشخیصش. بعد یک تستِ سی‌ثانیه‌ای که به‌تنهایی نیمی از حالت‌ها را کنار می‌گذارد، و در آخر یک سیاههٔ هشت‌ردیفه که آن چهارتا را هم در خودش دارد و هم چند علتِ دیگر را که کدشان به این فصل نمی‌رسید.

یک تابلوی عیب‌یابی با چند سرنخِ متفاوت که هر کدام به یک مسیر می‌رود

آخر این فصل می‌توانی:

  • انفجارِ گرادیان و nan را تشخیص بدهی و درست کنی
  • اثرِ نرمال‌نکردنِ ورودی را با عدد ببینی
  • خطای خاموشِ شکلِ tensor را بگیری
  • یک ترتیبِ بررسیِ منظم برای هر آموزشِ خراب داشته باشی

قبل از شروع#

از فصل ۷: حلقهٔ استاندارد و train/eval.

از فصل‌های ۲، ۳ و ۶: zero_grad، نرخِ یادگیری، logit.

📓 نوت‌بوک: نوت‌بوک این فصل را در Colab باز کن — همهٔ کدهای این فصل آماده و به‌ترتیب داخلش هست.

۱. علتِ اول: نرخِ یادگیریِ خیلی بزرگ#

import torch
import torch.nn as nn

X = torch.tensor([[1.0], [2.0], [3.0]])
Y = torch.tensor([[2.0], [4.0], [6.0]])
loss_fn = nn.MSELoss()

torch.manual_seed(0)
model = nn.Linear(1, 1)
opt = torch.optim.SGD(model.parameters(), lr=10.0)

for step in range(6):
    loss = loss_fn(model(X), Y)
    opt.zero_grad()
    loss.backward()
    opt.step()
    print(f"قدم {step + 1}: loss = {loss.item():.4g}")
قدم 1: loss = 14.79
قدم 2: loss = 1.762e+05
قدم 3: loss = 2.13e+09
قدم 4: loss = 2.574e+13
قدم 5: loss = 3.11e+17
قدم 6: loss = 3.758e+21

loss هر قدم حدودِ ده هزار برابر می‌شود. چند قدمِ دیگر به inf و بعد nan می‌رسد.

نشانه تشخیص
loss بالا می‌رود، آن هم به‌سرعت نرخِ یادگیری خیلی بزرگ
loss می‌شود nan بعد از چند قدم همان، یا تقسیم بر صفر در محاسبه
loss می‌شود inf همان

اولین کار همیشه یکی است: نرخ را ده برابر کم کن. اگر بازهم بالا رفت، ده برابرِ دیگر. این ارزان‌ترین آزمایشِ ممکن است و بیشترِ اوقات جواب می‌دهد.

۲. علتِ دوم: zero_grad جامانده#

torch.manual_seed(0)
model = nn.Linear(1, 1)
opt = torch.optim.SGD(model.parameters(), lr=0.05)

for step in range(4):
    loss = loss_fn(model(X), Y)
    loss.backward()                                   # zero_grad یادمان رفت
    opt.step()
    print(f"قدم {step + 1}: loss = {loss.item():.4f} | grad انباشته = {model.weight.grad.item():.4f}")
قدم 1: loss = 14.7868 | grad انباشته = -16.5908
قدم 2: loss = 3.0906 | grad انباشته = -24.0478
قدم 3: loss = 2.0401 | grad انباشته = -18.3024
قدم 4: loss = 13.6251 | grad انباشته = -2.6072

اینجا نکتهٔ خطرناک است: loss اول پایین می‌آید. ۱۴٫۸ به ۳٫۱ به ۲٫۰ — همه‌چیز عالی به‌نظر می‌رسد. و بعد ناگهان به ۱۳٫۶ برمی‌گردد.

اگر فقط سه قدم می‌دیدی، فکر می‌کردی همه‌چیز درست است. و در یک آموزشِ واقعی با هزاران قدم، این رفتار به‌صورتِ «loss نوسان می‌کند و هیچ‌وقت پایین نمی‌آید» ظاهر می‌شود.

⚠️ مواظب باش: نشانهٔ این باگ «نوسانِ بی‌قاعده» است، نه «واگرایی». و چون درمانش یک خط است ولی تشخیصش سخت، عادت کن opt.zero_grad() را همیشه اولین خطِ داخلِ حلقه بنویسی — نه آخرین. آن‌وقت فراموش کردنش فوراً به چشم می‌آید.

۳. علتِ سوم: دادهٔ نرمال‌نشده#

raw = torch.tensor([[2015.0, 4.0], [2020.0, 16.0], [2024.0, 32.0], [2018.0, 8.0]])
target = torch.tensor([[10.0], [40.0], [70.0], [22.0]])


def run(data, steps=200, lr=0.01):
    torch.manual_seed(0)
    net = nn.Sequential(nn.Linear(2, 8), nn.ReLU(), nn.Linear(8, 1))
    opt = torch.optim.SGD(net.parameters(), lr=lr)
    fn = nn.MSELoss()
    for _ in range(steps):
        loss = fn(net(data), target)
        opt.zero_grad()
        loss.backward()
        opt.step()
    return loss.item()


normed = (raw - raw.mean(0)) / raw.std(0)
print("دادهٔ خام        :", f"{run(raw):.4g}")
print("دادهٔ نرمال‌شده  :", f"{run(normed):.4g}")
دادهٔ خام        : 511.2
دادهٔ نرمال‌شده  : 0.6372

۵۱۱ در برابرِ ۰٫۶۴ — هشتصد برابر. همان مدل، همان نرخ، همان تعدادِ قدم. تنها تفاوت: دو خط نرمال‌سازی.

چرا این‌قدر اثر دارد؟ ستونِ year عددهایی حولِ ۲۰۲۰ دارد و ستونِ ram حولِ ۱۵. گرادیانِ ستونِ بزرگ‌تر خیلی بزرگ‌تر است، پس نرخی که برای آن مناسب است برای دیگری فاجعه است — و برعکس. با یک نرخِ مشترک، هیچ‌کدام درست آموزش نمی‌بینند.

📏 اندازه بگیر: نرمال‌سازیِ ورودی ارزان‌ترین بهبودِ کلِ یادگیری عمیق است و اولین کاری است که باید بکنی، نه آخرین. و یک قاعدهٔ سخت: میانگین و انحرافِ معیار را فقط از دادهٔ آموزش حساب کن، وگرنه همان نشتیِ سرنخ ترمِ ۳ فصلِ ۴ را ساخته‌ای. ترمِ ۲ فصلِ ۳ کلِ ماجرا را باز می‌کند.

۴. علتِ چهارم: شکلِ اشتباه، بدونِ خطا#

out = torch.tensor([[1.0], [2.0], [3.0], [4.0]])        # شکل (4, 1) — خروجی مدل
bad = torch.tensor([10.0, 20.0, 30.0, 40.0])            # شکل (4,)  — هدف، اشتباه

print("شکل خروجی مدل :", tuple(out.shape))
print("شکل هدف اشتباه:", tuple(bad.shape))
print("out - bad شکل :", tuple((out - bad).shape))
print("MSE با شکل اشتباه:", round(nn.MSELoss()(out, bad).item(), 2))
print("MSE با شکل درست  :", round(nn.MSELoss()(out, bad.unsqueeze(1)).item(), 2))
شکل خروجی مدل : (4, 1)
شکل هدف اشتباه: (4,)
out - bad شکل : (4, 4)
MSE با شکل اشتباه: 632.5
MSE با شکل درست  : 607.5

به سطرِ سوم نگاه کن: تفریقِ یک tensorِ (4,1) از یک (4,) نتیجه‌ای به شکلِ (4,4) داد.

broadcasting هر نمونه را با همهٔ نمونه‌های دیگر مقایسه کرد — شانزده مقایسه به‌جای چهار. lossی که درمی‌آید عددِ بی‌معنایی است که به عددِ درست هم نزدیک است و همین بدترش می‌کند.

در یک batch بزرگ‌تر، تفاوت بزرگ‌تر می‌شود و مدل چیزی کاملاً اشتباه یاد می‌گیرد — بدونِ اینکه هیچ خطایی ببینی.

چک کن: PyTorch در این حالت یک UserWarning می‌دهد که می‌گوید شکلِ هدف با شکلِ ورودی یکی نیست. هشدارها را نادیده نگیر. و قاعدهٔ سختِ خودت: قبل از اولین backward، شکلِ خروجیِ مدل و شکلِ هدف را چاپ کن. پنج ثانیه وقت می‌گیرد و یکی از پرهزینه‌ترین باگ‌های این حوزه را حذف می‌کند.

۵. تستِ طلایی: آیا اصلاً می‌تواند حفظ کند؟#

torch.manual_seed(0)
tiny_x = normed[:2]
tiny_y = target[:2]

net = nn.Sequential(nn.Linear(2, 16), nn.ReLU(), nn.Linear(16, 1))
opt = torch.optim.Adam(net.parameters(), lr=0.05)
fn = nn.MSELoss()

for step in range(400):
    loss = fn(net(tiny_x), tiny_y)
    opt.zero_grad()
    loss.backward()
    opt.step()

print("loss روی دو نمونه بعد از ۴۰۰ قدم:", f"{loss.item():.6f}")
loss روی دو نمونه بعد از ۴۰۰ قدم: 0.000000

این مفیدترین تستِ اشکال‌زداییِ کلِ دوره است، و سی ثانیه وقت می‌گیرد.

منطقش: دو نمونه را بردار و بگذار مدل حفظشان کند. اگر شبکه‌ای با شانزده نورون نتواند دو نقطه را حفظ کند، مشکل قطعاً از داده یا معماری نیست — از خودِ حلقه است: یا optimizer قدم نمی‌زند، یا گرادیان نمی‌رسد، یا شکل‌ها به‌هم ریخته‌اند.

و اگر توانست حفظ کند ولی روی دادهٔ کامل بد است، مشکل جای دیگری است — و همان جایی است که ترمِ ۲ کاملاً دربارهٔ آن است.

۶. ترتیبِ بررسی#

وقتی آموزشت کار نمی‌کند، به این ترتیب برو — از ارزان‌ترین به گران‌ترین:

# بررسی چطور
۱ شکلِ خروجی و هدف یکی است؟ چاپشان کن
۲ loss بالا می‌رود یا nan است؟ نرخ را ده برابر کم کن
۳ zero_grad هست؟ اولین خطِ داخلِ حلقه
۴ ورودی نرمال شده؟ میانگین و انحرافِ هر ستون را چاپ کن
۵ می‌تواند دو نمونه را حفظ کند؟ تستِ بخشِ ۵
۶ برچسب‌ها درست‌اند؟ چند نمونه را با چشم نگاه کن
۷ model.train() هست؟ بالای حلقه
۸ لایهٔ آخر فعال‌سازیِ اضافه ندارد؟ فصلِ ۶ بخشِ ۳

ردیفِ ششم را دستِ‌کم نگیر. جابه‌جا بودنِ برچسب‌ها، یا هم‌راستا نبودنِ ترتیبِ X و y بعد از یک sort، رایج‌تر از چیزی است که فکر می‌کنی — و هیچ ابزاری پیدایش نمی‌کند جز نگاه کردنِ خودت به چند نمونه.

🔧 اگر کار نکرد: سه پیامِ خطای واقعی که در این ترم می‌بینی و معنایشان. Expected all tensors to be on the same device: مدل روی GPU است و داده روی CPU (یا برعکس). CUDA out of memory: batch_size را نصف کن، و اگر باز هم شد، مدل را کوچک‌تر. Expected floating point type for target with class probabilities: به CrossEntropyLoss برچسبِ اعشاری داده‌ای؛ باید int64 باشد.

🤖 از دستیارت بپرس: یک حلقهٔ آموزشِ خودت را که کار می‌کند به دستیارت بده و بگو: «سه باگِ ظریف در این کد بگذار که خطا ندهند ولی آموزش را خراب کنند.» بعد بدونِ نگاه به جوابش، خودت پیدایشان کن. این تمرین از خواندنِ صد صفحه دربارهٔ اشکال‌زدایی مؤثرتر است، چون مهارتِ واقعی، تشخیص است، نه دانستنِ فهرست.

واژه‌های تازهٔ این فصل#

کلمه تلفظ به حروف فارسی یعنی چه
exploding gradient اکسپلودینگ گرِیدینت بزرگ شدنِ مهارنشدنیِ گرادیان‌ها
NaN نَن «عدد نیست» — نتیجهٔ محاسبهٔ نامعتبر
normalization نرمالیزیشن آوردنِ ستون‌ها به مقیاسِ مشترک
sanity check سنیتی چک آزمایشِ کوچکی که سلامتِ کد را ثابت می‌کند
overfit a batch اورفیت اِ بچ تستِ «آیا می‌تواند چند نمونه را حفظ کند»

تمرین‌ها

اول خودت فکر کن یا امتحان کن — بعد اینجا را باز کن.

در فصل بعد#

فصلِ آخرِ ترم: همهٔ این‌ها را روی دادهٔ واقعیِ سرنخ به‌کار می‌بری — همان جدولِ قیمتِ لپ‌تاپ، این‌بار با یک شبکهٔ عصبی. و نتیجه‌ای می‌گیری که احتمالاً انتظارش را نداری، و دقیقاً همان چیزی است که ترمِ ۲ را لازم می‌کند.

به آخر این فصل رسیدی!

اگر ساختی و جواب داد، این دکمه مال توست.