انتقل إلى المحتوى الرئيسي

الوحدة 8 — نقاط الحفظ واستئناف التدريب

تدريب حديث قد يستمرّ ساعات على GPU وحيدة. عطل في الشبكة، انقطاع كهرباء، ذاكرة تنفد قبل النهاية: كلّها أحداث واردة. نقطة الحفظ هي ملفّ يحتوي كلّ ما يلزم لاستئناف التدريب في الحقبة نفسها التي توقّف عندها، بلا انحراف رقميّ. هذه الوحدة تُبيّن ما نحفظ، وكيف، وكيف نرجع من جديد.

ما الذي نحفظه، ولماذا كلّ ذلك

قد يظنّ المبتدئ أنّ حفظ modele.state_dict() كافٍ. هذا خطأ: يفتقد نصف المعلومات اللازمة لاستئناف مطابق. القائمة الكاملة:

العنصرما يحملهنتيجة النسيان
model.state_dict()الأوزانتدريب من الصفر
optimizer.state_dict()لحظات Adam، الاندفاعحساب خاطئ للتحديث الأوّل
scheduler.state_dict()معدّل التعلّم الحالي، رقم الخطوةجدولة تعود لبدايتها
scaler.state_dict()عامل تضخيم GradScalerذبذبة قصيرة قبل الاستقرار
epochآخر حقبة أُنجِزتإعادة حقب سابقة
torch.get_rng_state()حالة مولّد الأعداد العشوائيةتباين نتائج التحقّق
best_metricأفضل مقياس رأيتهفقدان معيار «الأفضل»

نقطة حفظ سليمة تجمع كلّ هذا في قاموس واحد.

بناء نقطة حفظ كاملة

نُغلّف الحفظ في دالّة مُنَظَّفة:

def enregistrer(chemin, modele, opt, plan, scaler, epoque, meilleur):
torch.save({
"epoque": epoque,
"modele": modele.state_dict(),
"opt": opt.state_dict(),
"plan": plan.state_dict(),
"scaler": scaler.state_dict(),
"rng": torch.get_rng_state(),
"rng_cuda": torch.cuda.get_rng_state_all(),
"meilleur": meilleur,
"version_torch": torch.__version__,
}, chemin)

نُضيف version_torch لأنّ صيغة state_dict تتبدّل نادرًا عبر الإصدارات، ومعرفة الإصدار الأصليّ تفيد في التصحيح. لا يستهلك ذلك مساحةً تُذكَر.

دالّة التحميل المقابلة

الرجوع بنفس الترتيب:

def reprendre(chemin, modele, opt, plan, scaler, appareil):
etat = torch.load(chemin, map_location=appareil)
modele.load_state_dict(etat["modele"])
opt.load_state_dict(etat["opt"])
plan.load_state_dict(etat["plan"])
scaler.load_state_dict(etat["scaler"])
torch.set_rng_state(etat["rng"])
if torch.cuda.is_available() and "rng_cuda" in etat:
torch.cuda.set_rng_state_all(etat["rng_cuda"])
return etat["epoque"], etat["meilleur"]

نلاحظ map_location وقيمته المهمّة: ملفّ حُفظ على GPU يستحيل تحميله مباشرةً على وحدة معالجة مركزية دون هذا الوسيط. نضبطه على جهاز الاستقبال الفعليّ، فيتولّى PyTorch الترجمة.

استئناف من الحقبة الأخيرة

نُنسّق الحلقة كي تعرف كيف تستأنف:

epoque_debut = 1
meilleur = float("inf")

if os.path.exists("dernier.pt"):
epoque_debut, meilleur = reprendre("dernier.pt", modele, opt, plan, scaler, appareil)
epoque_debut += 1 # نبدأ من الحقبة التالية
print(f"استؤنف التدريب من الحقبة {epoque_debut}")

