Team Ai
Apppublic

jone/Music_Source_Separation

sourceHugging Faceupdated 4y agoView on Hugging Face
3likes
utils.py190 linesDownload Raw Back to bytesep
1import datetime2import logging3import os4import pickle5from typing import Dict, NoReturn6 7import librosa8import numpy as np9import yaml10 11 12def create_logging(log_dir: str, filemode: str) -> logging:13    r"""Create logging to write out log files.14 15    Args:16        logs_dir, str, directory to write out logs17        filemode: str, e.g., "w"18 19    Returns:20        logging21    """22    os.makedirs(log_dir, exist_ok=True)23    i1 = 024 25    while os.path.isfile(os.path.join(log_dir, "{:04d}.log".format(i1))):26        i1 += 127 28    log_path = os.path.join(log_dir, "{:04d}.log".format(i1))29    logging.basicConfig(30        level=logging.DEBUG,31        format="%(asctime)s %(filename)s[line:%(lineno)d] %(levelname)s %(message)s",32        datefmt="%a, %d %b %Y %H:%M:%S",33        filename=log_path,34        filemode=filemode,35    )36 37    # Print to console38    console = logging.StreamHandler()39    console.setLevel(logging.INFO)40    formatter = logging.Formatter("%(name)-12s: %(levelname)-8s %(message)s")41    console.setFormatter(formatter)42    logging.getLogger("").addHandler(console)43 44    return logging45 46 47def load_audio(48    audio_path: str,49    mono: bool,50    sample_rate: float,51    offset: float = 0.0,52    duration: float = None,53) -> np.array:54    r"""Load audio.55 56    Args:57        audio_path: str58        mono: bool59        sample_rate: float60    """61    audio, _ = librosa.core.load(62        audio_path, sr=sample_rate, mono=mono, offset=offset, duration=duration63    )64    # (audio_samples,) | (channels_num, audio_samples)65 66    if audio.ndim == 1:67        audio = audio[None, :]68        # (1, audio_samples,)69 70    return audio71 72 73def load_random_segment(74    audio_path: str, random_state, segment_seconds: float, mono: bool, sample_rate: int75) -> np.array:76    r"""Randomly select an audio segment from a recording."""77 78    duration = librosa.get_duration(filename=audio_path)79 80    start_time = random_state.uniform(0.0, duration - segment_seconds)81 82    audio = load_audio(83        audio_path=audio_path,84        mono=mono,85        sample_rate=sample_rate,86        offset=start_time,87        duration=segment_seconds,88    )89    # (channels_num, audio_samples)90 91    return audio92 93 94def float32_to_int16(x: np.float32) -> np.int16:95 96    x = np.clip(x, a_min=-1, a_max=1)97 98    return (x * 32767.0).astype(np.int16)99 100 101def int16_to_float32(x: np.int16) -> np.float32:102 103    return (x / 32767.0).astype(np.float32)104 105 106def read_yaml(config_yaml: str):107 108    with open(config_yaml, "r") as fr:109        configs = yaml.load(fr, Loader=yaml.FullLoader)110 111    return configs112 113 114def check_configs_gramma(configs: Dict) -> NoReturn:115    r"""Check if the gramma of the config dictionary for training is legal."""116    input_source_types = configs['train']['input_source_types']117 118    for augmentation_type in configs['train']['augmentations'].keys():119        augmentation_dict = configs['train']['augmentations'][augmentation_type]120 121        for source_type in augmentation_dict.keys():122            if source_type not in input_source_types:123                error_msg = (124                    "The source type '{}'' in configs['train']['augmentations']['{}'] "125                    "must be one of input_source_types {}".format(126                        source_type, augmentation_type, input_source_types127                    )128                )129                raise Exception(error_msg)130 131 132def magnitude_to_db(x: float) -> float:133    eps = 1e-10134    return 20.0 * np.log10(max(x, eps))135 136 137def db_to_magnitude(x: float) -> float:138    return 10.0 ** (x / 20)139 140 141def get_pitch_shift_factor(shift_pitch: float) -> float:142    r"""The factor of the audio length to be scaled."""143    return 2 ** (shift_pitch / 12)144 145 146class StatisticsContainer(object):147    def __init__(self, statistics_path):148        self.statistics_path = statistics_path149 150        self.backup_statistics_path = "{}_{}.pkl".format(151            os.path.splitext(self.statistics_path)[0],152            datetime.datetime.now().strftime("%Y-%m-%d_%H-%M-%S"),153        )154 155        self.statistics_dict = {"train": [], "test": []}156 157    def append(self, steps, statistics, split):158        statistics["steps"] = steps159        self.statistics_dict[split].append(statistics)160 161    def dump(self):162        pickle.dump(self.statistics_dict, open(self.statistics_path, "wb"))163        pickle.dump(self.statistics_dict, open(self.backup_statistics_path, "wb"))164        logging.info("    Dump statistics to {}".format(self.statistics_path))165        logging.info("    Dump statistics to {}".format(self.backup_statistics_path))166 167    '''168    def load_state_dict(self, resume_steps):169        self.statistics_dict = pickle.load(open(self.statistics_path, "rb"))170 171        resume_statistics_dict = {"train": [], "test": []}172 173        for key in self.statistics_dict.keys():174            for statistics in self.statistics_dict[key]:175                if statistics["steps"] <= resume_steps:176                    resume_statistics_dict[key].append(statistics)177 178        self.statistics_dict = resume_statistics_dict179    '''180 181 182def calculate_sdr(ref: np.array, est: np.array) -> float:183    s_true = ref184    s_artif = est - ref185    sdr = 10.0 * (186        np.log10(np.clip(np.mean(s_true ** 2), 1e-8, np.inf))187        - np.log10(np.clip(np.mean(s_artif ** 2), 1e-8, np.inf))188    )189    return sdr190