Download train.py from YMmim/object-detection-scratch: direct link, hf CLI and curl.
- Browser
- Download file 6.8 kB
-
https://huggingface.co/YMmim/object-detection-scratch/resolve/main/train.py
- Command line
-
hf download hf://YMmim/object-detection-scratch/train.py
-
curl -L -o train.py https://huggingface.co/YMmim/object-detection-scratch/resolve/main/train.py
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 | |
| # --------------------------------------------------------------- | |
| 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() | |