skytnt/anime-segmentation
Updated • 343 • 51
アニメイラストに特化した高精度な背景削除(セグメンテーション)モデルです。 強力なエンコーダーと軽量なデコーダーを組み合わせ、INT8量子化を施すことで、性能と速度を両立させました。
以下のコードで簡単に推論を試すことができます。
import cv2
import numpy as np
import onnxruntime as ort
from PIL import Image
# 設定
MODEL_PATH = "birefnext-aniseg-int8-v0.1.onnx"
INPUT_IMAGE = "input.jpg"
OUTPUT_IMAGE = "output_mask.png"
# ImageNetの正規化パラメータ
MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
providers = ['CUDAExecutionProvider', 'CPUExecutionProvider'] if ort.get_device() == 'GPU' else ['CPUExecutionProvider']
session = ort.InferenceSession(MODEL_PATH, providers=providers)
input_name = session.get_inputs()[0].name
output_name = session.get_outputs()[0].name
img = Image.open(INPUT_IMAGE).convert('RGB')
w0, h0 = img.size
# 32の倍数にリサイズ
w, h = (w0 // 32) * 32, (h0 // 32) * 32
img_resized = img.resize((w, h), Image.BILINEAR)
# 前処理
img_array = np.array(img_resized, dtype=np.float32) / 255.0
img_array = (img_array - MEAN) / STD
input_tensor = img_array.transpose(2, 0, 1)[None]
# 推論
result = session.run([output_name], {input_name: input_tensor})[0]
# 後処理
mask = result[0, 0]
mask_resized = cv2.resize(mask, (w0, h0), interpolation=cv2.INTER_LINEAR)
mask_uint8 = (np.clip(mask_resized, 0, 1) * 255).astype(np.uint8)
cv2.imwrite(OUTPUT_IMAGE, mask_uint8)
print(f"Saved: {OUTPUT_IMAGE}")