TensorFlow.js中loadGraphModel模型如何实现类似summary()的输入输出打印?
获取TensorFlow.js GraphModel的输入输出信息
通过loadGraphModel加载的GraphModel基于计算图结构,没有LayersModel自带的summary()方法,但可以通过以下方式获取输入输出的核心信息:
- 直接访问模型的
inputs和outputs属性,这两个属性返回包含输入/输出张量元信息的数组,每个元素包含名称、形状、数据类型等关键内容。
代码示例:打印输入输出详情
// 加载GraphModel const model = await tf.loadGraphModel('path/to/your/model.json'); // 打印输入信息 console.log('=== 模型输入信息 ==='); model.inputs.forEach((input, index) => { console.log(`输入 ${index + 1}`); console.log(`名称: ${input.name}`); console.log(`形状: ${input.shape}`); console.log(`数据类型: ${input.dtype}\n`); }); // 打印输出信息 console.log('=== 模型输出信息 ==='); model.outputs.forEach((output, index) => { console.log(`输出 ${index + 1}`); console.log(`名称: ${output.name}`); console.log(`形状: ${output.shape}`); console.log(`数据类型: ${output.dtype}\n`); });
验证实际输出形状(可选)
如果需要确认模型运行时的输出形状,可以用零张量作为测试输入执行模型:
// 根据第一个输入的形状创建测试张量 const testInput = tf.zeros(model.inputs[0].shape); // 执行模型获取输出 const testOutput = model.execute(testInput); console.log('测试输出形状:', testOutput.shape); // 释放张量内存,避免内存泄漏 testInput.dispose(); testOutput.dispose();
若需要查看完整计算图结构,可使用model.getOperations()获取所有运算节点,但该方法会返回大量细节,一般调试输入输出用前面的方法即可满足需求。
内容的提问来源于stack exchange,提问作者windev92
相关产品推荐
相关产品推荐

