Views
No views yet
1import onnxruntime as ort
2import numpy as np
3from transformers import AutoFeatureExtractor
4from PIL import Image
5
6# ONNX 모델 경로
7onnx_model_path = r'C:\mobilevit_model.onnx'
8
9# ONNX 런타임 세션 초기화
10ort_session = ort.InferenceSession(onnx_model_path)
11
12# 새로운 이미지 예측 함수 정의
13def predict_image(image_path):
14 # MobileViT 모델에 맞는 특징 추출기 로드
15 feature_extractor = AutoFeatureExtractor.from_pretrained("apple/mobilevit-small")
16
17 # 이미지를 로드하고 RGB로 변환
18 image = Image.open(image_path).convert("RGB")
19
20 # 이미지를 특징 추출기로 전처리
21 inputs = feature_extractor(images=image, return_tensors="np")
22 input_array = inputs['pixel_values'] # ONNX는 Numpy 형식을 사용
23
24 # ONNX 모델에 입력 전달 및 추론
25 ort_inputs = {ort_session.get_inputs()[0].name: input_array}
26 ort_outputs = ort_session.run(None, ort_inputs)
27
28 # 결과 해석
29 logits = ort_outputs[0]
30 predicted_class = np.argmax(logits, axis=-1).item()
31
32 return "그냥 사진" if predicted_class == 1 else "로맨스 스캠 사진"
33
34# 예측 예시
35image_path = r'C:\1234567.jpg'
36result = predict_image(image_path)
37print(result)
38