Faza 04 · lecția 08
Segmentarea instanțelor — Mask R-CNN
Scopul lecției: Adăugați o mică ramură pentru măști la un detector Faster R-CNN și obțineți segmentarea instanțelor. Partea dificilă este RoIAlign, iar ea este mai complicată decât pare.
Versiunea curentă AlexBred.com: primele 100 de lecții ale programului în limba română.
Cuprinsul lecției
- Obiective de învățare
- Problema
- Conceptul
- Arhitectura
- De ce RoIAlign, nu RoIPool
- RPN într-un paragraf
- Capul pentru măști
- Pierderi
- Formatul ieșirii
- Construiți
- Pasul 1: RoIAlign de la zero
- Pasul 2: Comparați cu RoIAlign din torchvision
- Pasul 3: Încărcați un Mask R-CNN preantrenat
- Pasul 4: Rulați inferența
- Pasul 5: Înlocuiți capetele pentru un număr de clase personalizat
- Pasul 6: Înghețați ceea ce nu trebuie antrenat
- Folosiți
- Livrați
- Exerciții
- Termeni-cheie
- Lecturi suplimentare
Adăugați o mică ramură pentru măști la un detector Faster R-CNN și obțineți segmentarea instanțelor. Partea dificilă este RoIAlign, iar ea este mai complicată decât pare.
Tip: Construiește + învață Limbaje: Python Cerințe preliminare: Faza 4 Lecția 06 (YOLO), Faza 4 Lecția 07 (U-Net) Timp: ~75 de minute
Obiective de învățare
- Urmăriți arhitectura Mask R-CNN de la un capăt la celălalt: backbone, FPN, RPN, RoIAlign, cap pentru casete, cap pentru măști
- Implementați RoIAlign de la zero și explicați de ce RoIPool nu mai este folosit în Mask R-CNN
- Folosiți modelul preantrenat torchvision
maskrcnn_resnet50_fpn_v2pentru măști de instanță de calitate de producție și interpretați corect formatul ieșirii sale - Ajustați fin Mask R-CNN pe un set de date personalizat mic, înlocuind capetele pentru casete și măști și păstrând backbone-ul înghețat
Problema
Segmentarea semantică oferă câte o mască pentru fiecare clasă. Segmentarea instanțelor oferă câte o mască pentru fiecare obiect, chiar și atunci când două obiecte aparțin aceleiași clase. Numărarea indivizilor, urmărirea între cadre și măsurarea obiectelor (caseta de încadrare a fiecărei cărămizi dintr-un zid, a fiecărei celule dintr-o imagine la microscop) necesită toate segmentarea instanțelor.
Mask R-CNN (He și colab., 2017) a rezolvat această problemă prin reformularea segmentării instanțelor ca detecție plus o mască. Proiectarea a fost atât de curată, încât în următorii cinci ani aproape fiecare lucrare despre segmentarea instanțelor a fost o variantă Mask R-CNN, iar implementarea torchvision oferă un punct de plecare solid pentru seturi de date mici și medii.
Notă tehnică a traducerii: Nu există o implementare universală implicită pentru producție. Alegerea frameworkului și a modelului depinde de cerințele de latență, memorie, licență, infrastructură și date; exemplele torchvision și Detectron2 sunt implementări de referință, nu o prescripție universală.
Problema inginerească dificilă este eșantionarea: cum extrageți o regiune de caracteristici de dimensiune fixă dintr-o casetă-propunere ale cărei colțuri nu se aliniază cu marginile pixelilor? O eroare aici costă zecimi de punct mAP pretutindeni. RoIAlign este răspunsul.
Conceptul
Arhitectura
Cinci componente de înțeles:
- Backbone — ResNet-50 sau ResNet-101 antrenat pe ImageNet. Produce o ierarhie de hărți de caracteristici cu stride-urile 4, 8, 16, 32.
- FPN (Feature Pyramid Network) — conexiuni de sus în jos plus laterale, care oferă fiecărui nivel C canale de caracteristici bogate semantic. Detecția interoghează nivelul FPN care corespunde dimensiunii obiectului.
- RPN (Region Proposal Network) — un cap convoluțional mic care, la fiecare poziție de ancoră, prezice „există un obiect aici?” și „cum rafinez caseta?”. Produce aproximativ 1 000 de propuneri pe imagine.
- RoIAlign — eșantionează un petic de caracteristici de dimensiune fixă (de exemplu, 7x7) din orice casetă de la orice nivel FPN. Eșantionare biliniară, fără cuantizare.
- Capete — un cap pentru casete cu două straturi, care rafinează caseta și alege o clasă, plus un cap convoluțional mic, care produce pentru fiecare propunere o mască binară
28x28.
De ce RoIAlign, nu RoIPool
Fast R-CNN original folosea RoIPool, care împarte o casetă-propunere într-o grilă, ia caracteristica maximă din fiecare celulă și rotunjește toate coordonatele la întregi. Acea rotunjire dezaliniază harta de caracteristici față de coordonatele pixelilor de intrare cu până la un pixel întreg din harta de caracteristici — puțin într-o imagine de 224x224, catastrofal atunci când harta de caracteristici are stride 32.
RoIPool:
caseta (34.7, 51.3, 98.2, 142.9)
rotunjire -> (34, 51, 98, 142)
împărțire în grilă -> rotunjește fiecare limită de celulă
dezalinierea se acumulează la fiecare pas
RoIAlign:
caseta (34.7, 51.3, 98.2, 142.9)
eșantionează la coordonate exacte în virgulă mobilă folosind interpolare biliniară
fără rotunjire nicăieri
În ablațiile raportate pentru Mask R-CNN, RoIAlign îmbunătățește AP față de RoIPool; mărimea câștigului depinde de backbone, stride, setul de date și protocolul de evaluare. Detectoarele în două etape care au nevoie de localizare precisă îl folosesc frecvent.
Notă tehnică a traducerii: Generalizarea originalului la toate detectoarele, inclusiv RT-DETR și Mask2Former, nu este corectă. RT-DETR este un detector Transformer end-to-end cu interogări de obiecte, iar Mask2Former folosește atenție mascată; niciuna dintre aceste arhitecturi nu folosește RoIAlign în rolul descris aici. RoIAlign rămâne însă componenta-cheie a Mask R-CNN și a multor detectoare în două etape.
RPN într-un paragraf
La fiecare poziție a unei hărți de caracteristici, plasați K casete-ancoră de dimensiuni și forme diferite. Preziceți pentru fiecare ancoră un scor de obiectualitate și un decalaj de regresie care transformă ancora într-o casetă potrivită mai bine. Într-o configurație uzuală torchvision, păstrați aproximativ 1 000 de casete cu scorul cel mai mare, aplicați NMS la IoU 0,7 și transmiteți supraviețuitoarele capetelor. RPN este antrenată cu pierderea proprie de clasificare binară a ancorelor și regresie a casetelor.
Notă tehnică a traducerii: Numărul de propuneri, pragul NMS și regulile de potrivire a ancorelor sunt hiperparametri ai implementării, nu constante ale RPN. Deși ambele includ clasificare și regresie de casetă, pierderea RPN bazată pe ancore nu este aceeași formulare ca pierderea YOLO bazată pe grilă.
Capul pentru măști
Pentru fiecare propunere (după RoIAlign), capul pentru măști este o FCN mică: patru convoluții 3x3, o deconvoluție 2x și o convoluție finală 1x1 care produce num_classes canale de ieșire la rezoluția 28x28. Se păstrează doar canalul corespunzător clasei prezise; celelalte sunt ignorate. Aceasta decuplează predicția măștii de clasificare.
Supraeșantionați masca 28x28 la dimensiunea originală în pixeli a propunerii pentru a obține masca binară finală.
Pierderi
Implementarea completă Mask R-CNN are cinci termeni de pierdere adunați:
L = L_rpn_cls + L_rpn_box + L_box_cls + L_box_reg + L_mask
L_rpn_cls,L_rpn_box— obiectualitatea plus regresia casetei pentru propunerile RPN.L_box_cls— entropie încrucișată pentru (C+1) clase (inclusiv fundalul) în clasificatorul capului.L_box_reg— smooth L1 pentru rafinarea casetei de către cap.L_mask— entropie încrucișată binară per pixel pentru ieșirea măștii 28x28.
În implementarea standard, termenii sunt însumați; ponderile și opțiunile specifice depind de implementare.
Notă tehnică a traducerii: Originalul spune „patru pierderi”, dar formula listează cinci termeni. Articolul Mask R-CNN definește trei termeni pe fiecare RoI (
L_cls,L_box,L_mask), iar detectorul complet adaugă cele două pierderi RPN. API-ul public torchvision nu expune un set general de ponderi ale pierderilor drept argumente ale constructorului Mask R-CNN.
Formatul ieșirii
torchvision.models.detection.maskrcnn_resnet50_fpn_v2 returnează o listă de dicționare, câte unul pentru fiecare imagine:
{
"boxes": (N, 4) în coordonate de pixeli (x1, y1, x2, y2),
"labels": (N,) ID-uri de clasă; 0 = fundal, deci indicii obiectelor încep de la 1,
"scores": (N,) scoruri de încredere,
"masks": (N, 1, H, W) măști în virgulă mobilă în [0, 1] — pragul 0.5 le face binare,
}
Masca este deja la rezoluția imaginii complete. Ieșirea de 28x28 a capului a fost supraeșantionată intern.
Notă tehnică a traducerii: Convenția cu id-ul de fundal
0, cele 91 de clase și ponderile COCO se aplică variantei preantrenate curente; în propriul set de date,num_classesși semantica etichetelor trebuie să corespundă exact anotărilor și versiunii de bibliotecă folosite.
Construiți
Pasul 1: RoIAlign de la zero
Aceasta este singura componentă din Mask R-CNN mai ușor de înțeles ca cod decât ca proză.
import torch
import torch.nn.functional as F
def roi_align_single(feature, box, output_size=7, spatial_scale=1 / 16.0):
"""
feature: (C, H, W) single-image feature map
box: (x1, y1, x2, y2) in original image pixel coordinates
output_size: side of the output grid (7 for box head, 14 for mask head)
spatial_scale: reciprocal of the feature map stride
"""
C, H, W = feature.shape
x1, y1, x2, y2 = [c * spatial_scale - 0.5 for c in box]
bin_w = (x2 - x1) / output_size
bin_h = (y2 - y1) / output_size
grid_y = torch.linspace(y1 + bin_h / 2, y2 - bin_h / 2, output_size)
grid_x = torch.linspace(x1 + bin_w / 2, x2 - bin_w / 2, output_size)
yy, xx = torch.meshgrid(grid_y, grid_x, indexing="ij")
gx = 2 * (xx + 0.5) / W - 1
gy = 2 * (yy + 0.5) / H - 1
grid = torch.stack([gx, gy], dim=-1).unsqueeze(0)
sampled = F.grid_sample(feature.unsqueeze(0), grid, mode="bilinear",
align_corners=False)
return sampled.squeeze(0)
Fiecare număr este într-o poziție eșantionată biliniar. Fără rotunjire, fără cuantizare, fără gradienți eliminați.
Pasul 2: Comparați cu RoIAlign din torchvision
from torchvision.ops import roi_align
feature = torch.randn(1, 16, 50, 50)
boxes = torch.tensor(0, 10, 20, 100, 90, dtype=torch.float32) # (batch_idx, x1, y1, x2, y2)
ours = roi_align_single(feature[0], boxes[0, 1:].tolist(), output_size=7, spatial_scale=1/4)
theirs = roi_align(feature, boxes, output_size=(7, 7), spatial_scale=1/4, sampling_ratio=1, aligned=True)[0]
print(f"shape ours: {tuple(ours.shape)}")
print(f"shape theirs: {tuple(theirs.shape)}")
print(f"max|diff|: {(ours - theirs).abs().max().item():.3e}")
Cu sampling_ratio=1 și aligned=True, cele două rezultate coincid până la 1e-5.
Notă tehnică a traducerii: Această funcție este o demonstrație cu o singură imagine și un singur RoI. Nu implementează eșantionarea adaptivă, loturile și mai multe RoI-uri, toate cazurile de margine sau propagarea dispozitivului și dtype-ului; fără
device=feature.device,torch.linspacefolosește dispozitivul implicit, care poate să nu coincidă cu cel al caracteristicilor. Pentru antrenare sau inferență reală folosițitorchvision.ops.roi_align.
Pasul 3: Încărcați un Mask R-CNN preantrenat
import torch
from torchvision.models.detection import maskrcnn_resnet50_fpn_v2, MaskRCNN_ResNet50_FPN_V2_Weights
model = maskrcnn_resnet50_fpn_v2(weights=MaskRCNN_ResNet50_FPN_V2_Weights.DEFAULT)
model.eval()
print(f"params: {sum(p.numel() for p in model.parameters()):,}")
print(f"classes (including background): {len(model.roi_heads.box_predictor.cls_score.out_features * [0])}")
Aproximativ 46M de parametri și 91 de clase COCO pentru varianta preantrenată consultată. Clasa internă cu id-ul 0 este fundalul; etichetele obiectelor prezise în această configurație încep de la 1.
Pasul 4: Rulați inferența
with torch.no_grad():
x = torch.randn(3, 400, 600)
predictions = model([x])
p = predictions[0]
print(f"boxes: {tuple(p['boxes'].shape)}")
print(f"labels: {tuple(p['labels'].shape)}")
print(f"scores: {tuple(p['scores'].shape)}")
print(f"masks: {tuple(p['masks'].shape)}")
Tensorul măștilor are forma (N, 1, H, W). Aplicați pragul 0,5 pentru a obține o mască binară pentru fiecare obiect:
binary_masks = (p['masks'] > 0.5).squeeze(1) # (N, H, W) boolean
Notă tehnică a traducerii: Exemplul cu
torch.randnverifică numai formele și nu este preprocesarea corectă pentru inferență. Modelele torchvision de detecție așteaptă imagini în intervalul[0, 1]; aplicați transformarea asociată ponderilor folosite.
Pasul 5: Înlocuiți capetele pentru un număr de clase personalizat
Rețeta obișnuită de ajustare fină: reutilizați backbone-ul, FPN și RPN; înlocuiți cele două capete de clasificare.
from torchvision.models.detection.faster_rcnn import FastRCNNPredictor
from torchvision.models.detection.mask_rcnn import MaskRCNNPredictor
def build_custom_maskrcnn(num_classes):
model = maskrcnn_resnet50_fpn_v2(weights=MaskRCNN_ResNet50_FPN_V2_Weights.DEFAULT)
in_features = model.roi_heads.box_predictor.cls_score.in_features
model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes)
in_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels
hidden_layer = 256
model.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask, hidden_layer, num_classes)
return model
custom = build_custom_maskrcnn(num_classes=5)
print(f"custom cls_score.out_features: {custom.roi_heads.box_predictor.cls_score.out_features}")
num_classes trebuie să includă clasa de fundal; așadar, un set de date cu 4 clase de obiecte folosește num_classes=5.
Pasul 6: Înghețați ceea ce nu trebuie antrenat
Pe seturi de date mici, înghețați backbone-ul și FPN. Doar obiectualitatea plus regresia RPN și cele două capete învață.
def freeze_backbone_and_fpn(model):
# torchvision Mask R-CNN packs the FPN inside `model.backbone` (as
# `model.backbone.fpn`), so iterating `model.backbone.parameters()` covers
# both the ResNet feature layers and the FPN lateral/output convs.
for p in model.backbone.parameters():
p.requires_grad = False
return model
custom = freeze_backbone_and_fpn(custom)
trainable = sum(p.numel() for p in custom.parameters() if p.requires_grad)
print(f"trainable after freeze: {trainable:,}")
Pe seturi de date cu 500 de imagini, înghețarea poate reduce numărul parametrilor ajustați și riscul de supraînvățare.
Notă tehnică a traducerii: Înghețarea backbone-ului și FPN nu garantează nici convergența, nici evitarea supraînvățării. Rezultatul depinde de dimensiunea și diversitatea datelor, de etichete, augmentare, ratele de învățare, numărul de epoci și evaluarea pe date separate.
Folosiți
Bucla completă de antrenare pentru Mask R-CNN în torchvision are 40 de linii și nu își schimbă semnificativ sensul între sarcini — înlocuiți seturile de date și porniți.
def train_step(model, images, targets, optimizer):
model.train()
loss_dict = model(images, targets)
losses = sum(loss for loss in loss_dict.values())
optimizer.zero_grad()
losses.backward()
optimizer.step()
return {k: v.item() for k, v in loss_dict.items()}
Lista targets trebuie să conțină, pentru fiecare imagine, dicționare cu boxes de tip FloatTensor[N, 4], labels de tip Int64Tensor[N] și masks binare de tip UInt8Tensor[N, H, W]. Modelul returnează la antrenare cele cinci componente ale pierderii — clasificare și regresie pentru RPN, clasificare și regresie pentru RoI, plus pierderea măștii — și o listă de predicții la evaluare, în funcție de model.training.
Notă tehnică a traducerii: Numărul „patru” din original este inexact pentru Mask R-CNN în torchvision: dicționarul de antrenare conține pierderile de clasificare și regresie pentru RPN, pierderile de clasificare și regresie pentru RoI și pierderea măștii.
Evaluatorul pycocotools produce mAP@IoU=0,5:0,95 atât pentru casete, cât și pentru măști; aveți nevoie de ambele numere pentru a afla dacă blocajul este capul pentru casete sau capul pentru măști.
Livrați
Această lecție produce:
outputs/prompt-instance-vs-semantic-router.md— un prompt care pune trei întrebări și alege segmentarea instanțelor, semantică sau panoptică, plus modelul exact cu care să începeți.outputs/skill-mask-rcnn-head-swapper.md— o abilitate care generează cele 10 linii de cod pentru înlocuirea capetelor în orice model torchvision de detecție, dat fiind noulnum_classes.
Exerciții
- (Ușor) Verificați RoIAlign față de
torchvision.ops.roi_alignpe 100 de casete aleatoare. Raportați diferența absolută maximă. Rulați și RoIPool (comportamentul de dinainte de 2017) și măsurați abaterea pentru casetele apropiate de margine.
Notă tehnică a traducerii: Nu există o abatere fixă de 1–2 pixeli din harta de caracteristici pentru RoIPool: ea depinde de coordonatele RoI, stride, dimensiunea grilei, convenția de rotunjire și implementare.
- (Mediu) Ajustați fin
maskrcnn_resnet50_fpn_v2pe un set de date personalizat cu 50 de imagini (oricare două clase: baloane, pești, gropi în asfalt, logouri). Înghețați backbone-ul, antrenați 20 de epoci și raportați mask AP@0,5. - (Dificil) Înlocuiți capul pentru măști al Mask R-CNN cu unul care prezice la 56x56 în loc de 28x28. Măsurați mAP@IoU=0,75 înainte și după. Explicați de ce câștigul (sau absența lui) corespunde compromisului preconizat dintre precizia conturului și memorie.
Termeni-cheie
| Termen | Ce spun oamenii | Ce înseamnă de fapt |
|---|---|---|
| Mask R-CNN | „Detecție plus măști” | Faster R-CNN plus un cap FCN mic, care prezice pentru fiecare propunere și clasă o mască de 28x28 |
| FPN | „Piramidă de caracteristici” | Conexiuni de sus în jos plus laterale, care oferă fiecărui nivel de stride C canale de caracteristici bogate semantic |
| RPN | „Propunător de regiuni” | Un cap convoluțional mic, care produce aproximativ 1 000 de propuneri obiect/fără obiect pe imagine |
| RoIAlign | „Decupare fără rotunjire” | Eșantionează biliniar o grilă de caracteristici de dimensiune fixă din orice casetă cu coordonate în virgulă mobilă |
| RoIPool | „Decupare de dinainte de 2017” | Are același scop ca RoIAlign, dar rotunjește coordonatele casetei; este înlocuit în Mask R-CNN |
| Mask AP | „mAP pentru instanțe” | Precizia medie calculată cu IoU al măștilor în loc de IoU al casetelor; metrica COCO pentru segmentarea instanțelor |
| Cap binar pentru măști | „Mască pe clasă” | Prezice pentru fiecare propunere câte o mască binară pentru fiecare clasă; se păstrează doar canalul clasei prezise |
| Clasă de fundal | „Clasa 0” | Clasa globală „fără obiect”; indicii claselor reale încep de la 1 |
Lecturi suplimentare
- Mask R-CNN (He și colab., 2017) — lucrarea; secțiunea 3 despre RoIAlign este lectura critică
- FPN: Feature Pyramid Networks (Lin și colab., 2017) — lucrarea despre FPN; este folosită de multe detectoare moderne, dar nu de toate
- tutorialul torchvision pentru Mask R-CNN — referința pentru bucla de ajustare fină
- grădina zoologică de modele Detectron2 — implementări de referință cu ponderi antrenate pentru numeroase variante de detecție și segmentare; consultați catalogul versiunii pentru disponibilitatea exactă
Sursă: Originalul în limba engleză
Navigare: ← Lecția 04.07 — Segmentarea semantică — U-Net · Faza 4 — Viziune computerizată · Lecția 04.09 — Generarea imaginilor — GAN-uri → · Catalog complet