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")