for epoque in range(epoque_debut, 21):
entrainer_une_epoque()
perte_val, exact_val = evaluer(modele, charg_val, critere, appareil)

enregistrer("dernier.pt", modele, opt, plan, scaler, epoque, meilleur)
if perte_val < meilleur:
meilleur = perte_val
enregistrer("meilleur.pt", modele, opt, plan, scaler, epoque, meilleur)

نحفظ دائمًا الحالة الأخيرة، ونحفظ بشرط الأفضلَ. الملفّان مقصودان: الأوّل يسمح بالمواصلة، والثاني يسمح باسترجاع أفضل مرور من دون فقدان القدرة على المتابعة.

أفضل حقبة أم آخر حقبة؟

سؤال شائع، جوابه يعتمد على القرار الذي تتخذه:

  • للنشر في الإنتاج: meilleur.pt — نريد أعلى دقّة تحقّق رأيناها.
  • لمواصلة التجربة أو التعديل: dernier.pt — نريد استمرار الاندفاع والجدولة كما هما.

هذه الازدواجية هي نفسها التي وردت في دورة 08 تحت اسم save_best_only مع restore_best_weights. الفكرة عابرة للأُطر.

توافق بين إصدارات PyTorch

state_dict مستقرّ إلى حدّ بعيد. الحفظ في PyTorch 2.1 والتحميل في 2.3 يعمل غالبًا بلا حزّة. مع ذلك، بعض النماذج المعقّدة قد تلقى تغييرات دقيقة في أسماء الوحدات الفرعية. عادتان دفاعيّتان:

  • ثبِّت الإصدار في requirements.txt بدل السماح بـ>=.
  • عند التحميل عبر إصدار مختلف، حمِّل بـstrict=False لعلمك بالأسماء المفقودة:
result = modele.load_state_dict(etat["modele"], strict=False)
print(result.missing_keys, result.unexpected_keys)

القائمتان يجب أن تكونا فارغتين. إن ظهرت مفاتيح غير متوقّعة، فتلك إشارة إلى تغيير معماريّ يستحقّ التحقّق.

استئناف عبر أجهزة مختلفة

سيناريو واقعيّ: نبدأ التدريب على GPU، ونستأنف على وحدة معالجة مركزية لتشخيص عابر. map_location="cpu" يُتيح ذلك بلا تعديل:

etat = torch.load("dernier.pt", map_location="cpu")
modele = ClassifieurMode()
modele.load_state_dict(etat["modele"])
# أعد النقل حين تعود إلى GPU
modele.to("cuda")

قاعدة الأمان: لا تحمِّل ملفًّا مصدره غير موثوق. torch.save يستعمل pickle بايثون، القابل لتنفيذ تعليمات عشوائية عند التحميل. هذا الخطر معروف رسميًّا في وثائق PyTorch، وهو خارج نطاق نقاط الحفظ الخاصّة بك، لكن يستحقّ الإدراك.

نموذج مغلَّف بـtorch.compile

torch.compile يُنشئ التفافًا حول النموذج. أسماء المفاتيح في state_dict تُصبح مسبوقة بـ_orig_mod.، وهذا يُربك التحميل حين نُريد نقل الأوزان إلى نموذج غير مُترجَم. حلّ نظيف:

def poids_purs(m):
"""يُعيد state_dict بلا سابقة _orig_mod. الناتجة عن torch.compile."""
from collections import OrderedDict
ordonne = OrderedDict()
for k, v in m.state_dict().items():
nk = k.removeprefix("_orig_mod.") if k.startswith("_orig_mod.") else k
ordonne[nk] = v
return ordonne

نحفظ الشكل النقيّ، فيبقى الملفّ قابلاً للاستعمال في أيّ سياق.

حجم الملفّ وتناوب النسخ

نموذج بمليون معلمة يزن نحو 4 ميغابايت في float32. شبكة أكبر تصل إلى مئات الميغابايت. لا تحفظ حقبةً بعد حقبة إلى ما لا نهاية:

