diff --git a/inference/xcodec_mini_infer/RepCodec/examples/data2vec_audio.py b/inference/xcodec_mini_infer/RepCodec/examples/data2vec_audio.py new file mode 100644 index 0000000..fb5b016 --- /dev/null +++ b/inference/xcodec_mini_infer/RepCodec/examples/data2vec_audio.py @@ -0,0 +1,541 @@ +# Copyright (c) ByteDance, Inc. and its affiliates. +# Copyright (c) Chutong Meng +# +# This source code is licensed under the MIT license found in the +# LICENSE file in the root directory of this source tree. +# Based on fairseq (https://github.com/facebookresearch/fairseq) + +# ref: https://github.com/facebookresearch/fairseq/blob/main/examples/data2vec/models/data2vec_audio.py + +import logging +import math +from dataclasses import dataclass, field +from typing import Optional + +from omegaconf import II + +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.distributed as dist + +from fairseq.modules import EMAModule, EMAModuleConfig +from fairseq.data.data_utils import compute_mask_indices +from fairseq.models import BaseFairseqModel, register_model +from fairseq.models.wav2vec import ( + ConvFeatureExtractionModel, + Wav2Vec2Config, + TransformerEncoder, +) +from fairseq.modules import ( + GradMultiply, + LayerNorm, +) +from fairseq.utils import index_put + + +logger = logging.getLogger(__name__) + + +@dataclass +class Data2VecAudioConfig(Wav2Vec2Config): + + loss_beta: float = field( + default=0, metadata={"help": "beta for smooth l1 loss. 0 means use l2 loss"} + ) + loss_scale: Optional[float] = field( + default=None, + metadata={ + "help": "scale the reconstruction loss by this constant. if None then scales by 1/sqrt(dim)" + }, + ) + average_top_k_layers: int = field( + default=8, metadata={"help": "how many layers to average"} + ) + + layer_norm_target_layer: bool = False + instance_norm_target_layer: bool = False + instance_norm_targets: bool = False + layer_norm_targets: bool = False + batch_norm_target_layer: bool = False + group_norm_target_layer: bool = False + + ema_decay: float = field(default=0.999, metadata={"help": "initial ema decay rate"}) + ema_end_decay: float = field( + default=0.9999, metadata={"help": "final ema decay rate"} + ) + + # when to finish annealing ema decay rate + ema_anneal_end_step: int = II("optimization.max_update") + + ema_transformer_only: bool = field( + default=True, + metadata={"help": "whether to momentum update only the transformer"}, + ) + ema_layers_only: bool = field( + default=True, + metadata={"help": "whether to momentum update only the transformer layers"}, + ) + + max_update: int = II("optimization.max_update") + + min_target_var: float = field( + default=0.1, metadata={"help": "stop training if target var falls below this"} + ) + min_pred_var: float = field( + default=0.01, + metadata={"help": "stop training if prediction var falls below this"}, + ) + + +def get_annealed_rate(start, end, curr_step, total_steps): + r = end - start + pct_remaining = 1 - curr_step / total_steps + return end - r * pct_remaining + + +@register_model("data2vec_audio", dataclass=Data2VecAudioConfig) +class Data2VecAudioModel(BaseFairseqModel): + def __init__(self, cfg: Data2VecAudioConfig): + super().__init__() + self.cfg = cfg + + feature_enc_layers = eval(cfg.conv_feature_layers) + self.extractor_embed = feature_enc_layers[-1][0] + + self.ema = None + self.embed = cfg.encoder_embed_dim + + self.average_top_k_layers = cfg.average_top_k_layers + self.loss_beta = cfg.loss_beta + self.loss_scale = cfg.loss_scale + + self.feature_extractor = ConvFeatureExtractionModel( + conv_layers=feature_enc_layers, + dropout=0.0, + mode=cfg.extractor_mode, + conv_bias=cfg.conv_bias, + ) + + self.post_extract_proj = nn.Linear(self.extractor_embed, cfg.encoder_embed_dim) + + self.mask_prob = cfg.mask_prob + self.mask_selection = cfg.mask_selection + self.mask_other = cfg.mask_other + self.mask_length = cfg.mask_length + self.no_mask_overlap = cfg.no_mask_overlap + self.mask_min_space = cfg.mask_min_space + + self.mask_channel_prob = cfg.mask_channel_prob + self.mask_channel_before = cfg.mask_channel_before + self.mask_channel_selection = cfg.mask_channel_selection + self.mask_channel_other = cfg.mask_channel_other + self.mask_channel_length = cfg.mask_channel_length + self.no_mask_channel_overlap = cfg.no_mask_channel_overlap + self.mask_channel_min_space = cfg.mask_channel_min_space + + self.dropout_input = nn.Dropout(cfg.dropout_input) + self.dropout_features = nn.Dropout(cfg.dropout_features) + + self.feature_grad_mult = cfg.feature_grad_mult + + self.mask_emb = nn.Parameter( + torch.FloatTensor(cfg.encoder_embed_dim).uniform_() + ) + + self.encoder = TransformerEncoder(cfg) + self.layer_norm = LayerNorm(self.extractor_embed) + + self.final_proj = nn.Linear(self.embed, self.embed) + + self.num_updates = 0 + + def make_ema_teacher(self): + ema_config = EMAModuleConfig( + ema_decay=self.cfg.ema_decay, + ema_fp32=True, + ) + skip_keys = set() + if self.cfg.ema_layers_only: + self.cfg.ema_transformer_only = True + for k, _ in self.encoder.pos_conv.named_parameters(): + skip_keys.add(f"pos_conv.{k}") + + self.ema = EMAModule( + self.encoder if self.cfg.ema_transformer_only else self, + ema_config, + skip_keys=skip_keys, + ) + + def set_num_updates(self, num_updates): + super().set_num_updates(num_updates) + + if self.ema is None and self.final_proj is not None: + logger.info(f"making ema teacher") + self.make_ema_teacher() + elif self.training and self.ema is not None: + if self.cfg.ema_decay != self.cfg.ema_end_decay: + if num_updates >= self.cfg.ema_anneal_end_step: + decay = self.cfg.ema_end_decay + else: + decay = get_annealed_rate( + self.cfg.ema_decay, + self.cfg.ema_end_decay, + num_updates, + self.cfg.ema_anneal_end_step, + ) + self.ema.set_decay(decay) + if self.ema.get_decay() < 1: + self.ema.step(self.encoder if self.cfg.ema_transformer_only else self) + + self.num_updates = num_updates + + def state_dict(self, destination=None, prefix="", keep_vars=False): + state = super().state_dict(destination, prefix, keep_vars) + + if self.ema is not None: + state[prefix + "_ema"] = self.ema.fp32_params + + return state + + def _load_from_state_dict(self, state_dict, prefix, *args, **kwargs): + if self.ema is not None: + k = prefix + "_ema" + assert k in state_dict + self.ema.restore(state_dict[k], True) + del state_dict[k] + return super()._load_from_state_dict(state_dict, prefix, *args, **kwargs) + + @classmethod + def build_model(cls, cfg: Data2VecAudioConfig, task=None): + """Build a new model instance.""" + + return cls(cfg) + + def apply_mask( + self, + x, + padding_mask, + mask_indices=None, + mask_channel_indices=None, + ): + B, T, C = x.shape + + if self.mask_channel_prob > 0 and self.mask_channel_before: + mask_channel_indices = compute_mask_indices( + (B, C), + None, + self.mask_channel_prob, + self.mask_channel_length, + self.mask_channel_selection, + self.mask_channel_other, + no_overlap=self.no_mask_channel_overlap, + min_space=self.mask_channel_min_space, + ) + mask_channel_indices = ( + torch.from_numpy(mask_channel_indices) + .to(x.device) + .unsqueeze(1) + .expand(-1, T, -1) + ) + x[mask_channel_indices] = 0 + + if self.mask_prob > 0: + if mask_indices is None: + mask_indices = compute_mask_indices( + (B, T), + padding_mask, + self.mask_prob, + self.mask_length, + self.mask_selection, + self.mask_other, + min_masks=1, + no_overlap=self.no_mask_overlap, + min_space=self.mask_min_space, + require_same_masks=self.cfg.require_same_masks, + mask_dropout=self.cfg.mask_dropout, + ) + mask_indices = torch.from_numpy(mask_indices).to(x.device) + x = index_put(x, mask_indices, self.mask_emb) + else: + mask_indices = None + + if self.mask_channel_prob > 0 and not self.mask_channel_before: + if mask_channel_indices is None: + mask_channel_indices = compute_mask_indices( + (B, C), + None, + self.mask_channel_prob, + self.mask_channel_length, + self.mask_channel_selection, + self.mask_channel_other, + no_overlap=self.no_mask_channel_overlap, + min_space=self.mask_channel_min_space, + ) + mask_channel_indices = ( + torch.from_numpy(mask_channel_indices) + .to(x.device) + .unsqueeze(1) + .expand(-1, T, -1) + ) + x = index_put(x, mask_channel_indices, 0) + + return x, mask_indices + + def _get_feat_extract_output_lengths(self, input_lengths: torch.LongTensor): + """ + Computes the output length of the convolutional layers + """ + + def _conv_out_length(input_length, kernel_size, stride): + return torch.floor((input_length - kernel_size) / stride + 1) + + conv_cfg_list = eval(self.cfg.conv_feature_layers) + + for i in range(len(conv_cfg_list)): + input_lengths = _conv_out_length( + input_lengths, conv_cfg_list[i][1], conv_cfg_list[i][2] + ) + + return input_lengths.to(torch.long) + + def forward( + self, + source, + padding_mask=None, + mask=True, + features_only=False, + layer=None, + mask_indices=None, + mask_channel_indices=None, + padding_count=None, + ): + features = source + + if self.feature_grad_mult > 0: + features = self.feature_extractor(features) + if self.feature_grad_mult != 1.0: + features = GradMultiply.apply(features, self.feature_grad_mult) + else: + with torch.no_grad(): + features = self.feature_extractor(features) + + features = features.transpose(1, 2) + + features = self.layer_norm(features) + + orig_padding_mask = padding_mask + + if padding_mask is not None and padding_mask.any(): + input_lengths = (1 - padding_mask.long()).sum(-1) + # apply conv formula to get real output_lengths + output_lengths = self._get_feat_extract_output_lengths(input_lengths) + + padding_mask = torch.zeros( + features.shape[:2], dtype=features.dtype, device=features.device + ) + + # these two operations makes sure that all values + # before the output lengths indices are attended to + padding_mask[ + ( + torch.arange(padding_mask.shape[0], device=padding_mask.device), + output_lengths - 1, + ) + ] = 1 + padding_mask = (1 - padding_mask.flip([-1]).cumsum(-1).flip([-1])).bool() + else: + padding_mask = None + + if self.post_extract_proj is not None: + features = self.post_extract_proj(features) + + pre_encoder_features = None + if self.cfg.ema_transformer_only: + pre_encoder_features = features.clone() + + features = self.dropout_input(features) + + if mask: + x, mask_indices = self.apply_mask( + features, + padding_mask, + mask_indices=mask_indices, + mask_channel_indices=mask_channel_indices, + ) + else: + x = features + mask_indices = None + + x, layer_results = self.encoder( + x, + padding_mask=padding_mask, + layer=layer, + ) + + if features_only: + return { + "x": x, + "padding_mask": padding_mask, + "layer_results": layer_results, + } + + result = { + "losses": {}, + } + + with torch.no_grad(): + self.ema.model.eval() + + if self.cfg.ema_transformer_only: + y, layer_results = self.ema.model.extract_features( + pre_encoder_features, + padding_mask=padding_mask, + min_layer=self.cfg.encoder_layers - self.average_top_k_layers, + ) + y = { + "x": y, + "padding_mask": padding_mask, + "layer_results": layer_results, + } + else: + y = self.ema.model.extract_features( + source=source, + padding_mask=orig_padding_mask, + mask=False, + ) + + target_layer_results = [l[2] for l in y["layer_results"]] + + permuted = False + if self.cfg.instance_norm_target_layer or self.cfg.batch_norm_target_layer: + target_layer_results = [ + tl.permute(1, 2, 0) for tl in target_layer_results # TBC -> BCT + ] + permuted = True + + if self.cfg.batch_norm_target_layer: + target_layer_results = [ + F.batch_norm( + tl.float(), running_mean=None, running_var=None, training=True + ) + for tl in target_layer_results + ] + + if self.cfg.instance_norm_target_layer: + target_layer_results = [ + F.instance_norm(tl.float()) for tl in target_layer_results + ] + + if permuted: + target_layer_results = [ + tl.transpose(1, 2) for tl in target_layer_results # BCT -> BTC + ] + + if self.cfg.group_norm_target_layer: + target_layer_results = [ + F.layer_norm(tl.float(), tl.shape[-2:]) + for tl in target_layer_results + ] + + if self.cfg.layer_norm_target_layer: + target_layer_results = [ + F.layer_norm(tl.float(), tl.shape[-1:]) + for tl in target_layer_results + ] + + y = sum(target_layer_results) / len(target_layer_results) + + if self.cfg.layer_norm_targets: + y = F.layer_norm(y.float(), y.shape[-1:]) + + if self.cfg.instance_norm_targets: + y = F.instance_norm(y.float().transpose(1, 2)).transpose(1, 2) + + if not permuted: + y = y.transpose(0, 1) + + y = y[mask_indices] + + x = x[mask_indices] + x = self.final_proj(x) + + sz = x.size(-1) + + if self.loss_beta == 0: + loss = F.mse_loss(x.float(), y.float(), reduction="none").sum(dim=-1) + else: + loss = F.smooth_l1_loss( + x.float(), y.float(), reduction="none", beta=self.loss_beta + ).sum(dim=-1) + + if self.loss_scale is not None: + scale = self.loss_scale + else: + scale = 1 / math.sqrt(sz) + + result["losses"]["regression"] = loss.sum() * scale + + if "sample_size" not in result: + result["sample_size"] = loss.numel() + + with torch.no_grad(): + result["target_var"] = self.compute_var(y) + result["pred_var"] = self.compute_var(x.float()) + + if self.num_updates > 5000 and result["target_var"] < self.cfg.min_target_var: + logger.error( + f"target var is {result['target_var'].item()} < {self.cfg.min_target_var}, exiting" + ) + raise Exception( + f"target var is {result['target_var'].item()} < {self.cfg.min_target_var}, exiting" + ) + if self.num_updates > 5000 and result["pred_var"] < self.cfg.min_pred_var: + logger.error( + f"pred var is {result['pred_var'].item()} < {self.cfg.min_pred_var}, exiting" + ) + raise Exception( + f"pred var is {result['pred_var'].item()} < {self.cfg.min_pred_var}, exiting" + ) + + if self.ema is not None: + result["ema_decay"] = self.ema.get_decay() * 1000 + + return result + + @staticmethod + def compute_var(y): + y = y.view(-1, y.size(-1)) + if dist.is_initialized(): + zc = torch.tensor(y.size(0)).cuda() + zs = y.sum(dim=0) + zss = (y ** 2).sum(dim=0) + + dist.all_reduce(zc) + dist.all_reduce(zs) + dist.all_reduce(zss) + + var = zss / (zc - 1) - (zs ** 2) / (zc * (zc - 1)) + return torch.sqrt(var + 1e-6).mean() + else: + return torch.sqrt(y.var(dim=0) + 1e-6).mean() + + def extract_features( + self, source, padding_mask, mask=False, layer=None + ): + res = self.forward( + source, + padding_mask, + mask=mask, + features_only=True, + layer=layer, + ) + return res + + def remove_pretraining_modules(self, last_layer=None): + self.final_proj = None + self.ema = None + if last_layer is not None: + self.encoder.layers = nn.ModuleList( + l for i, l in enumerate(self.encoder.layers) if i <= last_layer + ) diff --git a/inference/xcodec_mini_infer/RepCodec/examples/data2vec_feature_reader.py b/inference/xcodec_mini_infer/RepCodec/examples/data2vec_feature_reader.py new file mode 100644 index 0000000..dfa4234 --- /dev/null +++ b/inference/xcodec_mini_infer/RepCodec/examples/data2vec_feature_reader.py @@ -0,0 +1,87 @@ +# Copyright (c) ByteDance, Inc. and its affiliates. +# Copyright (c) Chutong Meng +# +# This source code is licensed under the MIT license found in the +# LICENSE file in the root directory of this source tree. +# Based on fairseq (https://github.com/facebookresearch/fairseq) + +import logging + +import torch +import torch.nn.functional as F +from fairseq import tasks +from fairseq.checkpoint_utils import load_checkpoint_to_cpu +from fairseq.data.audio.audio_utils import get_features_or_waveform +from omegaconf import OmegaConf + +from data2vec_audio import Data2VecAudioModel + +logger = logging.getLogger("dump_feature") + + +class Data2vecFeatureReader(object): + def __init__(self, ckpt_path: str, layer: int, device: str, max_chunk=1600000): + state = load_checkpoint_to_cpu(ckpt_path) + cfg = state["cfg"] + # load task + task = tasks.setup_task(cfg.task, from_checkpoint=True) + task.load_state_dict(state["task_state"]) + # load model config + if "layer_type" not in cfg.model: + # fix a missing key + model_config = {k: v for k, v in cfg.model.items()} + model_config["layer_type"] = "transformer" + model_config = OmegaConf.create(model_config) + else: + model_config = cfg.model + + # fix param name in the state + state["model"]["final_proj.weight"] = state["model"].pop("final_proj.0.weight") + state["model"]["final_proj.bias"] = state["model"].pop("final_proj.0.bias") + del state["model"]["_ema"] + + # load model + model = Data2VecAudioModel.build_model(model_config) + model.load_state_dict( + state["model"], strict=True, model_cfg=model_config + ) + + self.device = device + logger.info(f"device = {self.device}") + + self.model = model.eval().to(self.device) + self.task = task + self.layer = layer - 1 # make it 1-based + self.max_chunk = max_chunk + logger.info(f"TASK CONFIG:\n{self.task.cfg}") + logger.info(f" max_chunk = {self.max_chunk}") + + def read_audio(self, path, ref_len=None): + wav = get_features_or_waveform(path, need_waveform=True, use_sample_rate=self.task.cfg.sample_rate) + if wav.ndim == 2: + wav = wav.mean(-1) + assert wav.ndim == 1, wav.ndim + if ref_len is not None and abs(ref_len - len(wav)) > 160: + logger.warning(f"ref {ref_len} != read {len(wav)} ({path})") + return wav + + def get_feats(self, path, ref_len=None): + x = self.read_audio(path, ref_len=ref_len) + with torch.no_grad(): + x = torch.from_numpy(x).float().to(self.device) + if self.task.cfg.normalize: + x = F.layer_norm(x, x.shape) + x = x.view(1, -1) + + feat = [] + for start in range(0, x.size(1), self.max_chunk): + x_chunk = x[:, start: start + self.max_chunk] + res = self.model.extract_features( + source=x_chunk, + padding_mask=None, + mask=False, + layer=self.layer, + ) + feat_chunk = res["x"] + feat.append(feat_chunk) + return torch.cat(feat, 1).squeeze(0) diff --git a/inference/xcodec_mini_infer/RepCodec/examples/dump_feature.py b/inference/xcodec_mini_infer/RepCodec/examples/dump_feature.py new file mode 100644 index 0000000..cfecae5 --- /dev/null +++ b/inference/xcodec_mini_infer/RepCodec/examples/dump_feature.py @@ -0,0 +1,142 @@ +# Copyright (c) ByteDance, Inc. and its affiliates. +# Copyright (c) Chutong Meng +# +# This source code is licensed under the MIT license found in the +# LICENSE file in the root directory of this source tree. +# Based on fairseq (https://github.com/facebookresearch/fairseq) + +import logging +import os +import sys + +from feature_utils import get_path_iterator, dump_feature + +logging.basicConfig( + format="%(asctime)s | %(levelname)s | %(name)s | %(message)s", + datefmt="%Y-%m-%d %H:%M:%S", + level=os.environ.get("LOGLEVEL", "INFO").upper(), + stream=sys.stdout, +) +logger = logging.getLogger("dump_feature") + + +def main( + model_type: str, + tsv_path: str, + ckpt_path: str, + whisper_root: str, + whisper_name: str, + layer: int, + nshard: int, + rank: int, + feat_dir: str, + max_chunk: int, + use_cpu: bool = False +): + device = "cpu" if use_cpu else "cuda" + + # some checks + if model_type in ["hubert", "data2vec"]: + assert ckpt_path and os.path.exists(ckpt_path) + elif model_type in ["whisper"]: + assert whisper_name and whisper_root + else: + raise ValueError(f"Unsupported model type {model_type}") + + reader = None + if model_type == "hubert": + from hubert_feature_reader import HubertFeatureReader + reader = HubertFeatureReader(ckpt_path, layer, device=device, max_chunk=max_chunk) + elif model_type == "data2vec": + from data2vec_feature_reader import Data2vecFeatureReader + reader = Data2vecFeatureReader(ckpt_path, layer, device=device, max_chunk=max_chunk) + elif model_type == "whisper": + from whisper_feature_reader import WhisperFeatureReader + reader = WhisperFeatureReader(whisper_root, whisper_name, layer, device=device) + + assert reader is not None + + generator, num = get_path_iterator(tsv_path, nshard, rank) + dump_feature(reader, generator, num, nshard, rank, feat_dir) + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser() + parser.add_argument( + "--model_type", + required=True, + type=str, + choices=["data2vec", "hubert", "whisper"], + help="the type of the speech encoder." + ) + parser.add_argument( + "--tsv_path", + required=True, + type=str, + help="the path to the tsv file." + ) + parser.add_argument( + "--ckpt_path", + required=False, + type=str, + default=None, + help="path to the speech model. must provide for HuBERT and data2vec" + ) + parser.add_argument( + "--whisper_root", + required=False, + type=str, + default=None, + help="root dir to download/store whisper model. must provide for whisper model." + ) + parser.add_argument( + "--whisper_name", + required=False, + type=str, + default=None, + help="name of whisper model. e.g., large-v2. must provide for whisper model." + ) + parser.add_argument( + "--layer", + required=True, + type=int, + help="which layer of the model. this is 1-based." + ) + parser.add_argument( + "--feat_dir", + required=True, + type=str, + help="the output dir to save the representations." + ) + parser.add_argument( + "--nshard", + required=False, + type=int, + default=1, + help="total number of shards." + ) + parser.add_argument( + "--rank", + required=False, + type=int, + default=0, + help="shard id of this process." + ) + parser.add_argument( + "--max_chunk", + type=int, + default=1600000, + help="max number of frames of each batch." + ) + parser.add_argument( + "--use_cpu", + default=False, + action="store_true", + help="whether use cpu instead of gpu." + ) + args = parser.parse_args() + logger.info(args) + + main(**vars(args)) diff --git a/inference/xcodec_mini_infer/RepCodec/examples/feature_utils.py b/inference/xcodec_mini_infer/RepCodec/examples/feature_utils.py new file mode 100644 index 0000000..d12dd8c --- /dev/null +++ b/inference/xcodec_mini_infer/RepCodec/examples/feature_utils.py @@ -0,0 +1,70 @@ +# Copyright (c) ByteDance, Inc. and its affiliates. +# Copyright (c) Chutong Meng +# +# This source code is licensed under the MIT license found in the +# LICENSE file in the root directory of this source tree. +# Based on fairseq (https://github.com/facebookresearch/fairseq) + +# ref: https://github.com/facebookresearch/fairseq/blob/main/examples/hubert/simple_kmeans/feature_utils.py + +import logging +import os +import sys + +import tqdm +from npy_append_array import NpyAppendArray + + +logging.basicConfig( + format="%(asctime)s | %(levelname)s | %(name)s | %(message)s", + datefmt="%Y-%m-%d %H:%M:%S", + level=os.environ.get("LOGLEVEL", "INFO").upper(), + stream=sys.stdout, +) +logger = logging.getLogger("feature_utils") + + +def get_shard_range(tot, nshard, rank): + assert rank < nshard and rank >= 0, f"invaid rank/nshard {rank}/{nshard}" + start = round(tot / nshard * rank) + end = round(tot / nshard * (rank + 1)) + assert start < end, f"start={start}, end={end}" + logger.info( + f"rank {rank} of {nshard}, process {end-start} " + f"({start}-{end}) out of {tot}" + ) + return start, end + + +def get_path_iterator(tsv, nshard, rank): + with open(tsv, "r") as f: + root = f.readline().rstrip() + lines = [line.rstrip() for line in f] + start, end = get_shard_range(len(lines), nshard, rank) + lines = lines[start:end] + def iterate(): + for line in lines: + subpath, nsample = line.split("\t") + yield f"{root}/{subpath}", int(nsample) + return iterate, len(lines) + + +def dump_feature(reader, generator, num, nshard, rank, feat_dir): + iterator = generator() + + feat_path = f"{feat_dir}/{rank}_{nshard}.npy" + leng_path = f"{feat_dir}/{rank}_{nshard}.len" + + os.makedirs(feat_dir, exist_ok=True) + if os.path.exists(feat_path): + os.remove(feat_path) + + feat_f = NpyAppendArray(feat_path) + with open(leng_path, "w") as leng_f: + for path, nsample in tqdm.tqdm(iterator, total=num): + feat = reader.get_feats(path, nsample) + feat_f.append(feat.cpu().numpy()) + leng_f.write(f"{len(feat)}\n") + logger.info("finished successfully") + + diff --git a/inference/xcodec_mini_infer/RepCodec/examples/hubert_feature_reader.py b/inference/xcodec_mini_infer/RepCodec/examples/hubert_feature_reader.py new file mode 100644 index 0000000..3535851 --- /dev/null +++ b/inference/xcodec_mini_infer/RepCodec/examples/hubert_feature_reader.py @@ -0,0 +1,64 @@ +# Copyright (c) ByteDance, Inc. and its affiliates. +# Copyright (c) Chutong Meng +# +# This source code is licensed under the MIT license found in the +# LICENSE file in the root directory of this source tree. +# Based on fairseq (https://github.com/facebookresearch/fairseq) + +import logging + +import fairseq +import torch +import torch.nn.functional as F + +from fairseq.data.audio.audio_utils import get_features_or_waveform + +logger = logging.getLogger("dump_feature") + + +class HubertFeatureReader(object): + def __init__(self, ckpt_path: str, layer: int, device: str, max_chunk=1600000): + ( + model, + cfg, + task, + ) = fairseq.checkpoint_utils.load_model_ensemble_and_task([ckpt_path]) + + self.device = device + logger.info(f"device = {self.device}") + + self.model = model[0].eval().to(self.device) + self.task = task + self.layer = layer + self.max_chunk = max_chunk + logger.info(f"TASK CONFIG:\n{self.task.cfg}") + logger.info(f" max_chunk = {self.max_chunk}") + + def read_audio(self, path, ref_len=None): + wav = get_features_or_waveform(path, need_waveform=True, use_sample_rate=self.task.cfg.sample_rate) + if wav.ndim == 2: + wav = wav.mean(-1) + assert wav.ndim == 1, wav.ndim + if ref_len is not None and abs(ref_len - len(wav)) > 160: + logger.warning(f"ref {ref_len} != read {len(wav)} ({path})") + return wav + + def get_feats(self, path, ref_len=None): + x = self.read_audio(path, ref_len=ref_len) + with torch.no_grad(): + x = torch.from_numpy(x).float().to(self.device) + if self.task.cfg.normalize: + x = F.layer_norm(x, x.shape) + x = x.view(1, -1) + + feat = [] + for start in range(0, x.size(1), self.max_chunk): + x_chunk = x[:, start: start + self.max_chunk] + feat_chunk, _ = self.model.extract_features( + source=x_chunk, + padding_mask=None, + mask=False, + output_layer=self.layer, + ) + feat.append(feat_chunk) + return torch.cat(feat, 1).squeeze(0) diff --git a/inference/xcodec_mini_infer/RepCodec/examples/whisper_feature_reader.py b/inference/xcodec_mini_infer/RepCodec/examples/whisper_feature_reader.py new file mode 100644 index 0000000..d275db4 --- /dev/null +++ b/inference/xcodec_mini_infer/RepCodec/examples/whisper_feature_reader.py @@ -0,0 +1,110 @@ +# Copyright (c) ByteDance, Inc. and its affiliates. +# Copyright (c) Chutong Meng +# +# This source code is licensed under the MIT license found in the +# LICENSE file in the root directory of this source tree. +# Based on fairseq (https://github.com/facebookresearch/fairseq) and +# Whisper (https://github.com/openai/whisper/) + +import io +import logging +import os +from typing import Optional, Union + +import soundfile as sf +import torch +from whisper import _MODELS, _download, _ALIGNMENT_HEADS, available_models +from whisper.audio import log_mel_spectrogram +from whisper.model import ModelDimensions + +from whisper_model import Whisper_ + +logger = logging.getLogger("dump_feature") + + +def load_model( + name: str, + device: Optional[Union[str, torch.device]] = None, + download_root: str = None, + in_memory: bool = False, +) -> Whisper_: + """ + Reference: https://github.com/openai/whisper/blob/main/whisper/__init__.py#L97 + But we will load a `Whisper_` model for feature extraction. + + Parameters + ---------- + name : str + one of the official model names listed by `whisper.available_models()`, or + path to a model checkpoint containing the model dimensions and the model state_dict. + device : Union[str, torch.device] + the PyTorch device to put the model into + download_root: str + path to download the model files; by default, it uses "~/.cache/whisper" + in_memory: bool + whether to preload the model weights into host memory + + Returns + ------- + model : Whisper + The Whisper ASR model instance + """ + + if device is None: + device = "cuda" if torch.cuda.is_available() else "cpu" + if download_root is None: + default = os.path.join(os.path.expanduser("~"), ".cache") + download_root = os.path.join(os.getenv("XDG_CACHE_HOME", default), "whisper") + + if name in _MODELS: + checkpoint_file = _download(_MODELS[name], download_root, in_memory) + alignment_heads = _ALIGNMENT_HEADS[name] + elif os.path.isfile(name): + checkpoint_file = open(name, "rb").read() if in_memory else name + alignment_heads = None + else: + raise RuntimeError( + f"Model {name} not found; available models = {available_models()}" + ) + + with ( + io.BytesIO(checkpoint_file) if in_memory else open(checkpoint_file, "rb") + ) as fp: + checkpoint = torch.load(fp, map_location=device) + del checkpoint_file + + dims = ModelDimensions(**checkpoint["dims"]) + model = Whisper_(dims) + model.load_state_dict(checkpoint["model_state_dict"]) + + if alignment_heads is not None: + model.set_alignment_heads(alignment_heads) + + return model.to(device) + + +class WhisperFeatureReader(object): + def __init__(self, root, ckpt, layer, device): + self.device = device + logger.info(f"device = {self.device}") + + self.model: Whisper_ = load_model(name=ckpt, device=self.device, download_root=root).eval() + self.model.decoder = None # to save some memory by deleting the decoder + self.layer = layer # one-based + + def read_audio(self, path, ref_len=None): + wav, sample_rate = sf.read(path) + assert sample_rate == 16000, sample_rate + if ref_len is not None and abs(ref_len - len(wav)) > 160: + logger.warning(f"ref {ref_len} != read {len(wav)} ({path})") + return wav + + def get_feats(self, path, ref_len=None): + wav = self.read_audio(path, ref_len) + audio_length = len(wav) + with torch.no_grad(): + mel = log_mel_spectrogram(torch.from_numpy(wav).float().to(self.device)) + hidden = self.model.extract_features(mel.unsqueeze(0), target_layer=self.layer) + feature_length = audio_length // 320 + hidden = hidden[0, :feature_length] + return hidden.contiguous() diff --git a/inference/xcodec_mini_infer/RepCodec/examples/whisper_model.py b/inference/xcodec_mini_infer/RepCodec/examples/whisper_model.py new file mode 100644 index 0000000..04c15b5 --- /dev/null +++ b/inference/xcodec_mini_infer/RepCodec/examples/whisper_model.py @@ -0,0 +1,58 @@ +# Copyright (c) ByteDance, Inc. and its affiliates. +# Copyright (c) Chutong Meng +# +# This source code is licensed under the MIT license found in the +# LICENSE file in the root directory of this source tree. +# Based on fairseq (https://github.com/facebookresearch/fairseq) and +# Whisper (https://github.com/openai/whisper/) + +from typing import Optional + +import torch +import torch.nn.functional as F +from torch import Tensor +from whisper.model import AudioEncoder, sinusoids, Whisper, ModelDimensions + + +class AudioEncoder_(AudioEncoder): + def __init__(self, *args, **kwargs): + super(AudioEncoder_, self).__init__(*args, **kwargs) + + def extract_feature(self, x: Tensor, target_layer: Optional[int] = None): + """ + x : torch.Tensor, shape = (batch_size, n_mels, n_ctx) + the mel spectrogram of the audio + """ + x = F.gelu(self.conv1(x)) + x = F.gelu(self.conv2(x)) + x = x.permute(0, 2, 1) + + length_x = x.shape[1] + if length_x > self.positional_embedding.shape[0]: + self.register_buffer("positional_embedding", sinusoids(length_x, self.positional_embedding.shape[1])) + self.positional_embedding = self.positional_embedding.to(x.device) + x = (x + self.positional_embedding[:length_x, :]).to(x.dtype) + + if target_layer is None: + target_layer = len(self.blocks) + + for block in self.blocks[:target_layer]: + x = block(x) + + return x + + +class Whisper_(Whisper): + def __init__(self, dims: ModelDimensions): + super(Whisper_, self).__init__(dims) + # replace audio encoder with our audio encoder + self.encoder = AudioEncoder_( + self.dims.n_mels, + self.dims.n_audio_ctx, + self.dims.n_audio_state, + self.dims.n_audio_head, + self.dims.n_audio_layer, + ) + + def extract_features(self, mel: torch.Tensor, target_layer: Optional[int] = None): + return self.encoder.extract_feature(mel, target_layer)