Use logger instead of print

This commit is contained in:
litagin02
2024-01-06 13:59:10 +09:00
parent 432689e29f
commit 100949d970
6 changed files with 24 additions and 19 deletions

View File

@@ -7,9 +7,10 @@ from tqdm import tqdm
import commons import commons
import utils import utils
from common.log import logger
from common.stdout_wrapper import SAFE_STDOUT
from config import config from config import config
from text import cleaned_text_to_sequence, get_bert from text import cleaned_text_to_sequence, get_bert
from common.stdout_wrapper import SAFE_STDOUT
def process_line(x): def process_line(x):
@@ -81,4 +82,4 @@ if __name__ == "__main__":
# 这里是缩进的代码块,表示循环体 # 这里是缩进的代码块,表示循环体
pass # 使用pass语句作为占位符 pass # 使用pass语句作为占位符
print(f"bert.pt is generated! total: {len(lines)} bert.pt files.") logger.info(f"bert.pt is generated! total: {len(lines)} bert.pt files.")

View File

@@ -208,10 +208,10 @@ class Config:
def __init__(self, config_path: str, path_config: dict[str, str]): def __init__(self, config_path: str, path_config: dict[str, str]):
if not os.path.isfile(config_path) and os.path.isfile("default_config.yml"): if not os.path.isfile(config_path) and os.path.isfile("default_config.yml"):
shutil.copy(src="default_config.yml", dst=config_path) shutil.copy(src="default_config.yml", dst=config_path)
print( logger.info(
f"A configuration file {config_path} has been generated based on the default configuration file default_config.yml." f"A configuration file {config_path} has been generated based on the default configuration file default_config.yml."
) )
print( logger.info(
"If you have no special needs, please do not modify default_config.yml." "If you have no special needs, please do not modify default_config.yml."
) )
# sys.exit(0) # sys.exit(0)

View File

@@ -345,7 +345,11 @@ class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):
assert all(len(bucket) > 0 for bucket in buckets) assert all(len(bucket) > 0 for bucket in buckets)
# When one bucket is not traversed # When one bucket is not traversed
except Exception as e: except Exception as e:
print("Bucket warning ", e) logger.info("Bucket warning ", e)
logger.info(
"Empty buckets at indices: ",
[i for i, bucket in enumerate(buckets) if len(bucket) == 0],
)
for i in range(len(buckets) - 1, -1, -1): for i in range(len(buckets) - 1, -1, -1):
if len(buckets[i]) == 0: if len(buckets[i]) == 0:
buckets.pop(i) buckets.pop(i)

View File

@@ -10,6 +10,7 @@ from tqdm import tqdm
from config import config from config import config
from text.cleaner import clean_text from text.cleaner import clean_text
from common.stdout_wrapper import SAFE_STDOUT from common.stdout_wrapper import SAFE_STDOUT
from common.log import logger
preprocess_text_config = config.preprocess_text_config preprocess_text_config = config.preprocess_text_config
@@ -70,11 +71,10 @@ def preprocess(
) )
) )
except Exception as e: except Exception as e:
print(line) logger.error(
print( f"An error occurred while generating the training set and validation set, at line:\n{line}\nDetails:\n{e}"
f"An error occurred while generating the training set and validation set! Details:\n{e}"
) )
raise e raise
transcription_path = cleaned_path transcription_path = cleaned_path
spk_utt_map = defaultdict(list) spk_utt_map = defaultdict(list)
@@ -89,12 +89,12 @@ def preprocess(
utt, spk, language, text, phones, tones, word2ph = line.strip().split("|") utt, spk, language, text, phones, tones, word2ph = line.strip().split("|")
if utt in audioPaths: if utt in audioPaths:
# 过滤数据集错误相同的音频匹配多个文本导致后续bert出问题 # 过滤数据集错误相同的音频匹配多个文本导致后续bert出问题
print(f"Same audio matches multiple texts: {line}") logger.warning(f"Same audio matches multiple texts: {line}")
countSame += 1 countSame += 1
continue continue
if not os.path.isfile(utt): if not os.path.isfile(utt):
# 过滤数据集错误:不存在对应音频 # 过滤数据集错误:不存在对应音频
print(f"Audio not found: {utt}") logger.warning(f"Audio not found: {utt}")
countNotFound += 1 countNotFound += 1
continue continue
audioPaths.add(utt) audioPaths.add(utt)
@@ -102,7 +102,7 @@ def preprocess(
if spk not in spk_id_map.keys(): if spk not in spk_id_map.keys():
spk_id_map[spk] = current_sid spk_id_map[spk] = current_sid
current_sid += 1 current_sid += 1
print( logger.info(
f"Total repeated audios: {countSame}, Total number of audio not found: {countNotFound}" f"Total repeated audios: {countSame}, Total number of audio not found: {countNotFound}"
) )
@@ -140,7 +140,7 @@ def preprocess(
) )
with open(config_path, "w", encoding="utf-8") as f: with open(config_path, "w", encoding="utf-8") as f:
json.dump(json_config, f, indent=2, ensure_ascii=False) json.dump(json_config, f, indent=2, ensure_ascii=False)
print("Training set and validation set generation from texts is complete!") logger.info("Training set and validation set generation from texts is complete!")
if __name__ == "__main__": if __name__ == "__main__":

View File

@@ -7,9 +7,9 @@ import pyloudnorm as pyln
import soundfile import soundfile
from tqdm import tqdm from tqdm import tqdm
from config import config
from common.log import logger from common.log import logger
from common.stdout_wrapper import SAFE_STDOUT from common.stdout_wrapper import SAFE_STDOUT
from config import config
DEFAULT_BLOCK_SIZE: float = 0.400 # seconds DEFAULT_BLOCK_SIZE: float = 0.400 # seconds
@@ -119,4 +119,4 @@ if __name__ == "__main__":
pool.close() pool.close()
pool.join() pool.join()
print("Resampling Done!") logger.info("Resampling Done!")

View File

@@ -7,8 +7,9 @@ import torch
from tqdm import tqdm from tqdm import tqdm
import utils import utils
from config import config from common.log import logger
from common.stdout_wrapper import SAFE_STDOUT from common.stdout_wrapper import SAFE_STDOUT
from config import config
warnings.filterwarnings("ignore", category=UserWarning) warnings.filterwarnings("ignore", category=UserWarning)
from pyannote.audio import Inference, Model from pyannote.audio import Inference, Model
@@ -23,8 +24,7 @@ def save_style_vector(wav_path):
try: try:
style_vec = inference(wav_path) style_vec = inference(wav_path)
except Exception as e: except Exception as e:
print(f"\nError occurred with file: {wav_path}") logger.error(f"\nError occurred with file: {wav_path}, Details:\n{e}\n")
print(e)
raise raise
np.save(f"{wav_path}.npy", style_vec) # `test.wav` -> `test.wav.npy` np.save(f"{wav_path}.npy", style_vec) # `test.wav` -> `test.wav.npy`
return style_vec return style_vec
@@ -69,4 +69,4 @@ if __name__ == "__main__":
) )
) )
print(f"Finished generating style vectors! total: {len(wavnames)} npy files.") logger.info(f"Finished generating style vectors! total: {len(wavnames)} npy files.")