YOLOv8 객체 감지 모델 학습 및 ONNX 추론 변환 실전 가이드

모델 학습 설정 및 실행

Ultralytics 프레임워크를 통해 YOLOv8 을 학습시키는 가장 효율적인 방법은 `YOLO` 클래스를 초기화하고 `train` 메소드를 호출하는 것입니다. 프로젝트 디렉토리 구조에 따라 데이터 경로와 하이퍼파라미터를 관리할 수 있는 스크립트를 구성합니다.

import sys
from pathlib import Path
from ultralytics import YOLO

def init_training_environment():
    # 로컬 환경 변수 추가 시 주의 (필요한 경우에만 사용)
    project_root = Path(__file__).resolve().parent
    sys.path.insert(0, str(project_root / 'ultralytics'))
    
def run_detection_training():
    """
    YOLOv8 N 버전 모델 학습 함수
    """
    # 1. 모델 로드 (사전 훈련된 가중치를 기반으로 한 전이 학습 권장)
    model_name = "yolov8n.pt" 
    model = YOLO(model_name)

    # 2. 학습 인자 정의
    train_config = {
        'data': 'data/coco_dataset.yaml',      # 데이터셋 설정 파일 경로
        'epochs': 100,                         # 총 학습 에포크 수
        'patience': 30,                        # 검증 성능 미개선에 따른 조기 중단 임계값
        'batch': 16,                           # 배치 사이즈
        'imgsz': 640,                          # 입력 이미지 크기
        'device': [0, 1],                      # GPU 장치 목록 (멀티 GPU 지원)
        'workers': 4,                          # 데이터 로딩 워커 스레드
        'project': './outputs/train_exp',      # 저장 경로
        'name': 'run_001',                     # 실험 이름
        'optimizer': 'AdamW',                  # 최적화 알고리즘 선택
        'lr0': 0.01,                           # 초기 학습률
        'lrf': 0.01,                           # 최종 학습률 비율
        'mosaic': 1.0,                         # 모자이크 증강 활성화 (최종 몇 에포크 제외)
        'cache': True,                         # 메모리 캐싱 활성화로 속도 향상
    }

    # 3. 학습 수행 및 체크포인트 자동 저장
    model.train(**train_config)
    
    # 4. 검증 세팅 (선택 사항)
    val_metrics = model.val()
    return model, val_metrics

if __name__ == "__main__":
    trained_model, metrics = run_detection_training()

주요 하이퍼파라미터 설명

  • data: YOLO 학습을 위한 데이터셋 YAML 파일의 절대 또는 상대 경로입니다.
  • patience: 지정된 에포크 동안 mAP 이 개선되지 않으면 학습을 자동으로 종료합니다. 큰 값을 설정하면 강제 학습이 가능합니다.
  • device: cuda device 지정을 합니다. 단일 GPU 는 '0', 다중은 ['0','1'] 형식으로 배열됩니다.
  • mosaic: 학습 초기 데이터 믹싱 기법을 조절합니다. 마지막 단계에서는 이 값을 0 으로 낮추어 미세 조정 효과를 높일 수 있습니다.
  • cache: True 로 설정하면 데이터 로딩 속도를 크게 단축시킬 수 있으나 메모리 용량 확보가 필요합니다.

모델 추론 엔진 전환 (ONNX Export)

학습이 완료된 PyTorch 모델을 타사 플랫폼이나 경량화 된 추론 엔진에서 사용하기 위해 ONNX 형식으로 변환해야 합니다. 이는 배포 파이프라인의 호환성을 높이는 표준 절차입니다.

파이썬 스크립트 기반 변환

import torch
import onnx
from ultralytics import YOLO

def convert_to_onnx(checkpoint_path='weights/best.pt'):
    """
    YOLOv8 모델을 ONNX 포맷으로 내보내는 함수
    """
    # 1. 모델 복원
    loaded_model = YOLO(checkpoint_path)
    
    # 2. 더미 입력 텐서 생성 (배치 크기, 채널, 높이, 너비)
    dummy_input_shape = (1, 3, 640, 640)
    test_tensor = torch.randn(*dummy_input_shape)
    
    # 3. 출력 파일 경로
    target_filename = 'converted_model.onnx'
    
    # 4. TorchScript 또는 ONNX 직접 추출 방식 사용
    with torch.no_grad():
        torch.onnx.export(
            loaded_model,
            test_tensor,
            target_filename,
            input_names=['image'],
            output_names=['dets'],
            dynamic_axes={'image': {0: 'batch'}, 'dets': {0: 'batch'}},
            opset_version=12
        )
    
    print(f"모델 변환 완료: {target_filename}")
    return target_filename

명령행 인터페이스 (CLI) 를 통한 변환

스프립트 작성 없이 유틸리티 기능을 활용해 빠르게 변환할 수 있습니다.

yolo export model='./best.pt' format=onnx imgsz=640 simplify=True

위 명령어는 simplify=True 옵션을 포함하여 ONNX 그래프를 최적화하고 연산자를 통합합니다.

모델 정제 및 최적화

변환된 ONNX 모델은 복잡성이 여전히 존재할 수 있으므로 onnx-simplifier 도구를 사용하여 불필요한 노드를 제거하고 성능을 극대화합니다.

import onnx
import onnxsim

def optimize_onnx_model(input_path, output_path):
    onnx_model, check = onnxsim.simplify(onnx.load(input_path))
    
    if not check:
        raise Exception("정제 후 모델과 원본 간 상호 호환성 오류 발생")
        
    onnx.save(onnx_model, output_path)
    print(f"최적화된 모델 저장됨: {output_path}")

변환 검증 (Validation)

변환된 모델이 원래 PyTorch 모델과 동일한 출력 결과를 제공하는지 확인하는 테스트 코드는 필수적입니다. 여기서는 onnxruntime 을 사용하여 추론을 수행하고 결과값을 비교합니다.

import onnxruntime as ort
import numpy as np

def verify_model_conversion(pytorch_model_path, onnx_model_path):
    """
    ONNX 변환 결과 검증을 위해 원본 PyTorch 모델 과 비교
    """
    # ONNX 세션 초기화
    sess_options = ort.SessionOptions()
    session = ort.InferenceSession(onnx_model_path, sess_options)
    
    # 동일한 랜덤 입력 데이터 준비
    dummy_data = np.random.rand(1, 3, 640, 640).astype(np.float32)
    ort_input = {session.get_inputs()[0].name: dummy_data}
    
    # ONNX 추론 실행
    ort_output = session.run(None, ort_input)[0]
    
    # PyTorch 추론 실행 (반환되는 결과는 CPU 에서 처리 필요)
    torch_out = pytorch_model(dummy_data)
    torch_out_np = torch_out.cpu().detach().numpy()
    
    # 수치 오차 허용 범위 내에서 일치 여부 확인
    assert np.allclose(torch_out_np, ort_output, rtol=1e-3, atol=1e-5), \
        "출력 값 매칭 실패 - ONNX 변환 실패 가능"
        
    print("검증 성공: 두 모델 간 출력이 동일합니다.")

태그: YOLOv8 onnx PyTorch 객체감지 딥러닝배포

8월 19일 20:39에 게시됨