FastAI自定义TTA替代predict时Python版本兼容问题(fasttransform/fastcore相关)
我最近在Cloud Run部署FastAI预测服务时,碰到了个棘手的问题:原本用load_learner加载模型、调用learn.predict做预测的服务运行一切正常,想换成自定义的tta_predict函数来用测试时数据增强(TTA)提升预测精度,结果一堆麻烦找上门来。
先给大家看看我写的那个想替代learn.predict的tta_predict函数,在Colab的测试环境里跑着完全没问题:
import random from fastai.vision.all import * # 自定义TTA预测函数,输出格式和learn.predict对齐 def tta_predict(learner, img): # 为单张图片创建测试DataLoader test_dl = learner.dls.test_dl([img]) # 执行TTA preds, _ = learner.tta(dl=test_dl) # 计算平均概率 avg_probs = preds.mean(dim=0) # 获取预测类别索引 pred_idx = avg_probs.argmax().item() # 获取类别标签 class_label = learner.dls.vocab[pred_idx] # 返回和learn.predict一致的格式:(类别标签, 索引, 概率数组) return (class_label, pred_idx, avg_probs) # 使用示例 prediction = tta_predict(learn, grayscale_img) # 验证输出格式 print(type(prediction)) print(prediction) print(prediction[0]) # 类别标签 print(prediction[2]) # 平均概率
但把这段代码加到生产脚本里,替换掉原来的learn.predict后,问题就开始了:
问题1:Python 3.9下构建失败
报错信息很长,核心问题是FastAI的依赖库用了Python 3.10才支持的PEP 604联合类型语法(比如PILBase | TensorImageBase),而Python 3.9无法识别这个|运算符:
Traceback (most recent call last): File "/app/main.py", line 11, in
from fastai.vision.all import PILImage, BCEWithLogitsLossFlat, load_learner
File "/usr/local/lib/python3.9/site-packages/fastai/vision/all.py", line 4,
infrom .augment import * File "/usr/local/lib/python3.9/
site-packages/fastai/vision/augment.py", line 8, infrom .core import * File "/usr/local/lib/python3.9/site-packages/fastai/vision/core.py", line 259, in class PointScaler(Transform): File "/usr/local/lib/python3.9/site-packages/fasttransform/transform.py", line 75, in new if funcs: setattr(new_cls, nm, _merge_funcs(*funcs)) File "/usr/local/lib/python3.9/site-packages/fasttransform/transform.py", line 42, in _merge_funcs res = Function(fs[-1].methods[0].implementation) File "/usr/local/lib/python3.9/site-packages/plum/function.py", line 181, in methods self._resolve_pending_registrations() File "/usr/local/lib/python3.9/site-packages/plum/function.py", line 280, in _resolve_pending_registrations signature = Signature.from_callable(f, precedence=precedence) File "/usr/local/lib/python3.9/site-packages/plum/signature.py", line 88, in from_callable types, varargs = _extract_signature(f) File "/usr/local/lib/python3.9/site-packages/plum/signature.py", line 346, in _extract_signature resolve_pep563(f) File "/usr/local/lib/python3.9/site-packages/plum/signature.py", line 329, in resolve_pep563 beartype_resolve_pep563(f) # This mutates f. File "/usr/local/lib/python3.9/site-packages/beartype/peps/_pep563.py", line 263, in resolve_pep563 arg_name_to_hint[arg_name] = resolve_hint( File "/usr/local/lib/python3.9/site-packages/beartype/_check/forward/fwdmain.py", line 308, in resolve_hint return _resolve_func_scope_forward_hint( File "/usr/local/lib/python3.9/site-packages/beartype/_check/forward/fwdmain.py", line 855, in _resolve_func_scope_forward_hint raise exception_cls(exception_message) from exception beartype.roar.BeartypeDecorHintPep604Exception: Stringified PEP 604 type hint 'PILBase | TensorImageBase' syntactically invalid under Python < 3.10 (i.e., TypeError("unsupported operand type(s) for |: 'BypassNewMeta' and 'torch._C._TensorMeta'")). Consider either:
* Requiring Python >= 3.10. Abandon Python < 3.10 all ye who code here.
* Refactoring PEP 604 type hints into equivalent PEP 484 type hints: e.g.,
# Instead of this...
from future import annotations
def bad_func() -> int | str: ...
# Do this. Ugly, yet it works. Worky >>> pretty.
from typing import Union
我仔细检查了自己的代码,根本没用到这种语法,显然是引入TTA相关的FastAI模块(报错里提到了augment)后,依赖库的版本和Python3.9不兼容导致的。
问题2:切换到Python 3.10后,模型加载失败
换成Python3.10后,构建终于通过了,但运行时加载模型又报错:
ERROR loading model.pkl: Could not import 'Pipeline' from fastcore.transform - this module has been moved to the fasttransform package.
To migrate your code, please see the migration guide at: https://answerdotai.github.io/fasttransform/fastcore_migration_guide.html
可我自己的代码从来没有直接导入Pipeline、Transform或者fastcore啊!这明显是模型文件(model.pkl)在保存的时候,依赖的是旧版的fastcore,现在部署环境里的FastAI用了新的fasttransform包,导致加载时找不到对应的模块。
备注:内容来源于stack exchange,提问作者Hack-R

