You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用hub.KerasLayer调用YAMNet构建模型报错求助

问题定位与解决建议

错误原因

你调用new_model.build(waveform)时传入了实际的波形numpy数组,但tf.keras.Model.build()需要的是输入形状的元组,不是具体数据。当传入数组时,Keras会误把数组里的0.0(float类型)当成维度值,而维度必须是整数或None,因此抛出TypeError。

补充说明:YAMNet固定要求输入是16kHz采样率的单通道音频,你的3秒(3*16000采样点)音频长度符合要求,但传递方式不对。

修正步骤

  • 修正build参数:传入输入形状元组,而非实际数据。单样本输入用(3*16000,),支持批量输入用(None, 3*16000)(None表示批量大小可变)。
  • 可选:添加全局平均池化:YAMNet输出的embeddings形状是[batch_size, num_frames, 1024],加一层GlobalAveragePooling1D()可以把时间维度平均,得到固定形状[batch_size, 1024],避免后续Dense层处理时的维度歧义。

修正后的代码:

import tensorflow as tf
import tensorflow_hub as hub
import numpy as np  # 补全之前遗漏的numpy导入

num_classes = 8
# 定义输入形状:(批量大小, 采样点数),None支持可变批量
input_shape = (None, 3 * 16000)
new_model = tf.keras.Sequential([
    hub.KerasLayer("https://tfhub.dev/google/yamnet/1", output_key='embeddings', trainable=False),
    # 显式池化时间维度,统一特征形状
    tf.keras.layers.GlobalAveragePooling1D(),
    tf.keras.layers.Dense(num_classes, activation='softmax')
])
new_model.build(input_shape)

new_model.summary()

额外提醒

  • 训练时要保证输入是**16kHz采样率、单通道、float32格式(值范围[-1,1])**的音频,Mini Speech Commands数据集刚好符合,预处理时注意格式转换即可。
  • 部署到Android前,要把训练好的模型转成TensorFlow Lite格式,确保和Android端的TFLite运行时兼容,同时确认输入输出形状符合移动端推理需求。

内容的提问来源于stack exchange,提问作者zheng

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 03:36:20