| """ |
| Faster R-CNN ๋ฐ๋ฐ๋ฅ ๊ตฌํ โ [5/5] ํ์ต + ํ๊ฐ |
| ================================================ |
| ์ ์ฒด ํ์ดํ๋ผ์ธ: |
| ํ์ต: ์ด๋ฏธ์ง โ ๋ชจ๋ธ(train) โ RPN ์์ค + RoI ์์ค โ ์ญ์ ํ |
| ํ๊ฐ: ์ด๋ฏธ์ง โ ๋ชจ๋ธ(eval) โ ํ์ง ๊ฒฐ๊ณผ โ mAP@0.5 ๊ณ์ฐ |
| |
| ์คํ ์: |
| python train.py --voc_root /path/VOCdevkit/VOC2007 --epochs 12 |
| python train.py --voc_root /path/VOCdevkit/VOC2007 --eval_only --ckpt frcnn.pth |
| |
| ์ฃผ์: |
| - batch_size=1 ๋ก ์ค๊ณ(์ด๋ฏธ์ง ํฌ๊ธฐ๊ฐ ์ ๊ฐ๊ฐ์ด๋ผ ๋จ์ํ). |
| - GPU ๊ถ์ฅ. CPU๋ก๋ ๋์ง๋ง ๋งค์ฐ ๋๋ฆฌ๋ค. |
| """ |
|
|
| import argparse |
| import torch |
| from torch.utils.data import DataLoader |
|
|
| from dataset import VOCDataset, collate_fn, NUM_CLASSES, VOC_CLASSES |
| from model import FasterRCNN |
| from losses import rpn_loss, roi_loss |
| from box_utils import box_iou |
|
|
|
|
| |
| |
| |
| def train_one_epoch(model, loader, optimizer, device, epoch): |
| model.train() |
| running = 0.0 |
| for i, (imgs, targets) in enumerate(loader): |
| img = imgs[0].to(device).unsqueeze(0) |
| gt_boxes = targets[0]["boxes"].to(device) |
| gt_labels = targets[0]["labels"].to(device) |
| if gt_boxes.numel() == 0: |
| continue |
|
|
| out = model(img) |
| img_hw = img.shape[-2:] |
|
|
| l_rpn = rpn_loss(out["rpn_logits"], out["rpn_deltas"], |
| out["anchors"], gt_boxes, img_hw) |
| l_roi = roi_loss(model.head, out["feat"], out["stride"], |
| out["proposals"], gt_boxes, gt_labels) |
| loss = l_rpn + l_roi |
|
|
| optimizer.zero_grad() |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 10.0) |
| optimizer.step() |
|
|
| running += loss.item() |
| if (i + 1) % 100 == 0: |
| print(f"[epoch {epoch}] iter {i+1}/{len(loader)} " |
| f"loss {running/(i+1):.4f} (rpn {l_rpn.item():.3f} roi {l_roi.item():.3f})") |
| return running / max(1, len(loader)) |
|
|
|
|
| |
| |
| |
| @torch.no_grad() |
| def evaluate(model, loader, device, iou_thresh=0.5): |
| model.eval() |
| |
| preds = {c: [] for c in range(1, NUM_CLASSES)} |
| n_gt = {c: 0 for c in range(1, NUM_CLASSES)} |
|
|
| for imgs, targets in loader: |
| img = imgs[0].to(device).unsqueeze(0) |
| det = model(img) |
| gt_boxes = targets[0]["boxes"].to(device) |
| gt_labels = targets[0]["labels"].to(device) |
|
|
| for c in range(1, NUM_CLASSES): |
| gmask = gt_labels == c |
| gboxes = gt_boxes[gmask] |
| n_gt[c] += gboxes.shape[0] |
|
|
| pmask = det["labels"] == c |
| pboxes = det["boxes"][pmask] |
| pscores = det["scores"][pmask] |
| if pboxes.numel() == 0: |
| continue |
|
|
| order = pscores.argsort(descending=True) |
| pboxes, pscores = pboxes[order], pscores[order] |
|
|
| matched = torch.zeros(gboxes.shape[0], dtype=torch.bool) |
| for k in range(pboxes.shape[0]): |
| if gboxes.numel() == 0: |
| preds[c].append((pscores[k].item(), 0)) |
| continue |
| ious = box_iou(pboxes[k:k+1], gboxes).squeeze(0) |
| best_iou, best_j = ious.max(0) |
| if best_iou >= iou_thresh and not matched[best_j]: |
| preds[c].append((pscores[k].item(), 1)) |
| matched[best_j] = True |
| else: |
| preds[c].append((pscores[k].item(), 0)) |
|
|
| |
| aps = [] |
| for c in range(1, NUM_CLASSES): |
| ap = _voc_ap(preds[c], n_gt[c]) |
| aps.append(ap) |
| print(f" {VOC_CLASSES[c-1]:12s} AP = {ap:.4f}") |
| mAP = sum(aps) / len(aps) |
| print(f" {'mAP@0.5':12s} = {mAP:.4f}") |
| return mAP |
|
|
|
|
| def _voc_ap(pred_list, n_gt): |
| """(์ ์, TP์ฌ๋ถ) ๋ชฉ๋ก์ผ๋ก precision-recall ๊ณก์ ์๋ ๋์ด(AP) ๊ณ์ฐ.""" |
| if n_gt == 0 or len(pred_list) == 0: |
| return 0.0 |
| pred_list.sort(key=lambda x: x[0], reverse=True) |
| tp = torch.tensor([p[1] for p in pred_list], dtype=torch.float32) |
| fp = 1 - tp |
| tp_cum = torch.cumsum(tp, 0) |
| fp_cum = torch.cumsum(fp, 0) |
| recall = tp_cum / n_gt |
| precision = tp_cum / (tp_cum + fp_cum).clamp(min=1e-6) |
|
|
| |
| ap = 0.0 |
| for t in torch.linspace(0, 1, 11): |
| mask = recall >= t |
| p = precision[mask].max().item() if mask.any() else 0.0 |
| ap += p / 11.0 |
| return ap |
|
|
|
|
| |
| |
| |
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--voc_root", required=True, help="VOCdevkit/VOC2007 ๊ฒฝ๋ก") |
| ap.add_argument("--epochs", type=int, default=12) |
| ap.add_argument("--lr", type=float, default=1e-3) |
| ap.add_argument("--ckpt", default="frcnn.pth") |
| ap.add_argument("--eval_only", action="store_true") |
| args = ap.parse_args() |
|
|
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| print("device:", device) |
|
|
| model = FasterRCNN(NUM_CLASSES).to(device) |
|
|
| if args.eval_only: |
| model.load_state_dict(torch.load(args.ckpt, map_location=device)) |
| test_ds = VOCDataset(args.voc_root, split="test") |
| test_loader = DataLoader(test_ds, batch_size=1, shuffle=False, |
| collate_fn=collate_fn, num_workers=4) |
| evaluate(model, test_loader, device) |
| return |
|
|
| |
| train_ds = VOCDataset(args.voc_root, split="trainval") |
| train_loader = DataLoader(train_ds, batch_size=1, shuffle=True, |
| collate_fn=collate_fn, num_workers=4) |
|
|
| params = [p for p in model.parameters() if p.requires_grad] |
| optimizer = torch.optim.SGD(params, lr=args.lr, momentum=0.9, |
| weight_decay=5e-4) |
| scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=8, gamma=0.1) |
|
|
| for epoch in range(1, args.epochs + 1): |
| avg = train_one_epoch(model, train_loader, optimizer, device, epoch) |
| scheduler.step() |
| print(f"[epoch {epoch}] avg loss = {avg:.4f}") |
| torch.save(model.state_dict(), args.ckpt) |
| print(f" checkpoint saved โ {args.ckpt}") |
|
|
| print("ํ์ต ์๋ฃ. ํ๊ฐํ๋ ค๋ฉด --eval_only ๋ก ์คํํ์ธ์.") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|