From ff8c6cd7cf7b5493c65c18f827c7c7553b24e12a Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 30 Dec 2023 00:48:53 +0900 Subject: [PATCH 01/12] Add share option to App --- app.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/app.py b/app.py index 09ced91..5f270c8 100644 --- a/app.py +++ b/app.py @@ -360,6 +360,9 @@ if __name__ == "__main__": parser.add_argument( "--dir", "-d", type=str, help="Model directory", default=config.out_dir ) + parser.add_argument( + "--share", action="store_true", help="Share this app publicly", default=False + ) args = parser.parse_args() model_dir = args.dir @@ -518,4 +521,4 @@ if __name__ == "__main__": outputs=[style, ref_audio_path], ) - app.launch(inbrowser=True) + app.launch(inbrowser=True, share=args.share) From 6a3e1ca94cabdbb178db4691b392b3046574f66b Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 30 Dec 2023 09:16:05 +0900 Subject: [PATCH 02/12] Feat: in colab, use Google drive path as dataset_path --- webui_train.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/webui_train.py b/webui_train.py index 41486f3..2b749b1 100644 --- a/webui_train.py +++ b/webui_train.py @@ -11,8 +11,15 @@ from tools.log import logger from tools.subprocess_utils import run_script_with_log, second_elem_of +is_colab = "google.colab" in sys.modules + + def get_path(model_name): assert model_name != "", "モデル名は空にできません" + if is_colab: + dataset_path = os.path.join( + "/content/drive/MyDrive/Style-Bert-VITS2/Data", model_name + ) dataset_path = os.path.join("Data", model_name) lbl_path = os.path.join(dataset_path, "esd.list") train_path = os.path.join(dataset_path, "train.list") From f1441c35f290ad6d30c38d97a107476c9d6d5456 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 30 Dec 2023 09:25:02 +0900 Subject: [PATCH 03/12] Fix --- webui_train.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/webui_train.py b/webui_train.py index 2b749b1..fb894cd 100644 --- a/webui_train.py +++ b/webui_train.py @@ -20,7 +20,8 @@ def get_path(model_name): dataset_path = os.path.join( "/content/drive/MyDrive/Style-Bert-VITS2/Data", model_name ) - dataset_path = os.path.join("Data", model_name) + else: + dataset_path = os.path.join("Data", model_name) lbl_path = os.path.join(dataset_path, "esd.list") train_path = os.path.join(dataset_path, "train.list") val_path = os.path.join(dataset_path, "val.list") From 35a8bb514f0e7779cf0d6f9e9d960bc02c9676f6 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 30 Dec 2023 09:58:32 +0900 Subject: [PATCH 04/12] Fix: colab stdout error, so replace with stdout wrapper --- bert_gen.py | 4 ++-- preprocess_text.py | 4 ++-- resample.py | 4 ++-- slice.py | 5 +++-- style_gen.py | 4 ++-- tools/log.py | 4 ++-- tools/subprocess_utils.py | 3 ++- train_ms.py | 4 ++-- transcribe.py | 4 +++- 9 files changed, 20 insertions(+), 16 deletions(-) diff --git a/bert_gen.py b/bert_gen.py index d72ce4e..7ddcb23 100644 --- a/bert_gen.py +++ b/bert_gen.py @@ -1,5 +1,4 @@ import argparse -import sys from multiprocessing import Pool import torch @@ -10,6 +9,7 @@ import commons import utils from config import config from text import cleaned_text_to_sequence, get_bert +from tools.stdout_wrapper import get_stdout def process_line(x): @@ -76,7 +76,7 @@ if __name__ == "__main__": for _ in tqdm( pool.imap_unordered(process_line, zip(lines, add_blank)), total=len(lines), - file=sys.stdout, + file=get_stdout(), ): # 这里是缩进的代码块,表示循环体 pass # 使用pass语句作为占位符 diff --git a/preprocess_text.py b/preprocess_text.py index 38683c8..de19331 100644 --- a/preprocess_text.py +++ b/preprocess_text.py @@ -1,6 +1,5 @@ import json import os -import sys from collections import defaultdict from random import shuffle from typing import Optional @@ -10,6 +9,7 @@ from tqdm import tqdm from config import config from text.cleaner import clean_text +from tools.stdout_wrapper import get_stdout preprocess_text_config = config.preprocess_text_config @@ -52,7 +52,7 @@ def preprocess( lines = trans_file.readlines() # print(lines, ' ', len(lines)) if len(lines) != 0: - for line in tqdm(lines, file=sys.stdout): + for line in tqdm(lines, file=get_stdout()): try: utt, spk, language, text = line.strip().split("|") norm_text, phones, tones, word2ph = clean_text( diff --git a/resample.py b/resample.py index a6a8143..f203ecb 100644 --- a/resample.py +++ b/resample.py @@ -1,6 +1,5 @@ import argparse import os -import sys from multiprocessing import Pool, cpu_count import librosa @@ -10,6 +9,7 @@ from tqdm import tqdm from config import config from tools.log import logger +from tools.stdout_wrapper import get_stdout def normalize_audio(data, sr): @@ -88,7 +88,7 @@ if __name__ == "__main__": pool = Pool(processes=processes) for _ in tqdm( - pool.imap_unordered(process, tasks), file=sys.stdout, total=len(tasks) + pool.imap_unordered(process, tasks), file=get_stdout(), total=len(tasks) ): pass diff --git a/slice.py b/slice.py index c4faf0a..aa5676b 100644 --- a/slice.py +++ b/slice.py @@ -1,12 +1,13 @@ import argparse import os import shutil -import sys import soundfile as sf import torch from tqdm import tqdm +from tools.stdout_wrapper import get_stdout + vad_model, utils = torch.hub.load( repo_or_dir="snakers4/silero-vad", model="silero_vad", @@ -106,7 +107,7 @@ if __name__ == "__main__": shutil.rmtree(output_dir) total_sec = 0 - for wav_file in tqdm(wav_files, file=sys.stdout): + for wav_file in tqdm(wav_files, file=get_stdout()): time_sec = split_wav( wav_file, output_dir, diff --git a/style_gen.py b/style_gen.py index 7d7f65e..7312bd2 100644 --- a/style_gen.py +++ b/style_gen.py @@ -1,6 +1,5 @@ import argparse import concurrent.futures -import sys import warnings import numpy as np @@ -9,6 +8,7 @@ from tqdm import tqdm import utils from config import config +from tools.stdout_wrapper import get_stdout warnings.filterwarnings("ignore", category=UserWarning) from pyannote.audio import Inference, Model @@ -59,7 +59,7 @@ if __name__ == "__main__": tqdm( executor.map(save_style_vector, wavnames), total=len(wavnames), - file=sys.stdout, + file=get_stdout(), ) ) diff --git a/tools/log.py b/tools/log.py index 85526cb..6d7990a 100644 --- a/tools/log.py +++ b/tools/log.py @@ -2,8 +2,8 @@ logger封装 """ from loguru import logger -import sys +from .stdout_wrapper import get_stdout # 移除所有默认的处理器 logger.remove() @@ -13,4 +13,4 @@ log_format = ( "{time:MM-DD HH:mm:ss} |{level:^8}| {file}:{line} | {message}" ) -logger.add(sys.stdout, format=log_format, backtrace=True, diagnose=True) +logger.add(get_stdout(), format=log_format, backtrace=True, diagnose=True) diff --git a/tools/subprocess_utils.py b/tools/subprocess_utils.py index dea2e0d..b377656 100644 --- a/tools/subprocess_utils.py +++ b/tools/subprocess_utils.py @@ -2,6 +2,7 @@ import subprocess import sys from .log import logger +from .stdout_wrapper import get_stdout python = sys.executable @@ -10,7 +11,7 @@ def run_script_with_log(cmd: list[str]) -> tuple[bool, str]: logger.info(f"Running: {' '.join(cmd)}") result = subprocess.run( [python] + cmd, - stdout=sys.stdout, + stdout=get_stdout(), # type: ignore stderr=subprocess.PIPE, text=True, ) diff --git a/train_ms.py b/train_ms.py index 8272d5e..92177bb 100644 --- a/train_ms.py +++ b/train_ms.py @@ -4,7 +4,6 @@ import gc import os import platform import shutil -import sys import torch import torch.distributed as dist @@ -29,6 +28,7 @@ from mel_processing import mel_spectrogram_torch, spec_to_mel_torch from models import DurationDiscriminator, MultiPeriodDiscriminator, SynthesizerTrn from text.symbols import symbols from tools.log import logger +from tools.stdout_wrapper import get_stdout torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = ( @@ -427,7 +427,7 @@ def train_and_evaluate( ja_bert, en_bert, style_vec, - ) in enumerate(tqdm(train_loader, file=sys.stdout)): + ) in enumerate(tqdm(train_loader, file=get_stdout())): if net_g.module.use_noise_scaled_mas: current_mas_noise_scale = ( net_g.module.mas_noise_scale_initial diff --git a/transcribe.py b/transcribe.py index 4f2b955..234359b 100644 --- a/transcribe.py +++ b/transcribe.py @@ -5,6 +5,8 @@ import sys from faster_whisper import WhisperModel from tqdm import tqdm +from tools.stdout_wrapper import get_stdout + def transcribe(wav_path, initial_prompt=None): segments, _ = model.transcribe( @@ -45,7 +47,7 @@ if __name__ == "__main__": os.rename(output_file, output_file + ".bak") with open(output_file, "w", encoding="utf-8") as f: - for wav_file in tqdm(wav_files, file=sys.stdout): + for wav_file in tqdm(wav_files, file=get_stdout()): file_name = os.path.basename(wav_file) text = transcribe(wav_file, initial_prompt=initial_prompt) f.write(f"{file_name}|{speaker_name}|JP|{text}\n") From 85e9374c04e24a35c91c74099ce9e887ad2eb281 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 30 Dec 2023 10:03:30 +0900 Subject: [PATCH 05/12] I'm stupid --- tools/stdout_wrapper.py | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) create mode 100644 tools/stdout_wrapper.py diff --git a/tools/stdout_wrapper.py b/tools/stdout_wrapper.py new file mode 100644 index 0000000..c2c9fe0 --- /dev/null +++ b/tools/stdout_wrapper.py @@ -0,0 +1,17 @@ +import sys + + +class StdoutWrapper: + def write(self, message: str): + print(message, end="") + + def flush(self): + pass + + +def get_stdout(): + # Colab 環境をチェックする + if "google.colab" in sys.modules: + return StdoutWrapper() + else: + return sys.stdout From 24ac11c564760191dba349d125bc9abe54a0d782 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 30 Dec 2023 10:18:54 +0900 Subject: [PATCH 06/12] Fix stdout wrapper? --- tools/stdout_wrapper.py | 21 +++++++++++++++++++-- 1 file changed, 19 insertions(+), 2 deletions(-) diff --git a/tools/stdout_wrapper.py b/tools/stdout_wrapper.py index c2c9fe0..7841756 100644 --- a/tools/stdout_wrapper.py +++ b/tools/stdout_wrapper.py @@ -1,12 +1,29 @@ import sys +import tempfile class StdoutWrapper: + def __init__(self): + # 一時ファイルの作成とオープン + self.temp_file = tempfile.NamedTemporaryFile(mode="w+", delete=False) + def write(self, message: str): - print(message, end="") + # メッセージを一時ファイルに書き込む + self.temp_file.write(message) + self.temp_file.flush() def flush(self): - pass + # 一時ファイルのフラッシュ + self.temp_file.flush() + + def read(self): + # ファイルの内容を読み出す + self.temp_file.seek(0) # ファイルの先頭にシーク + return self.temp_file.read() + + def close(self): + # 一時ファイルを閉じる + self.temp_file.close() def get_stdout(): From e77fb5383bfaa1a628c59c0ff11acf13439cae21 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 30 Dec 2023 10:48:48 +0900 Subject: [PATCH 07/12] Try to fix... --- tools/stdout_wrapper.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/tools/stdout_wrapper.py b/tools/stdout_wrapper.py index 7841756..c3f4959 100644 --- a/tools/stdout_wrapper.py +++ b/tools/stdout_wrapper.py @@ -25,6 +25,10 @@ class StdoutWrapper: # 一時ファイルを閉じる self.temp_file.close() + def fileno(self): + # 一時ファイルのファイルディスクリプタを返す + return self.temp_file.fileno() + def get_stdout(): # Colab 環境をチェックする From 255ac994ec1580a6c9a42566bd8c46498a767ed2 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 30 Dec 2023 11:02:07 +0900 Subject: [PATCH 08/12] Add print to custom stdout --- tools/stdout_wrapper.py | 10 +++------- 1 file changed, 3 insertions(+), 7 deletions(-) diff --git a/tools/stdout_wrapper.py b/tools/stdout_wrapper.py index c3f4959..24589c6 100644 --- a/tools/stdout_wrapper.py +++ b/tools/stdout_wrapper.py @@ -4,29 +4,25 @@ import tempfile class StdoutWrapper: def __init__(self): - # 一時ファイルの作成とオープン self.temp_file = tempfile.NamedTemporaryFile(mode="w+", delete=False) + self.original_stdout = sys.stdout def write(self, message: str): - # メッセージを一時ファイルに書き込む self.temp_file.write(message) self.temp_file.flush() + print(message, end="", file=self.original_stdout) def flush(self): - # 一時ファイルのフラッシュ self.temp_file.flush() def read(self): - # ファイルの内容を読み出す - self.temp_file.seek(0) # ファイルの先頭にシーク + self.temp_file.seek(0) return self.temp_file.read() def close(self): - # 一時ファイルを閉じる self.temp_file.close() def fileno(self): - # 一時ファイルのファイルディスクリプタを返す return self.temp_file.fileno() From 67b417fdf93d84e4c2b2f54a63d0639a86922ef2 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 30 Dec 2023 14:33:58 +0900 Subject: [PATCH 09/12] Feat: make default style when training (WIP) --- make_default_style.py | 17 +++++++++++++++++ style_gen.py | 9 +++++++-- 2 files changed, 24 insertions(+), 2 deletions(-) create mode 100644 make_default_style.py diff --git a/make_default_style.py b/make_default_style.py new file mode 100644 index 0000000..ca708f6 --- /dev/null +++ b/make_default_style.py @@ -0,0 +1,17 @@ +import os +import numpy as np +import argparse + +parser = argparse.ArgumentParser() +parser.add_argument("--wav_dir", type=str, default="data/wav") + +embs = [] +names = [] +for file in os.listdir(wav_dir): + if file.endswith(".npy"): + xvec = np.load(os.path.join(wav_dir, file)) + embs.append(np.expand_dims(xvec, axis=0)) + names.append(file) + +x = np.concatenate(embs, axis=0) +x = np.squeeze(x) diff --git a/style_gen.py b/style_gen.py index 7312bd2..0d5006d 100644 --- a/style_gen.py +++ b/style_gen.py @@ -25,8 +25,13 @@ def extract_style_vector(wav_path): def save_style_vector(wav_path): style_vec = extract_style_vector(wav_path) - # `test.wav` -> `test.wav.npy` - np.save(f"{wav_path}.npy", style_vec) + np.save(f"{wav_path}.npy", style_vec) # `test.wav` -> `test.wav.npy` + return style_vec + + +def save_average_style_vector(style_vectors, filename="style_vectors.npy"): + average_vector = np.mean(style_vectors, axis=0) + np.save(filename, average_vector) if __name__ == "__main__": From 5ecf6c66bd0f51d583cfae9fd3a33cc5927789e9 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 30 Dec 2023 17:28:07 +0900 Subject: [PATCH 10/12] Feat: save default style, and colab train support (maybe) --- default_style.py | 29 +++++++++++++++++++++++++++++ make_default_style.py | 17 ----------------- train_ms.py | 37 ++++++++++++++++++++++++++++++++++++- webui_train.py | 7 ++++--- 4 files changed, 69 insertions(+), 21 deletions(-) create mode 100644 default_style.py delete mode 100644 make_default_style.py diff --git a/default_style.py b/default_style.py new file mode 100644 index 0000000..424f8eb --- /dev/null +++ b/default_style.py @@ -0,0 +1,29 @@ +import os +from tools.log import logger + +import numpy as np +import json + + +def set_style_config(json_path, output_path): + with open(json_path, "r") as f: + json_dict = json.load(f) + json_dict["data"]["num_styles"] = 1 + json_dict["data"]["style2id"] = {"Neutral": 0} + with open(output_path, "w") as f: + json.dump(json_dict, f, indent=2) + logger.info(f"Update style config (only Neutral style) to {output_path}") + + +def save_mean_vector(wav_dir, output_path): + embs = [] + for file in os.listdir(wav_dir): + if file.endswith(".npy"): + xvec = np.load(os.path.join(wav_dir, file)) + embs.append(np.expand_dims(xvec, axis=0)) + + x = np.concatenate(embs, axis=0) # (N, 256) + mean = np.mean(x, axis=0) # (256,) + only_mean = np.stack([mean]) # (1, 256) + np.save(output_path, only_mean) + logger.info(f"Saved mean style vector to {output_path}") diff --git a/make_default_style.py b/make_default_style.py deleted file mode 100644 index ca708f6..0000000 --- a/make_default_style.py +++ /dev/null @@ -1,17 +0,0 @@ -import os -import numpy as np -import argparse - -parser = argparse.ArgumentParser() -parser.add_argument("--wav_dir", type=str, default="data/wav") - -embs = [] -names = [] -for file in os.listdir(wav_dir): - if file.endswith(".npy"): - xvec = np.load(os.path.join(wav_dir, file)) - embs.append(np.expand_dims(xvec, axis=0)) - names.append(file) - -x = np.concatenate(embs, axis=0) -x = np.squeeze(x) diff --git a/train_ms.py b/train_ms.py index 92177bb..cb49a82 100644 --- a/train_ms.py +++ b/train_ms.py @@ -4,6 +4,7 @@ import gc import os import platform import shutil +import sys import torch import torch.distributed as dist @@ -29,6 +30,7 @@ from models import DurationDiscriminator, MultiPeriodDiscriminator, SynthesizerT from text.symbols import symbols from tools.log import logger from tools.stdout_wrapper import get_stdout +import default_style torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = ( @@ -41,6 +43,9 @@ torch.backends.cuda.enable_mem_efficient_sdp( True ) # Not available if torch version is lower than 2.0 torch.backends.cuda.enable_math_sdp(True) + +IS_COLAB = "google.colab" in sys.modules + global_step = 0 @@ -106,9 +111,39 @@ def run(): data = f.read() with open(config.train_ms_config.config_path, "w", encoding="utf-8") as f: f.write(data) + + """ + Path constants are a bit complicated... + TODO: Refactor or rename these? + (Both `config.yml` and `config.json` are used, which is confusing I think.) + + args.model: For saving all info needed for training. + default: `Data/{model_name}`. + hps.model_dir = model_dir: For saving checkpoints (for resuming training). + default: `Data/{model_name}/models`. + + config.out_dir: Root directory of model assets needed for inference. + default: `model_assets`. + + out_dir: For saving resulting models (for inference). + default: `model_assets/{model_name}`, which is used for inference. + """ + if IS_COLAB: + config.out_dir = "/content/drive/MyDrive/Style-Bert-VITS2/model_assets" + logger.info( + "Colab detected, so use mounted Google Drive as directory for saving resulting models:" + ) + logger.info(config.out_dir) + os.makedirs(config.out_dir, exist_ok=True) out_dir = os.path.join(config.out_dir, config.model_name) os.makedirs(out_dir, exist_ok=True) - shutil.copy(args.config, os.path.join(out_dir, "config.json")) + + # Save default style to out_dir + default_style.set_style_config(args.config, os.path.join(out_dir, "config.json")) + default_style.save_mean_vector( + os.path.join(args.model, "wavs"), + os.path.join(out_dir, "style_vectors.npy"), + ) torch.manual_seed(hps.train.seed) torch.cuda.set_device(local_rank) diff --git a/webui_train.py b/webui_train.py index fb894cd..db72bdd 100644 --- a/webui_train.py +++ b/webui_train.py @@ -10,16 +10,17 @@ import yaml from tools.log import logger from tools.subprocess_utils import run_script_with_log, second_elem_of - -is_colab = "google.colab" in sys.modules +IS_COLAB = "google.colab" in sys.modules def get_path(model_name): assert model_name != "", "モデル名は空にできません" - if is_colab: + if IS_COLAB: + logger.info("Colab detected, so use mounted Google Drive as dataset path:") dataset_path = os.path.join( "/content/drive/MyDrive/Style-Bert-VITS2/Data", model_name ) + logger.info(dataset_path) else: dataset_path = os.path.join("Data", model_name) lbl_path = os.path.join(dataset_path, "esd.list") From 6e8bcc2219ab075dcb7bf822f268c5e4dabc8b5a Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 30 Dec 2023 18:05:40 +0900 Subject: [PATCH 11/12] Fix colab check --- train_ms.py | 9 +++++++-- webui_train.py | 14 +++++++++++--- 2 files changed, 18 insertions(+), 5 deletions(-) diff --git a/train_ms.py b/train_ms.py index cb49a82..a917569 100644 --- a/train_ms.py +++ b/train_ms.py @@ -17,6 +17,7 @@ from tqdm import tqdm # logging.getLogger("numba").setLevel(logging.WARNING) import commons +import default_style import utils from config import config from data_utils import ( @@ -30,7 +31,6 @@ from models import DurationDiscriminator, MultiPeriodDiscriminator, SynthesizerT from text.symbols import symbols from tools.log import logger from tools.stdout_wrapper import get_stdout -import default_style torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = ( @@ -44,7 +44,12 @@ torch.backends.cuda.enable_mem_efficient_sdp( ) # Not available if torch version is lower than 2.0 torch.backends.cuda.enable_math_sdp(True) -IS_COLAB = "google.colab" in sys.modules +try: + import google.colab + + IS_COLAB = True +except ImportError: + IS_COLAB = False global_step = 0 diff --git a/webui_train.py b/webui_train.py index db72bdd..c44f464 100644 --- a/webui_train.py +++ b/webui_train.py @@ -10,7 +10,12 @@ import yaml from tools.log import logger from tools.subprocess_utils import run_script_with_log, second_elem_of -IS_COLAB = "google.colab" in sys.modules +try: + import google.colab + + IS_COLAB = True +except ImportError: + IS_COLAB = False def get_path(model_name): @@ -48,9 +53,12 @@ def initialize(model_name, batch_size, epochs, save_every_steps, bf16_run): model_path = os.path.join(dataset_path, "models") try: - shutil.copytree(src="pretrained", dst=model_path) + shutil.copytree( + src="pretrained", + dst=model_path, + ) except FileExistsError: - logger.error(f"Step 1: {model_path} already exists.") + logger.warning(f"Step 1: {model_path} already exists.") return False, f"Step1, Error: モデルフォルダ {model_path} が既に存在します。問題なければ削除してください。" except FileNotFoundError: logger.error("Step 1: `pretrained` folder not found.") From a9652979ccdcc02a0df2de76dcaf911eb40ccf73 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 30 Dec 2023 18:45:04 +0900 Subject: [PATCH 12/12] Replace get_stdout() with SAFE_STDOUT --- bert_gen.py | 4 ++-- preprocess_text.py | 4 ++-- resample.py | 4 ++-- slice.py | 4 ++-- style_gen.py | 4 ++-- tools/log.py | 4 ++-- tools/stdout_wrapper.py | 12 ++++++------ tools/subprocess_utils.py | 4 ++-- train_ms.py | 4 ++-- transcribe.py | 4 ++-- 10 files changed, 24 insertions(+), 24 deletions(-) diff --git a/bert_gen.py b/bert_gen.py index 7ddcb23..e41484d 100644 --- a/bert_gen.py +++ b/bert_gen.py @@ -9,7 +9,7 @@ import commons import utils from config import config from text import cleaned_text_to_sequence, get_bert -from tools.stdout_wrapper import get_stdout +from tools.stdout_wrapper import SAFE_STDOUT def process_line(x): @@ -76,7 +76,7 @@ if __name__ == "__main__": for _ in tqdm( pool.imap_unordered(process_line, zip(lines, add_blank)), total=len(lines), - file=get_stdout(), + file=SAFE_STDOUT, ): # 这里是缩进的代码块,表示循环体 pass # 使用pass语句作为占位符 diff --git a/preprocess_text.py b/preprocess_text.py index de19331..d905f19 100644 --- a/preprocess_text.py +++ b/preprocess_text.py @@ -9,7 +9,7 @@ from tqdm import tqdm from config import config from text.cleaner import clean_text -from tools.stdout_wrapper import get_stdout +from tools.stdout_wrapper import SAFE_STDOUT preprocess_text_config = config.preprocess_text_config @@ -52,7 +52,7 @@ def preprocess( lines = trans_file.readlines() # print(lines, ' ', len(lines)) if len(lines) != 0: - for line in tqdm(lines, file=get_stdout()): + for line in tqdm(lines, file=SAFE_STDOUT): try: utt, spk, language, text = line.strip().split("|") norm_text, phones, tones, word2ph = clean_text( diff --git a/resample.py b/resample.py index f203ecb..229b492 100644 --- a/resample.py +++ b/resample.py @@ -9,7 +9,7 @@ from tqdm import tqdm from config import config from tools.log import logger -from tools.stdout_wrapper import get_stdout +from tools.stdout_wrapper import SAFE_STDOUT def normalize_audio(data, sr): @@ -88,7 +88,7 @@ if __name__ == "__main__": pool = Pool(processes=processes) for _ in tqdm( - pool.imap_unordered(process, tasks), file=get_stdout(), total=len(tasks) + pool.imap_unordered(process, tasks), file=SAFE_STDOUT, total=len(tasks) ): pass diff --git a/slice.py b/slice.py index aa5676b..3cd2399 100644 --- a/slice.py +++ b/slice.py @@ -6,7 +6,7 @@ import soundfile as sf import torch from tqdm import tqdm -from tools.stdout_wrapper import get_stdout +from tools.stdout_wrapper import SAFE_STDOUT vad_model, utils = torch.hub.load( repo_or_dir="snakers4/silero-vad", @@ -107,7 +107,7 @@ if __name__ == "__main__": shutil.rmtree(output_dir) total_sec = 0 - for wav_file in tqdm(wav_files, file=get_stdout()): + for wav_file in tqdm(wav_files, file=SAFE_STDOUT): time_sec = split_wav( wav_file, output_dir, diff --git a/style_gen.py b/style_gen.py index 0d5006d..6fdf9e3 100644 --- a/style_gen.py +++ b/style_gen.py @@ -8,7 +8,7 @@ from tqdm import tqdm import utils from config import config -from tools.stdout_wrapper import get_stdout +from tools.stdout_wrapper import SAFE_STDOUT warnings.filterwarnings("ignore", category=UserWarning) from pyannote.audio import Inference, Model @@ -64,7 +64,7 @@ if __name__ == "__main__": tqdm( executor.map(save_style_vector, wavnames), total=len(wavnames), - file=get_stdout(), + file=SAFE_STDOUT, ) ) diff --git a/tools/log.py b/tools/log.py index 6d7990a..51dca5f 100644 --- a/tools/log.py +++ b/tools/log.py @@ -3,7 +3,7 @@ logger封装 """ from loguru import logger -from .stdout_wrapper import get_stdout +from .stdout_wrapper import SAFE_STDOUT # 移除所有默认的处理器 logger.remove() @@ -13,4 +13,4 @@ log_format = ( "{time:MM-DD HH:mm:ss} |{level:^8}| {file}:{line} | {message}" ) -logger.add(get_stdout(), format=log_format, backtrace=True, diagnose=True) +logger.add(SAFE_STDOUT, format=log_format, backtrace=True, diagnose=True) diff --git a/tools/stdout_wrapper.py b/tools/stdout_wrapper.py index 24589c6..23c6e76 100644 --- a/tools/stdout_wrapper.py +++ b/tools/stdout_wrapper.py @@ -26,9 +26,9 @@ class StdoutWrapper: return self.temp_file.fileno() -def get_stdout(): - # Colab 環境をチェックする - if "google.colab" in sys.modules: - return StdoutWrapper() - else: - return sys.stdout +try: + import google.colab + + SAFE_STDOUT = StdoutWrapper() +except ImportError: + SAFE_STDOUT = sys.stdout diff --git a/tools/subprocess_utils.py b/tools/subprocess_utils.py index b377656..ad15575 100644 --- a/tools/subprocess_utils.py +++ b/tools/subprocess_utils.py @@ -2,7 +2,7 @@ import subprocess import sys from .log import logger -from .stdout_wrapper import get_stdout +from .stdout_wrapper import SAFE_STDOUT python = sys.executable @@ -11,7 +11,7 @@ def run_script_with_log(cmd: list[str]) -> tuple[bool, str]: logger.info(f"Running: {' '.join(cmd)}") result = subprocess.run( [python] + cmd, - stdout=get_stdout(), # type: ignore + stdout=SAFE_STDOUT, # type: ignore stderr=subprocess.PIPE, text=True, ) diff --git a/train_ms.py b/train_ms.py index a917569..6cd8978 100644 --- a/train_ms.py +++ b/train_ms.py @@ -30,7 +30,7 @@ from mel_processing import mel_spectrogram_torch, spec_to_mel_torch from models import DurationDiscriminator, MultiPeriodDiscriminator, SynthesizerTrn from text.symbols import symbols from tools.log import logger -from tools.stdout_wrapper import get_stdout +from tools.stdout_wrapper import SAFE_STDOUT torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = ( @@ -467,7 +467,7 @@ def train_and_evaluate( ja_bert, en_bert, style_vec, - ) in enumerate(tqdm(train_loader, file=get_stdout())): + ) in enumerate(tqdm(train_loader, file=SAFE_STDOUT)): if net_g.module.use_noise_scaled_mas: current_mas_noise_scale = ( net_g.module.mas_noise_scale_initial diff --git a/transcribe.py b/transcribe.py index 234359b..7d2dd1a 100644 --- a/transcribe.py +++ b/transcribe.py @@ -5,7 +5,7 @@ import sys from faster_whisper import WhisperModel from tqdm import tqdm -from tools.stdout_wrapper import get_stdout +from tools.stdout_wrapper import SAFE_STDOUT def transcribe(wav_path, initial_prompt=None): @@ -47,7 +47,7 @@ if __name__ == "__main__": os.rename(output_file, output_file + ".bak") with open(output_file, "w", encoding="utf-8") as f: - for wav_file in tqdm(wav_files, file=get_stdout()): + for wav_file in tqdm(wav_files, file=SAFE_STDOUT): file_name = os.path.basename(wav_file) text = transcribe(wav_file, initial_prompt=initial_prompt) f.write(f"{file_name}|{speaker_name}|JP|{text}\n")