SongGen / songgen /xcodec_wrapper /modeling_xcodec.py
Vansh Chugh
initial deploy
ad73d02
Raw
History Blame Contribute Delete
6.64 kB
import torch
from transformers import PreTrainedModel
from transformers.models.encodec.modeling_encodec import EncodecDecoderOutput, EncodecEncoderOutput
from .configuration_xcodec import XCodecConfig
import torch.nn as nn
import sys
import os
XCODEC_INFER_PATH = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'xcodec_infer')
sys.path.append(XCODEC_INFER_PATH)
from omegaconf import OmegaConf
from models.soundstream_semantic import SoundStream
# model doesn't support batching yet
def build_codec_model(config):
model = eval(config.generator.name)(**config.generator.config)
return model
class XCodecModel(nn.Module):
config_class = XCodecConfig
main_input_name = "input_values"
def __init__(self):
super().__init__()
ckpt_path = os.path.join(XCODEC_INFER_PATH , 'ckpts/general_more/xcodec_hubert_general_audio_v2.pth')
config_path = os.path.join(XCODEC_INFER_PATH ,'ckpts/general_more/config_hubert_general.yaml')
config = OmegaConf.load(config_path)
self.model = build_codec_model(config)
parameter_dict = torch.load(ckpt_path)
self.model.load_state_dict(parameter_dict)
self.model.eval()
self.num_codebooks = 8
self.codebook_size = 1024
self.config = config
self.config.codebook_size = 1024
self.config.n_codebooks = 8
self.config.return_dict = True
self.config.frame_rate = 50
def encode(
self, input_values, padding_mask=None, bandwidth=None, return_dict=None, n_quantizers=None, sample_rate=None
):
"""
Encodes the input audio waveform into discrete codes.
Args:
input_values (`torch.Tensor` of shape `(batch_size, channels, sequence_length)`):
Float values of the input audio waveform.
padding_mask (`torch.Tensor` of shape `(batch_size, channels, sequence_length)`):
Padding mask used to pad the `input_values`.
bandwidth (`float`, *optional*):
Not used, kept to have the same inferface as HF encodec.
n_quantizers (`int`, *optional*) :
Number of quantizers to use, by default None
If None, all quantizers are used.
sample_rate (`int`, *optional*) :
Signal sampling_rate
Returns:
A list of frames containing the discrete encoded codes for the input audio waveform, along with rescaling
factors for each chunk when `normalize` is True. Each frames is a tuple `(codebook, scale)`, with
`codebook` of shape `[batch_size, num_codebooks, frames]`.
Scale is not used here.
"""
bsz, channels, input_length = input_values.shape
if channels < 1 or channels > 2:
raise ValueError(f"Number of audio channels must be 1 or 2, but got {channels}")
if bsz != 1:
raise ValueError(f"Number of audio batch_size must be 1 in XCodec Encode() : {input_values.shape}")
# audio_data = self.model.preprocess(input_values, sample_rate)
if channels != 1:
input_values = input_values.mean(1, keepdim=True)
if sample_rate is not None and sample_rate != 16000:
input_values = torchaudio.transforms.Resample(sample_rate, 16000)(input_values)
audio_data = input_values
return_dict = return_dict if return_dict is not None else self.config.return_dict
# TODO: for now, no chunk length
chunk_length = None # self.config.chunk_length
if chunk_length is None:
chunk_length = input_length
stride = input_length
else:
stride = self.config.chunk_stride
if padding_mask is None:
padding_mask = torch.ones_like(input_values).bool()
encoded_frames = []
scales = []
step = chunk_length - stride
if (input_length % stride) - step != 0:
raise ValueError(
"The input length is not properly padded for batched chunked decoding. Make sure to pad the input correctly."
)
for offset in range(0, input_length - step, stride):
mask = padding_mask[..., offset : offset + chunk_length].bool()
frame = audio_data[:, :, offset : offset + chunk_length] #[1,1, squeeze_len]
scale = None
encoded_frame= self.model.encode(frame, target_bw=4) # [8, 1, seq_len]
encoded_frame = encoded_frame.transpose(0 , 1) # [1, 8, seq_len]
encoded_frames.append(encoded_frame)
scales.append(scale)
encoded_frames = torch.stack(encoded_frames)
if not return_dict:
return (encoded_frames, scales)
return EncodecEncoderOutput(encoded_frames, scales)
def decode(
self,
audio_codes,
audio_scales=None,
padding_mask=None,
return_dict=None,
):
"""
Decodes the given frames into an output audio waveform.
Note that the output might be a bit bigger than the input. In that case, any extra steps at the end can be
trimmed.
Args:
audio_codes (`torch.FloatTensor` of shape `(batch_size, nb_chunks, chunk_length)`, *optional*):
Discret code embeddings computed using `model.encode`.
audio_scales (`torch.Tensor` of shape `(batch_size, nb_chunks)`, *optional*):
Not used, kept to have the same inferface as HF encodec.
padding_mask (`torch.Tensor` of shape `(batch_size, channels, sequence_length)`):
Padding mask used to pad the `input_values`.
Not used yet, kept to have the same inferface as HF encodec.
return_dict (`bool`, *optional*):
Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
"""
return_dict = return_dict or self.config.return_dict
# TODO: for now, no chunk length
# input shape [1, 1, 8, seq_len]
if len(audio_codes) != 1:
raise ValueError(f"Expected one frame, got {len(audio_codes)}")
audio_codes = audio_codes.transpose(1, 2) # [1, 8, 1, seq_len]
# audio_values = self.model.quantizer.from_codes(audio_codes.squeeze(0))[0]
audio_values = self.model.decode(audio_codes.squeeze(0))
if not return_dict:
return (audio_values,)
return EncodecDecoderOutput(audio_values)
def forward(self, tensor):
raise ValueError("`XCodecModel.forward` not implemented yet")