使用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
相关产品推荐
相关产品推荐

