بعد تدريب نموذج (في Training Loop)، الخطوة التالية هي حفظه ليستخدمه تطبيق آخر لاحقًا، أو لاستئناف التدريب من حيث توقّف.
فصل مهمّ: المعمارية vs الأوزان
المعمارية (architecture) = البنية (أي طبقات بأي أحجام)
الأوزان (weights) = القيم العددية المتعلّمة
state_dict() يحفظ الأوزان فقط. المعمارية تبقى في كود Python (class MyModel(nn.Module)).
import torch
import torch.nn as nn
class TinyModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(3, 2)
def forward(self, x):
return self.fc(x)
model = TinyModel()
print(model.state_dict())
OrderedDict([
('fc.weight', tensor([[...]])),
('fc.bias', tensor([...])),
])
الحفظ الأساسي
# حفظ
torch.save(model.state_dict(), "model_weights.pt")
# تحميل
model2 = TinyModel() # أنشئ النموذج مرّة أخرى
model2.load_state_dict(torch.load("model_weights.pt"))
model2.eval() # للاستخدام في inference
مثال كامل محقّق
import torch
import torch.nn as nn
class TinyModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(3, 2)
def forward(self, x):
return self.fc(x)
torch.manual_seed(42)
model = TinyModel()
w_before = model.fc.weight.detach().clone()
# === حفظ ===
torch.save(model.state_dict(), "model.pt")
# === استرجاع ===
model2 = TinyModel() # أوزان عشوائية جديدة
model2.load_state_dict(torch.load("model.pt"))
model2.eval()
# التحقّق
print("الأوزان متطابقة:", torch.equal(model.fc.weight, model2.fc.weight))
x = torch.randn(5, 3)
print("المخرجات متطابقة:", torch.allclose(model(x), model2(x)))
النتيجة:
الأوزان متطابقة: True
المخرجات متطابقة: True
state_dict — ماذا يُحفظ؟
for name, param in model.state_dict().items():
print(f"{name:20s} {tuple(param.shape)}")
fc.weight (2, 3)
fc.bias (2,)
فقط الـ Parameters المسجّلة (والـ Buffers مثل running_mean لـ BatchNorm).
حفظ النموذج كاملًا (أبسط لكن أقلّ مرونة)
# حفظ كامل (المعمارية + الأوزان)
torch.save(model, "model_full.pt")
# تحميل
model2 = torch.load("model_full.pt")
تحذير: يحفظ المعمارية مع pickle، وهذا يعتمد على بنية الكلاس. لا ينصح به عمليًا (إذا غيّرت الكود، قد لا يُحمَّل). استخدم
state_dictدائمًا إلا لحفظ النماذج النهائية في تطبيقات بسيطة.
Checkpoint كامل (للاستئناف)
# === حفظ checkpoint ===
torch.save({
"epoch": epoch,
"model_state_dict": model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"loss": loss,
"best_val_acc": 0.95,
}, "checkpoint.pt")
# === تحميل checkpoint ===
checkpoint = torch.load("checkpoint.pt", weights_only=False)
model.load_state_dict(checkpoint["model_state_dict"])
optimizer.load_state_dict(checkpoint["optimizer_state_dict"])
start_epoch = checkpoint["epoch"] + 1
best_val_acc = checkpoint["best_val_acc"]
weights_only=False: ضروري عند تحميل checkpoint كامل لأنoptimizer.state_dict()ليس مجرد tensors. PyTorch يحذّر افتراضيًا — هذا التحذير متوقّع وآمن إذا كنت تثق في مصدر الملف.
حفظ أفضل نموذج أثناء التدريب
best_val_loss = float("inf")
for epoch in range(max_epochs):
# ... تدريب ...
val_loss = evaluate(model, val_loader)
if val_loss < best_val_loss:
best_val_loss = val_loss
torch.save(model.state_dict(), "best.pt")
print(f"saved best at epoch {epoch}")
تفصيل أفضل في Optimizers و Regularization.
أمان torch.load — اقرأ هذا
torch.load يستخدم داخليًا pickle (Python serialization). pickle يستطيع تنفيذ أي كود Python أثناء التحميل.
القاعدة:
- لا تحمّل ملفات
.pt/.pthمن مصادر غير موثوقة — تعامل معها مثل تشغيل كود Python من إنترنت. - إذا كنت متأكّدًا أن الملف من عملك، يمكنك تجاهل تحذير
weights_only. - PyTorch 2.6+ يتغيّر الافتراضي نحو
weights_only=Trueللأمان.
الوضع الحالي
في PyTorch 2.x، torch.load يقبل وسيط weights_only:
# حفظ/تحميل tensors فقط (آمن من pickle deserialization)
torch.save(model.state_dict(), "weights.pt")
state_dict = torch.load("weights.pt", weights_only=True)
# عند الحاجة إلى checkpoint كامل (optimizer state, custom objects)
checkpoint = torch.load("checkpoint.pt", weights_only=False)
توصية: إذا كنت تحفظ أوزانًا فقط، استخدم
weights_only=True. إذا كنت تحفظ checkpoint كاملًا، فأنت تثق في الملف بنفسك.
INFEMore: خريطة الجهاز (device map)
عند تحميل ملف محفوظ على GPU إلى CPU (أو العكس):
# تحميل إلى نفس الجهاز الذي يحفظ
state_dict = torch.load("model.pt", map_location="cpu")
model.load_state_dict(state_dict)
# أو
state_dict = torch.load("model.pt", map_location=torch.device("cpu"))
# تحميل من GPU إلى CPU
state_dict = torch.load("model_gpu.pt", map_location="cpu")
# تحميل من CPU إلى GPU
state_dict = torch.load("model.pt", map_location="cuda")
model.to("cuda")
model.load_state_dict(state_dict)
صيغ بديلة (للتطبيقات الإنتاجية)
في الإنتاج، قد تفضّل:
| الصيغة | الاستخدام |
|---|---|
state_dict (.pt) | مرن، اعتمادي على كود Python |
TorchScript (torch.jit.save) | قابل للنشر بدون كود Python الأصلي |
| ONNX | متعدد اللغات (C++, Rust, mobile) |
| Safetensors (HuggingFace) | سريع، آمن من pickle |
# TorchScript (مثال مختصر)
scripted = torch.jit.script(model)
scripted.save("model_scripted.pt")
# ONNX (يحتاج onnx package)
dummy = torch.randn(1, 3)
torch.onnx.export(model, dummy, "model.onnx")
هذه خارج نطاق هذا الدرس للمبتدئين — احفظها كـ state_dict للآن.
مثال كامل محقّق: التدريب + checkpoint + استئناف
import torch
import torch.nn as nn
from torch.utils.data import TensorDataset, DataLoader
torch.manual_seed(0)
X = torch.randn(200, 4)
y = (X[:, 0] + X[:, 1] > 0).long()
ds = TensorDataset(X, y)
loader = DataLoader(ds, batch_size=16, shuffle=True)
model = nn.Sequential(nn.Linear(4, 16), nn.ReLU(), nn.Linear(16, 2))
optimizer = torch.optim.Adam(model.parameters(), lr=1e-2)
criterion = nn.CrossEntropyLoss()
# تدريب قصير
for epoch in range(3):
for xb, yb in loader:
optimizer.zero_grad()
loss = criterion(model(xb), yb)
loss.backward()
optimizer.step()
# === حفظ checkpoint ===
torch.save({
"model_state_dict": model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"loss": loss.item(),
}, "ckpt.pt")
# === محاكاة: إعادة إنشاء النموذج (training جديد) ===
model_new = nn.Sequential(nn.Linear(4, 16), nn.ReLU(), nn.Linear(16, 2))
opt_new = torch.optim.Adam(model_new.parameters(), lr=1e-2)
ckpt = torch.load("ckpt.pt", weights_only=False)
model_new.load_state_dict(ckpt["model_state_dict"])
opt_new.load_state_dict(ckpt["optimizer_state_dict"])
print("loss المحفوظ:", ckpt["loss"])
# تحقّق من تطابق المخرجات
x_test = torch.randn(5, 4)
print("مخرجات متطابقة:", torch.allclose(model(x_test), model_new(x_test)))
أخطاء شائعة
torch.load("model.pt")بدونweights_only=Trueللملفات غير الموثوقة: خطورة أمنية. افتراضيًا PyTorch 2.6+ سيطلب صراحة.- نسيان
model.eval()بعد التحميل: الطبقات مثل Dropout و BatchNorm تبقى في وضع التدريب. load_state_dictيفشل بـMissing keysأوUnexpected keys: المعمارية تغيّرت. تأكّد أنّmodelله نفس البنية.- حفظ المعمارية مع
torch.save(model, ...)ثم تغيير الـ class: قد لا يُحمَّل أبدًا. استخدمstate_dict. - نسيان
map_locationعند النقل بين GPU/CPU: RuntimeError إذا كانت الأوزان على CUDA والجهاز CPU. - عدم حفظ optimizer state: عند الاستئناف، التعلّم يتباطأ لأن optimizer يبدأ من جديد (مهم بشكل خاص لـ Adam).
الخطوات التالية
- Training Loop في PyTorch — أين تُدمج checkpointing.
- PyTorch: Model و nn.Module — معمارية النموذج.
- مشروع: مصنّف صور — تطبيق كامل مع save/load.
🔍 لحفظ النماذج بصيغة آمنة وقابلة للتشغيل المتبادل، اطّلع على
safetensors(مذكور للاكتمال).