YMmim's picture
Object detection from scratch: Faster R-CNN + YOLO comparison
017c046 verified
Raw
History Blame Contribute Delete
6.8 kB
"""
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) # [1,3,H,W]
gt_boxes = targets[0]["boxes"].to(device)
gt_labels = targets[0]["labels"].to(device)
if gt_boxes.numel() == 0:
continue
out = model(img) # training=True โ†’ ์ค‘๊ฐ„ ์‚ฐ์ถœ๋ฌผ ๋ฐ˜ํ™˜
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))
# ---------------------------------------------------------------
# ํ‰๊ฐ€: VOC ์Šคํƒ€์ผ mAP@0.5
# ---------------------------------------------------------------
@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) # eval โ†’ {boxes, labels, scores}
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)) # TP
matched[best_j] = True
else:
preds[c].append((pscores[k].item(), 0)) # FP
# ํด๋ž˜์Šค๋ณ„ AP โ†’ mAP
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)
# 11-point ๋ณด๊ฐ„ (VOC2007 ๋ฐฉ์‹)
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()