From 486241261c1d3c9c31911da39480f164a656f76a Mon Sep 17 00:00:00 2001 From: zhangpu <1250681871@qq.com> Date: Fri, 12 Dec 2025 14:17:21 +0800 Subject: [PATCH] =?UTF-8?q?=E5=9B=BE=E5=83=8F=E5=8E=8B=E7=BC=A9=E5=88=B025?= =?UTF-8?q?6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- SegFormer/inference/test_model.py | 16 +++++++++++++--- 1 file changed, 13 insertions(+), 3 deletions(-) diff --git a/SegFormer/inference/test_model.py b/SegFormer/inference/test_model.py index 5770d8c..ae5f4f8 100644 --- a/SegFormer/inference/test_model.py +++ b/SegFormer/inference/test_model.py @@ -43,6 +43,7 @@ class SegFormerInference: model_path: Optional[str] = None, pretrained_model: str = "nvidia/segformer-b0-finetuned-ade-512-512", num_classes: int = 2, + image_size: int = 256, # ⚠️ 重要:必须与训练时一致! device: str = "auto" ): """ @@ -53,9 +54,11 @@ class SegFormerInference: 如果为None,则使用预训练模型 pretrained_model: 预训练模型名称(用于加载processor) num_classes: 类别数 + image_size: 输入图像尺寸(必须与训练时一致!) device: 设备 ('cpu', 'cuda', 'auto') """ self.num_classes = num_classes + self.image_size = image_size # 设置设备 if device == "auto": @@ -68,7 +71,11 @@ class SegFormerInference: # ⚠️ 重要:使用与训练时完全一致的预处理 # 不再使用 SegformerImageProcessor,而是手动构建预处理pipeline print(f"构建预处理Pipeline(与训练时一致)") + print(f" 图像尺寸: {image_size}×{image_size}") + print(f" 归一化: ImageNet标准(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])") + self.transform = A.Compose([ + A.Resize(image_size, image_size), # ⚠️ 关键:必须resize到训练时的尺寸 A.Normalize( mean=[0.485, 0.456, 0.406], # ImageNet标准均值 std=[0.229, 0.224, 0.225], # ImageNet标准标准差 @@ -371,7 +378,8 @@ def compare_models( pretrained_inference = SegFormerInference( model_path=None, pretrained_model=pretrained_model, - num_classes=150 # ADE20K的类别数 + num_classes=150, # ADE20K的类别数 + image_size=256 # 与Fine-tune模型保持一致 ) # 加载Fine-tune模型 @@ -379,7 +387,8 @@ def compare_models( finetuned_inference = SegFormerInference( model_path=finetuned_model_path, pretrained_model=pretrained_model, - num_classes=2 + num_classes=2, + image_size=256 # ⚠️ 必须与训练时一致 ) # 读取图像 @@ -469,7 +478,8 @@ def main(): # 创建推理实例 inference = SegFormerInference( model_path=FINETUNED_MODEL_PATH, - num_classes=2 + num_classes=2, + image_size=256 # ⚠️ 必须与训练时一致(见config.py第272行) ) # 单张图像测试