Fine Tuning Detr Custom Dataset
커스텀 데이터셋으로 객체 탐지 모델 파인튜닝하기 🖼, Spaces에 배포하기, 그리고 Gradio API 연동하기
작성자: Sergio Paniego 번역: 한지욱
이 노트북에서는 object detection 모델, 그 중에서도 DETR을 커스텀 데이터셋으로 파인튜닝해볼 거예요. 이 작업은 Hugging Face 생태계를 활용해서 진행할 거예요.
우리는 사전 학습된 DETR 모델을 시작점으로 삼고, Fashionpedia라는 주석이 달린 패션 이미지 데이터셋으로 모델을 파인튜닝할 거예요. 이렇게 하면 모델이 패션 분야에서 더 잘 객체를 인식하고 감지할 수 있도록 조정하게 돼요.
모델을 성공적으로 파인튜닝한 후에는 Hugging Face의 Gradio Space에 배포할 거예요. 그리고 나서 Gradio API를 통해 배포된 모델과 상호작용하는 방법도 알아볼 거예요. 이를 통해 실제 응용에도 활용할 수 있는 다양한 가능성을 열 수 있어요.

1. 라이브러리 설치
먼저 객체 탐지 모델을 파인튜닝하기 위해 필요한 라이브러리들을 설치해볼게요.
2. 데이터셋 로드 📁
📁 우리가 사용할 데이터셋은 Fashionpedia예요. 이 데이터셋은 논문 Fashionpedia: Ontology, Segmentation, and an Attribute Localization Dataset에서 소개된 거고, 저자들은 이렇게 설명하고 있어요:
Fashionpedia is a dataset which consists of two parts: (1) an ontology built by fashion experts containing 27 main apparel categories, 19 apparel parts, 294 fine-grained attributes and their relationships; (2) a dataset with 48k everyday and celebrity event fashion images annotated with segmentation masks and their associated per-mask fine-grained attributes, built upon the Fashionpedia ontology.
Fashionpedia는 두 가지로 구성된 데이터셋입니다: (1) 패션 전문가들이 만든 온톨로지로, 27개의 주요 의류 카테고리, 19개의 의류 부위, 294개의 세분화된 속성과 이들 간의 관계를 담고 있습니다; (2) Fashionpedia 온톨로지를 기반으로, 분할 마스크와 마스크별 세부 속성 주석이 달린 일상 및 연예인 이벤트 패션 이미지 48,000여 장으로 구성된 데이터셋입니다.
이 데이터셋에는 다음과 같은 내용이 포함돼 있어요:
- 이미지 46,781장 🖼
- 바운딩 박스 342,182개 📦
Hugging Face에서 다운로드할 수 있어요: Fashionpedia Dataset
DatasetDict({
, train: Dataset({
, features: ['image_id', 'image', 'width', 'height', 'objects'],
, num_rows: 45623
, })
, val: Dataset({
, features: ['image_id', 'image', 'width', 'height', 'objects'],
, num_rows: 1158
, })
,}) 예제 중 하나의 내부 구조를 한번 살펴볼게요.
{'image_id': 23,
, 'image': <PIL.JpegImagePlugin.JpegImageFile image mode=RGB size=682x1024>,
, 'width': 682,
, 'height': 1024,
, 'objects': {'bbox_id': [150311, 150312, 150313, 150314],
, 'category': [23, 23, 33, 10],
, 'bbox': [[445.0, 910.0, 505.0, 983.0],
, [239.0, 940.0, 284.0, 994.0],
, [298.0, 282.0, 386.0, 352.0],
, [210.0, 282.0, 448.0, 665.0]],
, 'area': [1422, 843, 373, 56375]}} 3. 모델 학습이랑 테스트를 위해 데이터셋 나누기 ➗
이 데이터셋은 train이랑 test 두 가지로 나뉘어 있어요. 모델을 파인튜닝할 때는 train 데이터를, 모델 성능을 확인할 때는 test 데이터를 사용할 거예요.
선택 사항
다음 주석 처리된 셀에서는 전체 데이터셋에서 훈련용과 테스트용으로 각각 1%만 랜덤으로 샘플링해요. 데이터셋이 워낙 크기 때문에 학습 속도를 빠르게 하기 위한 방법이에요.
최상의 결과를 원한다면 이 두 셀은 건너뛰고 전체 데이터셋을 사용하는 걸 추천해요. 하지만 필요하다면 주석을 해제해서 사용할 수도 있어요.
4. 데이터셋에서 예제 하나를 골라, 그 안에 있는 객체들 시각화해보기👀
이제 데이터셋을 불러왔으니, 주석이 달린 객체들과 함께 예제 하나를 시각화해볼게요.
id2label과 label2id 생성하기
이 변수들은 객체 ID와 해당 레이블 간의 매핑 정보를 담고 있어요. id2label은 ID를 레이블로 바꿔주고, label2id는 레이블을 ID로 바꿔주는 역할을 해요.
이미지 하나 그려보기! 🎨
이제 데이터셋에 있는 이미지 하나를 시각화해서 어떤 모습인지 직접 확인해볼게요.
{'bbox_id': [158977, 158978, 158979, 158980, 158981, 158982, 158983], 'category': [1, 23, 23, 6, 31, 31, 33], 'bbox': [[210.0, 225.0, 536.0, 784.0], [290.0, 897.0, 350.0, 1015.0], [464.0, 950.0, 534.0, 1021.0], [313.0, 407.0, 524.0, 954.0], [268.0, 229.0, 333.0, 563.0], [489.0, 247.0, 528.0, 591.0], [387.0, 225.0, 450.0, 253.0]], 'area': [69960, 2449, 1788, 75418, 15149, 5998, 479]}
이미지 몇 개 더 살펴보기 📸
이제 데이터셋에 있는 이미지를 몇 개 더 보면서 전체적인 데이터 분위기를 파악해볼게요.
{'bbox_id': [150311, 150312, 150313, 150314], 'category': [23, 23, 33, 10], 'bbox': [[445.0, 910.0, 505.0, 983.0], [239.0, 940.0, 284.0, 994.0], [298.0, 282.0, 386.0, 352.0], [210.0, 282.0, 448.0, 665.0]], 'area': [1422, 843, 373, 56375]}
{'bbox_id': [158953, 158954, 158955, 158956, 158957, 158958, 158959, 158960, 158961, 158962], 'category': [2, 33, 31, 31, 13, 7, 22, 22, 23, 23], 'bbox': [[182.0, 220.0, 472.0, 647.0], [294.0, 221.0, 407.0, 257.0], [405.0, 297.0, 472.0, 647.0], [182.0, 264.0, 266.0, 621.0], [284.0, 135.0, 372.0, 169.0], [238.0, 537.0, 414.0, 606.0], [351.0, 732.0, 417.0, 922.0], [202.0, 749.0, 270.0, 930.0], [200.0, 921.0, 256.0, 979.0], [373.0, 903.0, 455.0, 966.0]], 'area': [87267, 1220, 16895, 18541, 1468, 9360, 8629, 8270, 2717, 3121]}
{'bbox_id': [169196, 169197, 169198, 169199, 169200, 169201, 169202, 169203, 169204, 169205, 169206, 169207, 169208, 169209, 169210], 'category': [13, 29, 28, 32, 32, 31, 31, 0, 31, 31, 18, 4, 6, 23, 23], 'bbox': [[441.0, 132.0, 499.0, 150.0], [412.0, 164.0, 494.0, 295.0], [427.0, 164.0, 476.0, 207.0], [406.0, 326.0, 448.0, 335.0], [484.0, 327.0, 508.0, 334.0], [366.0, 323.0, 395.0, 372.0], [496.0, 271.0, 523.0, 302.0], [366.0, 164.0, 523.0, 372.0], [360.0, 186.0, 406.0, 332.0], [502.0, 201.0, 534.0, 321.0], [496.0, 259.0, 515.0, 278.0], [360.0, 164.0, 534.0, 411.0], [403.0, 384.0, 510.0, 638.0], [393.0, 584.0, 430.0, 663.0], [449.0, 638.0, 518.0, 681.0]], 'area': [587, 2922, 931, 262, 111, 1171, 540, 3981, 4457, 1724, 188, 26621, 16954, 2167, 1773]}
{'bbox_id': [167967, 167968, 167969, 167970, 167971, 167972, 167973, 167974, 167975, 167976, 167977, 167978, 167979, 167980, 167981, 167982, 167983, 167984, 167985, 167986, 167987, 167988, 167989, 167990, 167991, 167992, 167993, 167994, 167995, 167996, 167997, 167998, 167999, 168000, 168001, 168002, 168003, 168004, 168005, 168006, 168007, 168008, 168009, 168010, 168011, 168012, 168013, 168014, 168015, 168016, 168017, 168018, 168019, 168020, 168021, 168022, 168023, 168024, 168025, 168026, 168027, 168028, 168029, 168030, 168031, 168032, 168033, 168034, 168035, 168036, 168037, 168038, 168039, 168040], 'category': [6, 23, 23, 31, 31, 4, 1, 35, 32, 35, 35, 35, 35, 28, 35, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 42, 33], 'bbox': [[300.0, 421.0, 460.0, 846.0], [383.0, 841.0, 432.0, 899.0], [304.0, 740.0, 347.0, 831.0], [246.0, 222.0, 295.0, 505.0], [456.0, 229.0, 492.0, 517.0], [246.0, 169.0, 492.0, 517.0], [355.0, 213.0, 450.0, 433.0], [289.0, 353.0, 303.0, 427.0], [442.0, 288.0, 460.0, 340.0], [451.0, 290.0, 458.0, 304.0], [407.0, 238.0, 473.0, 486.0], [487.0, 501.0, 491.0, 517.0], [246.0, 455.0, 252.0, 505.0], [340.0, 169.0, 442.0, 238.0], [348.0, 230.0, 372.0, 476.0], [411.0, 179.0, 414.0, 182.0], [414.0, 183.0, 418.0, 186.0], [418.0, 187.0, 421.0, 190.0], [421.0, 192.0, 425.0, 195.0], [424.0, 196.0, 428.0, 199.0], [426.0, 200.0, 430.0, 204.0], [429.0, 204.0, 433.0, 208.0], [431.0, 209.0, 435.0, 213.0], [433.0, 214.0, 437.0, 218.0], [434.0, 218.0, 438.0, 222.0], [436.0, 223.0, 440.0, 226.0], [437.0, 227.0, 441.0, 231.0], [438.0, 232.0, 442.0, 235.0], [433.0, 232.0, 437.0, 236.0], [429.0, 233.0, 432.0, 237.0], [423.0, 233.0, 426.0, 237.0], [417.0, 233.0, 421.0, 237.0], [353.0, 172.0, 355.0, 174.0], [353.0, 175.0, 354.0, 177.0], [351.0, 178.0, 353.0, 181.0], [350.0, 182.0, 351.0, 184.0], [347.0, 187.0, 350.0, 189.0], [346.0, 190.0, 349.0, 193.0], [345.0, 194.0, 348.0, 197.0], [344.0, 199.0, 347.0, 202.0], [342.0, 204.0, 346.0, 207.0], [342.0, 208.0, 345.0, 211.0], [342.0, 212.0, 344.0, 215.0], [342.0, 217.0, 345.0, 220.0], [344.0, 221.0, 346.0, 224.0], [348.0, 222.0, 350.0, 225.0], [353.0, 223.0, 356.0, 226.0], [359.0, 223.0, 361.0, 226.0], [364.0, 223.0, 366.0, 226.0], [247.0, 448.0, 253.0, 454.0], [251.0, 454.0, 254.0, 456.0], [252.0, 460.0, 255.0, 463.0], [252.0, 466.0, 255.0, 469.0], [253.0, 471.0, 255.0, 475.0], [253.0, 478.0, 255.0, 481.0], [253.0, 483.0, 256.0, 486.0], [254.0, 489.0, 256.0, 492.0], [254.0, 495.0, 256.0, 497.0], [247.0, 457.0, 249.0, 460.0], [247.0, 463.0, 249.0, 466.0], [248.0, 469.0, 249.0, 471.0], [248.0, 476.0, 250.0, 478.0], [248.0, 481.0, 250.0, 483.0], [249.0, 486.0, 250.0, 488.0], [487.0, 459.0, 490.0, 461.0], [487.0, 465.0, 490.0, 467.0], [487.0, 471.0, 490.0, 472.0], [487.0, 476.0, 489.0, 478.0], [486.0, 482.0, 489.0, 484.0], [486.0, 488.0, 489.0, 490.0], [486.0, 494.0, 488.0, 496.0], [486.0, 500.0, 488.0, 501.0], [485.0, 505.0, 487.0, 507.0], [365.0, 213.0, 409.0, 226.0]], 'area': [44062, 2140, 2633, 9206, 5905, 44791, 12948, 211, 335, 43, 691, 62, 104, 2169, 439, 9, 10, 9, 8, 9, 14, 10, 13, 13, 11, 11, 10, 10, 12, 10, 10, 14, 4, 2, 4, 2, 5, 6, 7, 7, 8, 7, 6, 7, 5, 5, 7, 6, 5, 12, 5, 7, 8, 6, 6, 6, 4, 4, 6, 5, 2, 4, 4, 2, 6, 6, 3, 4, 6, 6, 4, 2, 4, 94]}
{'bbox_id': [168041, 168042, 168043, 168044, 168045, 168046, 168047], 'category': [10, 32, 35, 31, 4, 29, 33], 'bbox': [[238.0, 309.0, 471.0, 1022.0], [234.0, 572.0, 331.0, 602.0], [235.0, 580.0, 324.0, 599.0], [119.0, 318.0, 343.0, 856.0], [111.0, 262.0, 518.0, 1022.0], [166.0, 262.0, 393.0, 492.0], [238.0, 309.0, 278.0, 324.0]], 'area': [12132, 1548, 755, 43926, 178328, 9316, 136]}
{'bbox_id': [160050, 160051, 160052, 160053, 160054, 160055], 'category': [10, 31, 31, 23, 23, 33], 'bbox': [[290.0, 364.0, 429.0, 665.0], [304.0, 369.0, 397.0, 508.0], [290.0, 468.0, 310.0, 522.0], [213.0, 842.0, 294.0, 905.0], [446.0, 840.0, 536.0, 896.0], [311.0, 364.0, 354.0, 379.0]], 'area': [26873, 5301, 747, 1438, 1677, 71]}
{'bbox_id': [160056, 160057, 160058, 160059, 160060, 160061, 160062, 160063, 160064, 160065, 160066], 'category': [10, 36, 42, 42, 42, 42, 42, 42, 42, 23, 33], 'bbox': [[127.0, 198.0, 451.0, 949.0], [277.0, 336.0, 319.0, 402.0], [340.0, 343.0, 344.0, 347.0], [321.0, 338.0, 327.0, 343.0], [336.0, 361.0, 342.0, 365.0], [329.0, 321.0, 333.0, 326.0], [313.0, 294.0, 319.0, 300.0], [330.0, 299.0, 334.0, 304.0], [295.0, 330.0, 300.0, 334.0], [332.0, 926.0, 376.0, 946.0], [284.0, 198.0, 412.0, 270.0]], 'area': [137575, 1915, 14, 24, 18, 15, 25, 16, 16, 740, 586]}
{'bbox_id': [158963, 158964, 158965, 158966, 158967, 158968, 158969, 158970, 158971], 'category': [1, 31, 31, 7, 22, 22, 23, 23, 33], 'bbox': [[262.0, 449.0, 435.0, 686.0], [399.0, 471.0, 435.0, 686.0], [262.0, 451.0, 294.0, 662.0], [276.0, 603.0, 423.0, 726.0], [291.0, 759.0, 343.0, 934.0], [341.0, 749.0, 401.0, 947.0], [302.0, 919.0, 337.0, 994.0], [323.0, 925.0, 374.0, 1005.0], [343.0, 456.0, 366.0, 467.0]], 'area': [22330, 4422, 4846, 14000, 6190, 6997, 1547, 2107, 49]}
{'bbox_id': [158972, 158973, 158974, 158975, 158976], 'category': [23, 23, 28, 10, 5], 'bbox': [[412.0, 588.0, 451.0, 631.0], [333.0, 585.0, 357.0, 627.0], [361.0, 243.0, 396.0, 257.0], [303.0, 243.0, 447.0, 517.0], [330.0, 259.0, 425.0, 324.0]], 'area': [949, 737, 133, 17839, 2916]}
5. 유효하지 않은 바운딩 박스 걸러내기 ❌
데이터셋 전처리의 첫 번째 단계로, 잘못된 바운딩 박스를 걸러낼 거예요. 데이터셋을 살펴보니까 몇몇 바운딩 박스가 제대로 된 구조를 갖고 있지 않더라고요. 그래서 이런 유효하지 않은 항목들은 제외할 예정이에요.
Map: 0%| | 0/45623 [00:00<?, ? examples/s]
Image with invalid bbox: 8396 Invalid bbox detected and discarded: [0.0, 0.0, 0.0, 0.0] - bbox_id: 139952 - category: 42 Image with invalid bbox: 19725 Invalid bbox detected and discarded: [0.0, 0.0, 0.0, 0.0] - bbox_id: 23298 - category: 42 Image with invalid bbox: 19725 Invalid bbox detected and discarded: [0.0, 0.0, 0.0, 0.0] - bbox_id: 23299 - category: 42 Image with invalid bbox: 21696 Invalid bbox detected and discarded: [0.0, 0.0, 0.0, 0.0] - bbox_id: 277148 - category: 42 Image with invalid bbox: 23055 Invalid bbox detected and discarded: [0.0, 0.0, 0.0, 0.0] - bbox_id: 287029 - category: 33 Image with invalid bbox: 23671 Invalid bbox detected and discarded: [0.0, 0.0, 0.0, 0.0] - bbox_id: 290142 - category: 42 Image with invalid bbox: 26549 Invalid bbox detected and discarded: [0.0, 0.0, 0.0, 0.0] - bbox_id: 311943 - category: 37 Image with invalid bbox: 26834 Invalid bbox detected and discarded: [0.0, 0.0, 0.0, 0.0] - bbox_id: 309141 - category: 37 Image with invalid bbox: 31748 Invalid bbox detected and discarded: [0.0, 0.0, 0.0, 0.0] - bbox_id: 262063 - category: 42 Image with invalid bbox: 34253 Invalid bbox detected and discarded: [0.0, 0.0, 0.0, 0.0] - bbox_id: 315750 - category: 19
Map: 0%| | 0/1158 [00:00<?, ? examples/s]
Dataset({
features: ['image_id', 'image', 'width', 'height', 'objects'],
num_rows: 45623
})
Dataset({
features: ['image_id', 'image', 'width', 'height', 'objects'],
num_rows: 1158
})
6. 클래스별 등장 횟수 시각화하기 👀
이번에는 각 클래스가 얼마나 자주 등장하는지 그래프로 확인해볼게요. 이걸 통해 클래스 분포를 파악하고, 혹시 치우친 부분은 없는지도 살펴볼 수 있어요.
<ipython-input-66-aa111d1e000d>:14: UserWarning: FixedFormatter should only be used together with FixedLocator ax.set_xticklabels(categories, rotation=90, ha='right')
데이터셋을 보면 "신발"이나 "소매" 같은 몇몇 클래스가 유난히 많이 등장하는 걸 확인할 수 있어요. 이건 특정 클래스들이 더 자주 등장하면서 데이터셋이 불균형할 수 있다는 걸 의미해요. 이런 불균형을 파악하는 건, 모델 학습 과정에서 생길 수 있는 편향을 줄이기 위해 정말 중요해요.
7. 데이터셋에 데이터 증강 추가하기
데이터 증강(data augmentation)🪄은 객체 탐지 작업에서 성능을 높이는 데 아주 중요한 역할을 해요. 이번 섹션에서는 Albumentations 라이브러리를 활용해서 효과적으로 데이터셋을 증강해볼 거예요.
Albumentations는 객체 탐지를 위해 특화된 다양한 강력한 증강 기법들을 제공해요. 여러 가지 변형을 적용하면서도 바운딩 박스가 정확하게 조정되도록 해주기 때문에, 데이터셋을 더 다양하게 만들고 모델의 견고함과 일반화 능력을 높이는 데 큰 도움이 돼요.
8. 모델 체크포인트에서 이미지 프로세서 초기화하기 🎆
이제 사전 학습된 모델 체크포인트를 이용해서 이미지 프로세서를 만들어볼 거예요. 여기서는 facebook/detr-resnet-50-dc5 모델을 사용할 거예요.
preprocessor_config.json: 0%| | 0.00/274 [00:00<?, ?B/s]
The `max_size` parameter is deprecated and will be removed in v4.26. Please specify in `size['longest_edge'] instead`.
데이터셋을 처리하기 위한 메서드 추가하기
이제 데이터셋을 처리할 수 있는 메서드들을 추가해볼 거예요. 이 메서드들은 이미지랑 주석(annotation)을 모델에 맞게 변환해주는 역할을 해요.
9. 증강된 예제 시각화하기 🎆
이제 곧 모델 학습 단계에 들어가요! 그 전에, 증강이 잘 적용됐는지 확인하기 위해 증강된 샘플 몇 개를 시각화해볼게요. 이 과정을 통해 학습에 적합하고 효과적인 증강이 이루어졌는지 한 번 더 점검할 수 있어요.
10. 체크포인트에서 모델 초기화하기
이미지 프로세서와 같은 체크포인트를 사용해서 모델도 초기화할 거예요. 사전 학습된 모델을 불러온 다음, 우리 데이터셋에 맞게 파인튜닝할 예정이에요.
10. Hugging Face Hub에 연결해서 파인튜닝한 모델 업로드하기 🔌
이제 Hugging Face Hub에 연결해서 우리가 파인튜닝한 모델을 업로드해볼 거예요. 이렇게 하면 모델을 다른 사람들과 공유하거나, 배포하거나, 추가로 평가하는 데 활용할 수 있어요.
11. 학습 파라미터 설정하고, W&B 연결한 다음, 학습 시작하기!
이번에는 학습에 사용할 파라미터들을 설정하고 Weights & Biases (W&B)랑 연결해볼 거예요. W&B는 실험 추적, 지표 시각화, 학습 과정 관리 등에 도움을 주는 도구예요. 이제 본격적으로 학습을 시작해볼게요!
학습 과정을 추적하기 위해 W&B에 연결하기
모델을 학습시켜보기! 🚀
드디어 모델 학습을 시작할 시간이에요. 학습을 돌리면서, 우리가 파인튜닝한 모델이 데이터를 어떻게 학습하는지 지켜봐요!
먼저, 평가할 때 사용할 지표들을 계산해줄 compute_metrics 메서드를 정의할 거예요.
12. 테스트 이미지로 모델 성능 확인해보기 📝
모델 학습이 끝났으니, 이제 테스트 이미지를 가지고 성능을 확인해볼 수 있어요. 우리가 훈련한 모델은 Hugging Face 모델로 등록돼 있어서, 예측하는 과정도 아주 간단해요. 아래 셀에서는 새로운 이미지에 대해 추론을 수행하고, 모델이 얼마나 잘 작동하는지 확인해볼 거예요.
결과 확인
테스트 이미지에 대한 모델의 예측 결과를 시각화해서 보여줄 거예요. 이걸 통해 모델이 얼마나 잘 작동하는지 확인할 수 있고, 잘한 점과 개선할 부분도 파악할 수 있어요.