图像压缩到256
This commit is contained in:
@@ -43,6 +43,7 @@ class SegFormerInference:
|
|||||||
model_path: Optional[str] = None,
|
model_path: Optional[str] = None,
|
||||||
pretrained_model: str = "nvidia/segformer-b0-finetuned-ade-512-512",
|
pretrained_model: str = "nvidia/segformer-b0-finetuned-ade-512-512",
|
||||||
num_classes: int = 2,
|
num_classes: int = 2,
|
||||||
|
image_size: int = 256, # ⚠️ 重要:必须与训练时一致!
|
||||||
device: str = "auto"
|
device: str = "auto"
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -53,9 +54,11 @@ class SegFormerInference:
|
|||||||
如果为None,则使用预训练模型
|
如果为None,则使用预训练模型
|
||||||
pretrained_model: 预训练模型名称(用于加载processor)
|
pretrained_model: 预训练模型名称(用于加载processor)
|
||||||
num_classes: 类别数
|
num_classes: 类别数
|
||||||
|
image_size: 输入图像尺寸(必须与训练时一致!)
|
||||||
device: 设备 ('cpu', 'cuda', 'auto')
|
device: 设备 ('cpu', 'cuda', 'auto')
|
||||||
"""
|
"""
|
||||||
self.num_classes = num_classes
|
self.num_classes = num_classes
|
||||||
|
self.image_size = image_size
|
||||||
|
|
||||||
# 设置设备
|
# 设置设备
|
||||||
if device == "auto":
|
if device == "auto":
|
||||||
@@ -68,7 +71,11 @@ class SegFormerInference:
|
|||||||
# ⚠️ 重要:使用与训练时完全一致的预处理
|
# ⚠️ 重要:使用与训练时完全一致的预处理
|
||||||
# 不再使用 SegformerImageProcessor,而是手动构建预处理pipeline
|
# 不再使用 SegformerImageProcessor,而是手动构建预处理pipeline
|
||||||
print(f"构建预处理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([
|
self.transform = A.Compose([
|
||||||
|
A.Resize(image_size, image_size), # ⚠️ 关键:必须resize到训练时的尺寸
|
||||||
A.Normalize(
|
A.Normalize(
|
||||||
mean=[0.485, 0.456, 0.406], # ImageNet标准均值
|
mean=[0.485, 0.456, 0.406], # ImageNet标准均值
|
||||||
std=[0.229, 0.224, 0.225], # ImageNet标准标准差
|
std=[0.229, 0.224, 0.225], # ImageNet标准标准差
|
||||||
@@ -371,7 +378,8 @@ def compare_models(
|
|||||||
pretrained_inference = SegFormerInference(
|
pretrained_inference = SegFormerInference(
|
||||||
model_path=None,
|
model_path=None,
|
||||||
pretrained_model=pretrained_model,
|
pretrained_model=pretrained_model,
|
||||||
num_classes=150 # ADE20K的类别数
|
num_classes=150, # ADE20K的类别数
|
||||||
|
image_size=256 # 与Fine-tune模型保持一致
|
||||||
)
|
)
|
||||||
|
|
||||||
# 加载Fine-tune模型
|
# 加载Fine-tune模型
|
||||||
@@ -379,7 +387,8 @@ def compare_models(
|
|||||||
finetuned_inference = SegFormerInference(
|
finetuned_inference = SegFormerInference(
|
||||||
model_path=finetuned_model_path,
|
model_path=finetuned_model_path,
|
||||||
pretrained_model=pretrained_model,
|
pretrained_model=pretrained_model,
|
||||||
num_classes=2
|
num_classes=2,
|
||||||
|
image_size=256 # ⚠️ 必须与训练时一致
|
||||||
)
|
)
|
||||||
|
|
||||||
# 读取图像
|
# 读取图像
|
||||||
@@ -469,7 +478,8 @@ def main():
|
|||||||
# 创建推理实例
|
# 创建推理实例
|
||||||
inference = SegFormerInference(
|
inference = SegFormerInference(
|
||||||
model_path=FINETUNED_MODEL_PATH,
|
model_path=FINETUNED_MODEL_PATH,
|
||||||
num_classes=2
|
num_classes=2,
|
||||||
|
image_size=256 # ⚠️ 必须与训练时一致(见config.py第272行)
|
||||||
)
|
)
|
||||||
|
|
||||||
# 单张图像测试
|
# 单张图像测试
|
||||||
|
|||||||
Reference in New Issue
Block a user