Flask集成CNN脑肿瘤检测模型时TypeError错误的解决求助
脑肿瘤检测CNN模型集成Flask网站的模型加载错误解决
问题场景
开发用于脑肿瘤检测的CNN模型,已完成训练并保存,尝试集成到Flask网站时出现加载错误,无法理解错误成因。
代码实现
from flask import Flask, render_template, request, redirect, url_for import os from werkzeug.utils import secure_filename import torch from torchvision import transforms from PIL import Image import torch.nn as nn import torch.nn.functional as F import cv2 import numpy as np import glob from torch.utils.data import Dataset, DataLoader, ConcatDataset class Dataset(object): """An abstract class representing a Dataset. All other datasets should subclass it. All subclasses should override ``__len__``, that provides the size of the dataset, and ``__getitem__``, supporting integer indexing in range from 0 to len(self) exclusive. """ def __getitem__(self, index): raise NotImplementedError def __len__(self): raise NotImplementedError def __add__(self, other): return ConcatDataset([self, other]) class MRI(Dataset): def __init__(self): tumor = [] healthy = [] # cv2 - It reads in BGR format by default for f in glob.iglob(r"C:\Users\ishan\Desktop\Ishan\Coding\Hackathon\Dataset\Yes_Data/*.jpg"): img = cv2.imread(f) img = cv2.resize(img,(128,128)) # I can add this later in the boot-camp for more adventure b, g, r = cv2.split(img) img = cv2.merge([r,g,b]) img = img.reshape((img.shape[2],img.shape[0],img.shape[1])) # otherwise the shape will be (h,w,#channels) tumor.append(img) for f in glob.iglob(r"C:\Users\ishan\Desktop\Ishan\Coding\Hackathon\Dataset\No_data/*.jpg"): img = cv2.imread(f) img = cv2.resize(img,(128,128)) b, g, r = cv2.split(img) img = cv2.merge([r,g,b]) img = img.reshape((img.shape[2],img.shape[0],img.shape[1])) healthy.append(img) # our images tumor = np.array(tumor,dtype=np.float32) healthy = np.array(healthy,dtype=np.float32) # our labels tumor_label = np.ones(tumor.shape[0], dtype=np.float32) healthy_label = np.zeros(healthy.shape[0], dtype=np.float32) # Concatenates self.images = np.concatenate((tumor, healthy), axis=0) self.labels = np.concatenate((tumor_label, healthy_label)) def __len__(self): return self.images.shape[0] def __getitem__(self, index): sample = {'image': self.images[index], 'label':self.labels[index]} return sample def normalize(self): self.images = self.images/255.0 class CNN(nn.Module): def __init__(self): super(CNN,self).__init__() self.cnn_model = nn.Sequential( nn.Conv2d(in_channels=3, out_channels=6, kernel_size=5), nn.Tanh(), nn.AvgPool2d(kernel_size=2, stride=5), nn.Conv2d(in_channels=6, out_channels=16, kernel_size=5), nn.Tanh(), nn.AvgPool2d(kernel_size=2, stride=5)) self.fc_model = nn.Sequential( nn.Linear(in_features=256, out_features=120), nn.Tanh(), nn.Linear(in_features=120, out_features=84), nn.Tanh(), nn.Linear(in_features=84, out_features=1)) def forward(self, x): x = self.cnn_model(x) x = x.view(x.size(0), -1) x = self.fc_model(x) x = torch.sigmoid(x) return x app = Flask(__name__) app.config['UPLOAD_FOLDER'] = 'static/uploads' app.config['ALLOWED_EXTENSIONS'] = {'jpg'} path = r'C:\Users\ishan\Desktop\Ishan\Coding\Project\FinalModel.pth' # Create an instance of the model model_instance = CNN() # Load your trained model model_instance.load_state_dict(torch.load('FinalModel.pth', map_location=torch.device('cpu'))) # Set the model in evaluation mode model_instance.eval() def allowed_file(filename): return '.' in filename and filename.rsplit('.', 1)[1].lower() in app.config['ALLOWED_EXTENSIONS'] def preprocess_image(image_path): transform = transforms.Compose([ transforms.Resize((128, 128)), transforms.ToTensor(), ]) image = Image.open(image_path).convert('RGB') return transform(image).unsqueeze(0) @app.route('/') def index(): return render_template('index.html') @app.route('/', methods=['POST']) def upload_file(): if 'file' not in request.files: return redirect(request.url) file = request.files['file'] if file.filename == '': return redirect(request.url) if file and allowed_file(file.filename): filename = secure_filename(file.filename) file_path = os.path.join(app.config['UPLOAD_FOLDER'], filename) file.save(file_path) # Make predictions input_tensor = preprocess_image(file_path) with torch.no_grad(): output = model(input_tensor) # Output is a tensor prediction = "Tumor Detected" if output.item() > 0.5 else "No Tumor" return render_template('index.html', prediction=prediction, image_path=file_path) return redirect(request.url) if __name__ == '__main__': app.run(debug=True)
报错信息
Traceback (most recent call last): File "C:\Users\ishan\Desktop\Ishan\Coding\Project\app.py", line 113, in <module> model_instance.load_state_dict(torch.load('FinalModel.pth', map_location=torch.device('cpu'))) File "C:\Users\ishan\AppData\Local\Programs\Python\Python311\Lib\site-packages\torch\nn\modules\module.py", line 2103, in load_state_dict raise TypeError(f"Expected state_dict to be dict-like, got {type(state_dict)}.") TypeError: Expected state_dict to be dict-like, got <class 'type'>
错误原因与解决方法
核心问题
torch.load('FinalModel.pth')加载返回的不是模型的state_dict字典,而是模型类本身(<class 'type'>),说明保存模型时错误地保存了模型类,而非模型实例的state_dict。
解决步骤
修正模型保存代码
训练完成后,必须保存模型实例的state_dict,而不是模型类或整个模型。正确的保存代码应为:# 训练时的保存逻辑 model_instance = CNN() # ...训练过程... torch.save(model_instance.state_dict(), 'FinalModel.pth')若之前误保存了模型类(如
torch.save(CNN, 'FinalModel.pth')),需重新训练并按上述方式保存,或找到正确的state_dict文件替换。修正模型加载路径
代码中定义了绝对路径变量path但未使用,加载时用的是相对路径'FinalModel.pth',需确保文件存在于Flask运行目录,或直接使用绝对路径:model_instance.load_state_dict(torch.load(path, map_location=torch.device('cpu')))修复预测时的模型实例引用错误
代码中预测部分使用了未定义的model变量,应改为已实例化的model_instance:with torch.no_grad(): output = model_instance(input_tensor) # 替换model为model_instance prediction = "Tumor Detected" if output.item() > 0.5 else "No Tumor"
内容的提问来源于stack exchange,提问作者Ishan Joshi
相关产品推荐
相关产品推荐

