Restore gc and emptycache when training

This commit is contained in:
litagin02
2024-02-05 16:02:16 +09:00
parent 48e24c020b
commit b0bb147d44
3 changed files with 10 additions and 6 deletions

View File

@@ -1,5 +1,6 @@
import argparse
import datetime
import gc
import os
import platform
@@ -757,8 +758,9 @@ def train_and_evaluate(
)
pbar.update()
# 本家ではこれをスピードアップのために消すと書かれていたので、一応消してみる
# gc.collect()
# torch.cuda.empty_cache()
# と思ったけどメモリ使用量が減るかもしれないのでつけてみる
gc.collect()
torch.cuda.empty_cache()
if pbar is None and rank == 0:
logger.info(f"====> Epoch: {epoch}, step: {global_step}")

View File

@@ -1,5 +1,6 @@
import argparse
import datetime
import gc
import os
import platform
@@ -589,7 +590,7 @@ def train_and_evaluate(
language,
bert,
style_vec,
) in enumerate(tqdm(train_loader)):
) in enumerate(train_loader):
if net_g.module.use_noise_scaled_mas:
current_mas_noise_scale = (
net_g.module.mas_noise_scale_initial
@@ -899,8 +900,9 @@ def train_and_evaluate(
)
pbar.update()
# 本家ではこれをスピードアップのために消すと書かれていたので、一応消してみる
# gc.collect()
# torch.cuda.empty_cache()
# と思ったけどメモリ使用量が減るかもしれないのでつけてみる
gc.collect()
torch.cuda.empty_cache()
if pbar is None and rank == 0:
logger.info(f"====> Epoch: {epoch}, step: {global_step}")

View File

@@ -302,7 +302,7 @@ def train(model_name, skip_style=False, use_jp_extra=True):
cmd = [train_py, "--config", config_path, "--model", dataset_path]
if skip_style:
cmd.append("--skip_default_style")
success, message = run_script_with_log(cmd)
success, message = run_script_with_log(cmd, ignore_warning=True)
if not success:
logger.error(f"Train failed.")
return False, f"Error: 学習に失敗しました:\n{message}"