从TFEvent还原音频tf.Tensor时的解码错误与张量不一致问题
从TensorBoard TFEvent还原音频tf.Tensor的正确方法
问题本质
用clu的summary_writer把音频张量写到TensorBoard日志后,还原时踩了两个坑:一是tf.audio.decode_wav要求输入必须是标量,直接喂批量数据会报错;二是解码参数和写入时不匹配,导致还原后的张量和原数据差太多,过不了单元测试的断言。
解决步骤
1. 先明确clu写音频的规则
clu写音频时默认会这么处理:
- 把输入张量转成float32格式
- 按16kHz采样率、单声道编码成WAV字节流
- 把字节流存在TFEvent的
summary.value的bytes_list里
如果写入时自定义了采样率、声道数,还原时必须完全对应,不然肯定出问题。
2. 正确读取并解码TFEvent里的音频数据
第一步:提取TFEvent中的音频字节流
import tensorflow as tf from tensorboard.backend.event_processing.event_accumulator import EventAccumulator # 加载日志目录 event_acc = EventAccumulator("你的日志路径") event_acc.Reload() # 取音频标签对应的事件数据(假设标签是"audio") audio_tags = event_acc.Tags()["audio"] audio_events = event_acc.Audio(audio_tags[0]) # 拿到第一个事件里的音频字节流 audio_bytes = audio_events[0].audio_string
第二步:满足标量要求,精准解码
tf.audio.decode_wav只认标量字符串张量,所以先把字节流转成标量,同时要指定和写入时一致的声道数、采样率,别让tf自动推断:
# 转成标量字符串张量 audio_bytes_scalar = tf.convert_to_tensor(audio_bytes, dtype=tf.string) # 解码时指定声道数(clu默认单声道),不限制采样点数 waveform, sample_rate = tf.audio.decode_wav( audio_bytes_scalar, desired_channels=1, desired_samples=None ) # 转成和原张量一致的数据类型,调整形状匹配 restored_tensor = tf.cast(waveform, dtype=tf.float32) restored_tensor = tf.squeeze(restored_tensor) # 原张量是一维的话就挤掉多余维度
3. 匹配写入时的自定义参数
如果写入时改了采样率或声道数,比如这样写:
from clu import summary_writer writer = summary_writer.create_summary_writer("你的日志路径") with writer.as_default(): # 自定义44100采样率、双声道 summary_writer.audio( "audio", audio_tensor, sample_rate=44100, max_outputs=1 )
那还原时就得对应改解码参数:
waveform, sample_rate = tf.audio.decode_wav( audio_bytes_scalar, desired_channels=2, desired_samples=None )
4. 搞定断言不通过的问题
如果还是过不了np.testing.assert_allclose,检查这几点:
- 数据类型是否一致:原张量是float16的话,还原后也要转成float16,别用默认的float32
- 编码精度损失:WAV是16位整数编码,转成float32后会有微小误差,断言时给个合理的误差范围,比如
np.testing.assert_allclose(original.numpy(), restored.numpy(), rtol=1e-4, atol=1e-6) - 形状是否匹配:原张量如果是
[batch, length, channels],还原后要调整到对应形状,比如写入时取了batch里的第一个样本,还原后也要对应处理
完整单元测试代码
import tensorflow as tf import numpy as np from clu import summary_writer from tensorboard.backend.event_processing.event_accumulator import EventAccumulator import tempfile def test_audio_summary_roundtrip(): # 生成测试用的1秒16kHz单声道音频张量 original_audio = tf.random.normal([16000], dtype=tf.float32) # 用临时目录存日志 with tempfile.TemporaryDirectory() as tmpdir: # 写入日志 writer = summary_writer.create_summary_writer(tmpdir) with writer.as_default(): # 要加batch和channel维度,符合clu的输入要求 summary_writer.audio("test_audio", original_audio[tf.newaxis, :, tf.newaxis], sample_rate=16000) writer.flush() # 读取并还原 event_acc = EventAccumulator(tmpdir) event_acc.Reload() audio_events = event_acc.Audio("test_audio") audio_bytes = audio_events[0].audio_string # 解码还原 audio_bytes_scalar = tf.convert_to_tensor(audio_bytes, dtype=tf.string) waveform, _ = tf.audio.decode_wav(audio_bytes_scalar, desired_channels=1) restored_audio = tf.squeeze(tf.cast(waveform, tf.float32)) # 断言,允许微小精度误差 np.testing.assert_allclose(original_audio.numpy(), restored_audio.numpy(), rtol=1e-4, atol=1e-6) if __name__ == "__main__": test_audio_summary_roundtrip()
内容的提问来源于stack exchange,提问作者BioGeek
相关产品推荐
相关产品推荐

