已训练转换的TensorFlowJS自定义模型运行时出现concat4D维度不匹配错误
错误触发原因
这个报错来自TensorFlowJS的4D张量拼接层(concat4D),4D张量默认采用[批量大小, 高度, 宽度, 通道数]的NHWC格式,拼接操作要求除拼接轴外其余所有维度的形状完全一致。
Error in concat4D: Shape of tensors[1] (1,30,40,256) does not match the shape of the rest (1,15,20,832) along the non-concatenated axis 1.
从报错信息可以看出,当前拼接轴为最后一维(通道数),但两个待拼接张量的第1维(高度)分别为30和15、第2维(宽度)分别为40和20,完全不匹配,因此触发报错。
根本原因是你输入模型的图像尺寸和模型训练阶段约定的输入尺寸不匹配,导致多分支结构(常见于Unet、YOLO、残差拼接类网络)的不同分支下采样后输出的特征图尺寸无法对齐,无法完成拼接操作。
修复方案
- 首先获取模型训练时的标准输入参数:可通过
console.log(model.inputs[0].shape)打印加载后模型的预期输入形状,格式为[批量大小, 标准高度, 标准宽度, 输入通道数],比如常见的输出为[null, 120, 160, 3],即要求输入图像高120、宽160、3通道RGB格式。 - 对输入图像做标准化预处理,不要直接传入
imageToRgbaMatrix的输出结果,参考修改后的代码:
const RGB = await imageToRgbaMatrix(imageUrl); const model = await tf.loadLayersModel(modelUrl); // 1. 转张量+去除多余的alpha通道(如果imageToRgbaMatrix返回4通道RGBA格式) let imgTensor = tf.tensor(RGB).slice([0,0,0], [-1,-1,3]); // 2. 缩放到模型要求的标准尺寸,这里的[120, 160]替换为你打印出来的标准高度、宽度 imgTensor = tf.image.resizeBilinear(imgTensor, [120, 160]); // 3. 和训练阶段对齐做归一化,比如训练时做了/255归一化到[0,1]区间,此处保持一致 imgTensor = imgTensor.div(tf.scalar(255)); // 4. 扩展batch维度 const ImageData = imgTensor.expandDims(0); const predictions = model.predict(ImageData);
- 如果打印模型输入形状出现动态维度(比如高度宽度为null),说明转换TFJS模型时未固定输入形状,重新使用tfjs-converter转换时添加参数
--input_shape [1, 标准高度, 标准宽度, 3]即可。
内容的提问来源于stack exchange,提问作者Dev
相关产品推荐
相关产品推荐

