File size: 6,464 Bytes
7902c8d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 | """
models/train.py
Classe Trainer pour PyTorch.
Fonctionnalitรฉs : early stopping, ReduceLROnPlateau,
sauvegarde du meilleur modรจle, courbes train/val.
"""
import torch
import torch.nn as nn
from tqdm import tqdm
import matplotlib.pyplot as plt
class Trainer:
def __init__(self, model, train_dataloader, test_dataloader,
lr=1e-3, epochs=30, device="cpu", patience=5):
self.model = model
self.train_dataloader = train_dataloader
self.test_dataloader = test_dataloader
self.epochs = epochs
self.patience = patience
self.device = device
self.criterion = nn.CrossEntropyLoss()
self.optimizer = torch.optim.Adam(model.parameters(), lr=lr)
self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
self.optimizer, mode="min",
factor=0.5, patience=3)
# โโ Entraรฎnement complet โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
def train(self, save_path=None, plot=False):
self.train_loss, self.train_acc = [], []
self.val_loss, self.val_acc = [], []
best_val_loss = float("inf")
epochs_no_improve = 0
best_state = None
for epoch in range(self.epochs):
tr_loss, tr_acc = self._train_one_epoch(epoch)
v_loss, v_acc = self._validate()
self.train_loss.append(tr_loss)
self.train_acc.append(tr_acc)
self.val_loss.append(v_loss)
self.val_acc.append(v_acc)
self.scheduler.step(v_loss)
lr = self.optimizer.param_groups[0]["lr"]
print(f"Epoch {epoch+1:02d}/{self.epochs} "
f"| Train loss={tr_loss:.4f} acc={tr_acc:.2f}% "
f"| Val loss={v_loss:.4f} acc={v_acc:.2f}% "
f"| LR={lr:.2e}")
# Early stopping
if v_loss < best_val_loss:
best_val_loss = v_loss
epochs_no_improve = 0
best_state = {k: v.clone() for k, v in self.model.state_dict().items()}
if save_path:
torch.save(best_state, save_path)
print(f" โ Best model saved (val_loss={v_loss:.4f})")
else:
epochs_no_improve += 1
print(f" โ No improvement {epochs_no_improve}/{self.patience}")
if epochs_no_improve >= self.patience:
print(f"\nโ Early stopping at epoch {epoch+1}")
break
if best_state:
self.model.load_state_dict(best_state)
if plot:
self.plot_history()
# โโ Une epoch de train โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
def _train_one_epoch(self, epoch):
self.model.train()
total_loss, total_correct, total_samples = 0, 0, 0
pbar = tqdm(self.train_dataloader,
desc=f"Epoch {epoch+1}/{self.epochs} [train]", leave=False)
for imgs, labels in pbar:
imgs, labels = imgs.to(self.device), labels.to(self.device)
self.optimizer.zero_grad()
out = self.model(imgs)
loss = self.criterion(out, labels)
loss.backward()
self.optimizer.step()
_, preds = out.max(1)
correct = (preds == labels).sum().item()
total = labels.size(0)
total_correct += correct
total_samples += total
total_loss += loss.item()
pbar.set_postfix({
"Batch Acc": f"{100.*correct/total:.1f}%",
"Avg Acc": f"{100.*total_correct/total_samples:.1f}%",
"Loss": f"{total_loss/total_samples:.4f}",
})
return total_loss / total_samples, 100. * total_correct / total_samples
# โโ Validation โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
@torch.no_grad()
def _validate(self):
self.model.eval()
total_loss, total_correct, total_samples = 0, 0, 0
for imgs, labels in self.test_dataloader:
imgs, labels = imgs.to(self.device), labels.to(self.device)
out = self.model(imgs)
loss = self.criterion(out, labels)
_, preds = out.max(1)
total_correct += (preds == labels).sum().item()
total_samples += labels.size(0)
total_loss += loss.item() * labels.size(0)
return total_loss / total_samples, 100. * total_correct / total_samples
# โโ รvaluation finale (public) โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
@torch.no_grad()
def evaluate(self):
loss, acc = self._validate()
print(f"\nTest Accuracy : {acc:.2f}% | Test Loss : {loss:.4f}")
return acc, loss
# โโ Courbes โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
def plot_history(self, save_path="/kaggle/working/history_pytorch.png"):
epochs = range(1, len(self.train_loss) + 1)
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 5))
ax1.plot(epochs, self.train_loss, label="Train", color="tab:blue")
ax1.plot(epochs, self.val_loss, label="Val", color="tab:orange")
ax1.set_title("Loss"); ax1.set_xlabel("Epoch")
ax1.legend(); ax1.grid(alpha=.3)
ax2.plot(epochs, self.train_acc, label="Train", color="tab:blue")
ax2.plot(epochs, self.val_acc, label="Val", color="tab:orange")
ax2.set_title("Accuracy (%)"); ax2.set_xlabel("Epoch")
ax2.legend(); ax2.grid(alpha=.3)
fig.suptitle("Training History โ PyTorch", fontsize=13)
fig.tight_layout()
plt.savefig(save_path, dpi=120)
plt.show()
print(f"โ Courbes sauvegardรฉes โ {save_path}")
|