jone/Music_Source_Separation
3
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 