Apache MXNet序列化模型反序列化报错问题咨询
解决MXNet加载SageMaker训练模型的报错问题
嘿,我看你踩了个常见的小坑——用mx.ndarray.load()加载SageMaker导出的model_algo-1肯定会报错的,因为这个方法是用来加载单独的NDArray数组文件的,而你手里的model_algo-1是整个MXNet模型的序列化文件(包含计算图+参数),得用专门的模型加载方法才行。
给你两种最常用的解决方案,对应不同的MXNet模型类型:
方案1:加载Gluon模型(大部分SageMaker训练的MXNet模型都是这种)
如果你的模型是用Gluon API训练的,直接用SymbolBlock.imports()就能搞定,步骤很简单:
- 先回忆下你训练时输入数据的变量名(比如通常是
data) - 运行这段代码:
import mxnet as mx # 替换成你训练时的输入变量名,比如'data' input_names = ['data'] # 加载模型,ctx选cpu或gpu都行 model = mx.gluon.SymbolBlock.imports( 'model_algo-1', input_names, ctx=mx.cpu() ) # 打印模型结构确认加载成功 print(model)
方案2:加载Symbol API训练的老模型
如果是用MXNet早期的Symbol API训练的模型,得用模块加载的方式:
import mxnet as mx # 先加载模型的计算图符号 sym = mx.sym.load('model_algo-1') # 创建模型模块 mod = mx.mod.Module(symbol=sym, context=mx.cpu()) # 绑定输入形状,这里的形状要和你训练时的输入一致,比如(1, 3, 224, 224)是单张RGB图的形状 mod.bind(for_training=False, data_shapes=[('data', (1, 3, 224, 224))]) # 加载模型参数 mod.load_params('model_algo-1')
为啥原来的代码不行?
再给你理清楚原因:mx.ndarray.load()只能读取用mx.ndarray.save()保存的单个或多个NDArray数组,而SageMaker导出的model_algo-1是把整个模型的计算图和参数序列化到一起的文件,格式完全不匹配,自然会报错啦。
另外提个小技巧:如果解压model.tar.gz后看到model-shapes.json文件,里面会明确写着模型的输入输出形状,照着这个来绑定形状就不会出错了。
内容的提问来源于stack exchange,提问作者Vasanti
相关产品推荐
相关产品推荐

