在Flask应用中调用fastai load_learner()遇命名空间错误求助
问题解决:fastai load_learner 找不到 label_func
问题原因
你遇到的错误源于pickle序列化模型时会记录label_func的模块命名空间:训练模型时,label_func大概率是在非__main__的模块中定义的,但用gunicorn部署Flask时,你的脚本被当作__main__模块加载。即便你在脚本里定义了label_func,要么是作用域错误(比如类内的方法属于实例,而非模块级),要么是gunicorn的加载逻辑导致pickle无法在当前__main__模块中找到该函数。
另外,你在ClassifySceenshot类里定义的label_func是实例方法,而pickle寻找的是模块级全局函数,这个类内定义完全无效,只会造成混淆。
解决方案
方案1:统一模块导入(推荐)
将label_func放到独立工具模块,训练和部署时统一从该模块导入,确保命名空间一致:
- 创建
src/utils.py,写入与训练时完全一致的label_func逻辑:
def label_func(f): # 必须和训练模型时的实现完全相同,不能随便返回"test" # 示例:从文件名父目录提取标签 # return f.parent.name return ...
- 训练模型时,导入该函数并使用:
from src.utils import label_func # 用该函数创建DataLoaders并训练 dls = ImageDataLoaders.from_name_func(..., label_func=label_func) learn = vision_learner(dls, ...) learn.export("testing.pkl")
- 部署代码中同样导入该函数,删除多余定义:
from flask import Flask from flask_restful import Api, Resource import os from fastai.vision.all import * from src.classifier import Classifier from src.utils import label_func # 关键:统一模块导入 app = Flask(__name__) api = Api(app) class ClassifySceenshot(Resource): def get(self): filename = "tmp_0.jpg" try: c = Classifier() c.save_capture(filename) with open(filename, "rb") as img: image = img.read() bytes_array = bytearray(image) learner = load_learner("testing.pkl") prediction, pred_idx, probability = learner.predict(bytes_array) os.remove(filename) return {'status': 'success', 'prediction': prediction, 'probability': float(probability[pred_idx]), # 转换为JSON可序列化类型 }, 200 except Exception as e: if os.path.exists(filename): # 先判断文件存在再删除,避免额外报错 os.remove(filename) return {'status': 'error', 'message': str(e)}, 500 api.add_resource(ClassifySceenshot, '/classify') if __name__ == '__main__': app.run(port=XXXX, debug=True)
方案2:确保模块级函数在gunicorn加载时可见
若不想修改训练代码,需保证label_func在gunicorn加载的__main__模块中存在:
- 删除类内的
label_func,仅保留全局作用域的定义,且逻辑必须与训练时完全一致:
from flask import Flask from flask_restful import Api, Resource import os from fastai.vision.all import * from src.classifier import Classifier app = Flask(__name__) api = Api(app) # 全局定义,必须和训练时的label_func逻辑一致 def label_func(f): # 原训练代码中的实现 return ... class ClassifySceenshot(Resource): def get(self): # 原有逻辑不变 ... api.add_resource(ClassifySceenshot, '/classify') if __name__ == '__main__': app.run(port=XXXX, debug=True)
- 用gunicorn启动时直接指定脚本文件,确保脚本以
__main__模块加载时函数存在:
gunicorn your_script_name:app
关键注意点
label_func的逻辑必须与训练时完全一致,不能随意返回测试值,否则会导致模型标签映射错误,预测结果失效。- 部署时需确保模型文件
testing.pkl的路径正确,gunicorn的工作目录需能访问到该文件。
内容的提问来源于stack exchange,提问作者ETray
相关产品推荐
相关产品推荐

