new train
This commit is contained in:
@@ -1,10 +1,12 @@
|
||||
import os
|
||||
import cv2
|
||||
import glob
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
import onnxruntime as ort # 【修改1】替换torch导入
|
||||
from torchvision import transforms # 保留transforms,仍用于预处理
|
||||
from core.const import label_name, input_size
|
||||
from core.const import mode, label_name, input_size
|
||||
|
||||
|
||||
def test():
|
||||
@@ -12,13 +14,18 @@ def test():
|
||||
providers = ['CUDAExecutionProvider', 'CPUExecutionProvider'] if 'CUDAExecutionProvider' in ort.get_available_providers() else ['CPUExecutionProvider']
|
||||
print(f"使用设备: {providers[0]}")
|
||||
|
||||
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
core_dir = os.path.dirname(script_dir)
|
||||
model_dir = os.path.join(core_dir, "models")
|
||||
dataset_dir = os.path.join(core_dir, "dataset", mode, "test")
|
||||
|
||||
# 【修改3】加载ONNX模型,替代原来的PyTorch模型加载
|
||||
session = ort.InferenceSession("./model/resnet_final.onnx", providers=providers)
|
||||
|
||||
session = ort.InferenceSession(os.path.join(model_dir, "resnet18_epoch_50_bak2.onnx"), providers=providers)
|
||||
|
||||
# 获取输入名称(用于后续推理时指定输入)
|
||||
input_name = session.get_inputs()[0].name
|
||||
|
||||
im_list = glob.glob("./dataset/test/*/*.jpg")
|
||||
im_list = glob.glob(os.path.join(dataset_dir, "*", "*.jpg"))
|
||||
np.random.shuffle(im_list)
|
||||
|
||||
# 预处理完全不变
|
||||
@@ -49,6 +56,7 @@ def test():
|
||||
|
||||
img = np.asarray(im_data)
|
||||
img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)
|
||||
img = cv2.resize(img, (200, int(img.shape[0] * 200 / img.shape[1])))
|
||||
cv2.imshow("img", img)
|
||||
cv2.waitKey(0)
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import os
|
||||
import cv2
|
||||
import glob
|
||||
import torch
|
||||
@@ -5,17 +6,26 @@ from torchvision import transforms
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
from core.nets.resnet18 import resnet18
|
||||
from core.const import label_name, input_size
|
||||
from core.const import mode, label_name, input_size
|
||||
|
||||
|
||||
def test():
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
print(device)
|
||||
|
||||
net = resnet18()
|
||||
net.load_state_dict(torch.load("./model/resnet_epoch_14.pth", weights_only=True))
|
||||
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
core_dir = os.path.dirname(script_dir)
|
||||
model_dir = os.path.join(core_dir, "models")
|
||||
dataset_dir = os.path.join(core_dir, "dataset", mode, "test")
|
||||
|
||||
im_list = glob.glob("./dataset/test/*/*.jpg")
|
||||
print("model_dir", model_dir)
|
||||
|
||||
net = resnet18()
|
||||
# net.load_state_dict(torch.load("./models/resnet18_epoch_100.pth", weights_only=True))
|
||||
net.load_state_dict(torch.load(os.path.join(model_dir, "resnet18_epoch_50_bak2.pth"), weights_only=True))
|
||||
|
||||
# im_list = glob.glob("./dataset/test/*/*.jpg")
|
||||
im_list = glob.glob(os.path.join(dataset_dir, "*", "*.jpg"))
|
||||
np.random.shuffle(im_list)
|
||||
|
||||
net.to(device)
|
||||
@@ -41,12 +51,13 @@ def test():
|
||||
_, pred = torch.max(outputs.data, dim=1)
|
||||
print(label_name[pred.cpu().numpy()[0]], " ", im_path)
|
||||
|
||||
# prob, pred = torch.topk(outputs.data, k=3, dim=1)
|
||||
prob, pred = torch.topk(outputs.data, k=3, dim=1)
|
||||
# for i in range(3):
|
||||
# print(label_name[pred[0, i].item()], " ", prob[0, i].item(), " ", im_path)
|
||||
|
||||
img = np.asarray(im_data)
|
||||
img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)
|
||||
img = cv2.resize(img, (200, 200))
|
||||
cv2.imshow("img", img)
|
||||
cv2.waitKey(0)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user