You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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。

解决步骤

  1. 修正模型保存代码
    训练完成后,必须保存模型实例的state_dict,而不是模型类或整个模型。正确的保存代码应为:

    # 训练时的保存逻辑
    model_instance = CNN()
    # ...训练过程...
    torch.save(model_instance.state_dict(), 'FinalModel.pth')
    

    若之前误保存了模型类(如torch.save(CNN, 'FinalModel.pth')),需重新训练并按上述方式保存,或找到正确的state_dict文件替换。

  2. 修正模型加载路径
    代码中定义了绝对路径变量path但未使用,加载时用的是相对路径'FinalModel.pth',需确保文件存在于Flask运行目录,或直接使用绝对路径:

    model_instance.load_state_dict(torch.load(path, map_location=torch.device('cpu')))
    
  3. 修复预测时的模型实例引用错误
    代码中预测部分使用了未定义的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.01 05:40:55