From adefc3ecb2df6729beac100033d81b3340e00867 Mon Sep 17 00:00:00 2001 From: zhanghuan <1262329256@qq.com> Date: Thu, 11 Sep 2025 13:39:49 +0800 Subject: [PATCH] =?UTF-8?q?=E8=8F=9C=E5=93=81=E8=AF=86=E5=88=AB=E7=A8=8B?= =?UTF-8?q?=E5=BA=8F=EF=BC=8C=E5=A2=9E=E5=8A=A0=E5=90=91=E9=87=8F=E8=BE=93?= =?UTF-8?q?=E5=87=BA=E5=88=B0=E6=8E=A7=E5=88=B6=E5=8F=B0=E7=9A=84=E5=8A=9F?= =?UTF-8?q?=E8=83=BD=EF=BC=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- classifier/food_classifier_app.py | 55 +++++++++++++++++++++---------- settings/settings.py | 7 ++-- toAndroid/toAndroid.py | 4 +-- train/train_food_classifier.py | 3 +- 4 files changed, 44 insertions(+), 25 deletions(-) diff --git a/classifier/food_classifier_app.py b/classifier/food_classifier_app.py index 277f248..16f99f4 100644 --- a/classifier/food_classifier_app.py +++ b/classifier/food_classifier_app.py @@ -66,7 +66,7 @@ class FoodClassifierApp: # 定义图像预处理(与训练时相同) self.transform = transforms.Compose([ - transforms.Resize((32, 32)), + transforms.Resize((32, 32),interpolation=transforms.InterpolationMode.BILINEAR), transforms.ToTensor(), transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)) ]) @@ -114,25 +114,26 @@ class FoodClassifierApp: image = cv2.imdecode(nparr, cv2.IMREAD_COLOR) if image is not None: + # image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) return image - # 方法2:如果方法1失败,尝试使用PIL - from PIL import Image as PILImage - pil_image = PILImage.open(file_path) - - # 转换为RGB(如果是RGBA) - if pil_image.mode == 'RGBA': - pil_image = pil_image.convert('RGB') - elif pil_image.mode == 'L': # 灰度图 - pil_image = pil_image.convert('RGB') - - # 转换为numpy数组 - image_array = np.array(pil_image) - - # PIL使用RGB,OpenCV使用BGR,需要转换 - image = cv2.cvtColor(image_array, cv2.COLOR_RGB2BGR) - - return image + # # 方法2:如果方法1失败,尝试使用PIL + # from PIL import Image as PILImage + # pil_image = PILImage.open(file_path) + # + # # 转换为RGB(如果是RGBA) + # if pil_image.mode == 'RGBA': + # pil_image = pil_image.convert('RGB') + # elif pil_image.mode == 'L': # 灰度图 + # pil_image = pil_image.convert('RGB') + # + # # 转换为numpy数组 + # image_array = np.array(pil_image) + # + # # PIL使用RGB,OpenCV使用BGR,需要转换 + # image = cv2.cvtColor(image_array, cv2.COLOR_RGB2BGR) + # + # return image except Exception as e: print(f"加载图片失败: {e}") @@ -549,12 +550,30 @@ class FoodClassifierApp: pil_image = Image.fromarray(image_rgb) # 应用预处理 + resize_transform = transforms.Resize((32,32),interpolation=transforms.InterpolationMode.BICUBIC) + resized_image = resize_transform(pil_image) + if isinstance(resized_image, Image.Image): + # 转换为tensor但不归一化 + to_tensor = transforms.ToTensor() + resized_tensor = to_tensor(resized_image) + print(f"缩放后tensor形状: {resized_tensor.shape}") + + # 打印前5个像素值(每个通道) + print("前5个像素值 (R, G, B):") + for i in range(min(5, resized_tensor.shape[1])): + r_val = resized_tensor[0, 0, i].item() * 255 # Red通道 (转换回0-255范围) + g_val = resized_tensor[1, 0, i].item() * 255 # Green通道 + b_val = resized_tensor[2, 0, i].item() * 255 # Blue通道 + print(f" 像素[0,{i}]: R={r_val:.2f}, G={g_val:.2f}, B={b_val:.2f}") + + input_tensor = self.transform(pil_image).unsqueeze(0) # 添加batch维度 input_tensor = input_tensor.to(self.device) # 进行预测 with torch.no_grad(): outputs = self.model(input_tensor) + print('outputs',outputs) probabilities = F.softmax(outputs, dim=1) confidence, predicted = torch.max(probabilities, 1) diff --git a/settings/settings.py b/settings/settings.py index f5a49f3..c92e99a 100644 --- a/settings/settings.py +++ b/settings/settings.py @@ -15,17 +15,16 @@ VAL_DATA_DIR = os.path.join(DATASET_DIR, 'val') TEST_DATA_DIR = os.path.join(DATASET_DIR, 'test') # 模型保存路径 -MODEL_DIR = os.path.join(BASE_DIR, 'model', '05') +MODEL_DIR = os.path.join(BASE_DIR, 'model', '06') BEST_MODEL_PATH = os.path.join(MODEL_DIR, 'best_food_model.pth') TRAINING_CURVES_PATH = os.path.join(MODEL_DIR, 'training_curves.png') TRAINING_RESULTS_PATH = os.path.join(MODEL_DIR, 'training_results.txt') -# INFERENCE_BEST_MODEL_PATH = os.path.join(BASE_DIR, 'model', '03','best_food_model.pth') +# INFERENCE_BEST_MODEL_PATH = os.path.join(BASE_DIR, 'model', '05','best_food_model.pth') INFERENCE_BEST_MODEL_PATH = BEST_MODEL_PATH # 训练参数 -# NUM_EPOCHS = 100 -NUM_EPOCHS = 13 +NUM_EPOCHS = 100 BATCH_SIZE = 32 LEARNING_RATE = 0.001 WEIGHT_DECAY = 1e-4 diff --git a/toAndroid/toAndroid.py b/toAndroid/toAndroid.py index 04fca23..f30b533 100644 --- a/toAndroid/toAndroid.py +++ b/toAndroid/toAndroid.py @@ -7,7 +7,7 @@ from net import create_food_cnn # 1. 初始化模型 model = create_food_cnn() # 2. 加载训练好的权重 -model.load_state_dict(torch.load("../model/02/best_food_model.pth", map_location='cpu')) +model.load_state_dict(torch.load("../model/06/best_food_model.pth", map_location='cpu')) model.eval() # 设置为推理模式 # 3. 创建示例输入 (假设输入是 3x224x224 的图片) @@ -15,4 +15,4 @@ example_input = torch.randn(1, 3, 224, 224) # 4. 转换为 TorchScript traced_script_module = torch.jit.trace(model, example_input) -traced_script_module.save("../model/02/best_food_model_mobile.pt") +traced_script_module.save("../model/06/best_food_model_mobile.pt") diff --git a/train/train_food_classifier.py b/train/train_food_classifier.py index fd9e203..2f99be3 100644 --- a/train/train_food_classifier.py +++ b/train/train_food_classifier.py @@ -14,8 +14,9 @@ import time # 添加net目录到路径 # sys.path.append(os.path.join(os.path.dirname(__file__), '..', 'net')) -# from food_net import create_food_cnn +sys.path.append(os.path.join(os.path.dirname(__file__), '..')) from net import create_food_cnn +# from net import create_food_cnn from settings import settings # 设置matplotlib支持中文显示