Pytorch 模型输入报错踩坑笔记

文章讲述了在使用PyTorch处理音频数据时,因不加索引地使用torch.squeeze()导致的错误。作者提醒在音频输入有额外声道维度时,应谨慎使用该函数,特别是在batch_size为1的情况下,可能导致维度丢失。解决方案是明确指定squeeze的维度,如audio_input.squeeze(1)。

前言

这一篇希望能长期记录一下写 pytorch 时踩到的各种的坑,有些时候偷个小懒一时半会儿不会遇到问题,但会在很久之后背刺(捂脸)。

CSDN 应该都是目的导向型的比较受欢迎,带着问题 google 搜到的文章。这种泛读的经验文章估计又吃灰了hhh。

正文

输入使用 squeeze()

谨慎使用 torch.squeeze,会将所有长度为 1 的维度都抹掉,所以注意补充参数。

下面的场景中,读入的 audio_input 的 shape 为 [batch_size, 1, audio_length] 。额外的维度1是音频加载库从文件中读取音频时自带的声道维度。但是使用的模型默认单声道,如果传入该维度会报错。于是使用下面的代码去除该维度。

audio_input = torch.squeeze(audio_input)

在很长的时间中没有出现任何问题。

直到有一次评估的时候,在最后一个样本的时候,他崩了。因为只有一个样本,所以输入的 shape 为 [1,1, audio_length] ,这导致前两个维度都被 squeeze 了。训练的时候没有遇到这个问题因为会 drop last 。所以建议加上 index 参数,确定 squeeze 第几个维度。例如切换为下面的写法。

audio_input = audio_input.squeeze(1)

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值