import os

# حافظ فقط على آخر ثلاث نقاط
if epoque > 3 and os.path.exists(f"epoque_{epoque - 3}.pt"):
os.remove(f"epoque_{epoque - 3}.pt")

قاعدة معقولة: dernier.pt (يُستبدَل)، meilleur.pt (يُستبدَل شرطًا)، وربّما نقاط عند بعض علامات مهمّة (نصف التدريب، ثلثه). لا تُترك عشرات النسخ تسدّ القرص.

مثال متكامل: الحلقة مع نقاط الحفظ

نجمع كلّ ما سبق في تعديل بسيط لحلقة الوحدة السابعة:

epoque_debut, meilleur = 1, float("inf")
if os.path.exists("dernier.pt"):
epoque_debut, meilleur = reprendre("dernier.pt", modele, opt, plan, scaler, appareil)
epoque_debut += 1

for epoque in range(epoque_debut, 21):
modele.train()
for lot_x, lot_y in charg_tr:
lot_x = lot_x.to(appareil, non_blocking=True)
lot_y = lot_y.to(appareil, non_blocking=True)
opt.zero_grad()
with torch.cuda.amp.autocast():
perte = critere(modele(lot_x), lot_y)
scaler.scale(perte).backward()
scaler.step(opt); scaler.update()
plan.step()

perte_val, _ = evaluer(modele, charg_val, critere, appareil)
enregistrer("dernier.pt", modele, opt, plan, scaler, epoque, meilleur)
if perte_val < meilleur:
meilleur = perte_val
enregistrer("meilleur.pt", modele, opt, plan, scaler, epoque, meilleur)

النتيجة: أيّ توقّف الآن يُكلّف حقبةً واحدةً على أقصى تقدير، لا التدريب كلّه.

torch.save على المسارات لا الأشياء

عادة سيّئة رأيتها في مشاريع كثيرة: torch.save(modele, ...) بدون state_dict. هذا يحفظ الكائن كاملاً بما فيه معمارية الصنف، ما يجعل التحميل يتطلّب استيراد الملفّ الأصليّ للصنف. تغيير مسار الملفّ يكسر التحميل بصمت. القاعدة: احفظ state_dict دومًا، وأعِد بناء النموذج بنفسك قبل التحميل.

اختبر الاستئناف قبل الحاجة إليه

بعد إنجاز أوّل ثلاث حقب، أوقف التدريب يدويًّا وأعِد تشغيله. راقب: هل تعود الجدولة إلى القيمة الصحيحة؟ هل يستمرّ عدّاد الحقب من مكانه؟ هل تُطابق خسارة الحقبة الرابعة ما كنت ستراه في تدريب متواصل؟ إن كان الجواب لا في إحدى النقاط، صحّح الآن — لا أثناء انقطاع فعليّ في منتصف الليل.

في الخلاصة

  • نقطة حفظ سليمة تحوي الأوزان، المُحسِّن، الجدولة، GradScaler، الحقبة، البذرة، الأفضل؛ نسيان أيّ منها يُنتج استئنافًا غير مطابق.
  • map_location يسمح بالتحميل على جهاز مختلف عن الأصل؛ لا تحمِّل ملفًّا مصدره غير موثوق.
  • meilleur.pt للنشر، dernier.pt للمواصلة؛ حافظ على النسختين ولا تتراكم النسخ إلى ما لا نهاية.
  • نموذج مغلَّف بـtorch.compile يُضيف بادئة _orig_mod. في المفاتيح؛ استخرج الشكل النقيّ قبل الحفظ لتفادي إرباك المستقبل.

الوحدة التالية: التعلّم بالنقل، حيث نستبدل شبكتنا بشبكة ResNet18 مُدرَّبة مسبقًا وتحوّل مشروع Fashion-MNIST إلى مسألة أشدّ طموحًا.