49 lines
1.3 KiB
Python
49 lines
1.3 KiB
Python
import av
|
|
import torch
|
|
|
|
|
|
def load_audio(filepath: str) -> tuple[torch.Tensor, int]:
|
|
"""
|
|
从本地路径加载音频文件
|
|
|
|
Args:
|
|
filepath: 音频文件路径
|
|
|
|
Returns:
|
|
wav: 音频张量
|
|
sr: 采样率
|
|
"""
|
|
with av.open(filepath) as af:
|
|
if not af.streams.audio:
|
|
raise ValueError("No audio stream found in the file.")
|
|
|
|
stream = af.streams.audio[0]
|
|
sr = stream.codec_context.sample_rate
|
|
n_channels = stream.channels
|
|
|
|
frames = []
|
|
length = 0
|
|
for frame in af.decode(streams=stream.index):
|
|
buf = torch.from_numpy(frame.to_ndarray())
|
|
if buf.shape[0] != n_channels:
|
|
buf = buf.view(-1, n_channels).t()
|
|
|
|
frames.append(buf)
|
|
length += buf.shape[1]
|
|
|
|
if not frames:
|
|
raise ValueError("No audio frames decoded.")
|
|
|
|
wav = torch.cat(frames, dim=1)
|
|
wav = f32_pcm(wav)
|
|
return wav, sr
|
|
|
|
def f32_pcm(wav: torch.Tensor) -> torch.Tensor:
|
|
"""Convert audio to float 32 bits PCM format."""
|
|
if wav.dtype.is_floating_point:
|
|
return wav
|
|
elif wav.dtype == torch.int16:
|
|
return wav.float() / (2 ** 15)
|
|
elif wav.dtype == torch.int32:
|
|
return wav.float() / (2 ** 31)
|
|
raise ValueError(f"Unsupported wav dtype: {wav.dtype}") |