From 391dea85de5982b039c8d8fe1b653494c0e12292 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=BA=90=E6=96=87=E9=9B=A8?= <41315874+fumiama@users.noreply.github.com> Date: Mon, 4 Sep 2023 00:47:33 +0800 Subject: [PATCH 1/3] =?UTF-8?q?feat:=20=E4=BC=98=E5=8C=96=E6=97=A5?= =?UTF-8?q?=E5=BF=97=E6=89=93=E5=8D=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- attentions.py | 9 +++++---- utils.py | 6 ++---- webui.py | 18 ++++++++++++++++-- 3 files changed, 23 insertions(+), 10 deletions(-) diff --git a/attentions.py b/attentions.py index b5f2fec..1192dd7 100644 --- a/attentions.py +++ b/attentions.py @@ -1,13 +1,14 @@ import copy import math -import numpy as np import torch from torch import nn from torch.nn import functional as F import commons -import modules -from torch.nn.utils import weight_norm, remove_weight_norm +import logging + +logger = logging.getLogger(__name__) + class LayerNorm(nn.Module): def __init__(self, channels, eps=1e-5): super().__init__() @@ -55,7 +56,7 @@ class Encoder(nn.Module): self.spk_emb_linear = nn.Linear(self.gin_channels, self.hidden_channels) # vits2 says 3rd block, so idx is 2 by default self.cond_layer_idx = kwargs['cond_layer_idx'] if 'cond_layer_idx' in kwargs else 2 - print(self.gin_channels, self.cond_layer_idx) + logging.debug(self.gin_channels, self.cond_layer_idx) assert self.cond_layer_idx < self.n_layers, 'cond_layer_idx should be less than n_layers' self.drop = nn.Dropout(p_dropout) self.attn_layers = nn.ModuleList() diff --git a/utils.py b/utils.py index fbd7149..bf825ea 100644 --- a/utils.py +++ b/utils.py @@ -11,8 +11,7 @@ import torch MATPLOTLIB_FLAG = False -logging.basicConfig(stream=sys.stdout, level=logging.DEBUG) -logger = logging +logger = logging.getLogger(__name__) def load_checkpoint(checkpoint_path, model, optimizer=None, skip_optimizer=False): @@ -42,13 +41,12 @@ def load_checkpoint(checkpoint_path, model, optimizer=None, skip_optimizer=False new_state_dict[k] = saved_state_dict[k] assert saved_state_dict[k].shape == v.shape, (saved_state_dict[k].shape, v.shape) except: - print("error, %s is not in the checkpoint" % k) + logger.error("%s is not in the checkpoint" % k) new_state_dict[k] = v if hasattr(model, 'module'): model.module.load_state_dict(new_state_dict, strict=False) else: model.load_state_dict(new_state_dict, strict=False) - print("load ") logger.info("Loaded checkpoint '{}' (iteration {})".format( checkpoint_path, iteration)) return model, optimizer, learning_rate, iteration diff --git a/webui.py b/webui.py index 7a58eea..ae387bb 100644 --- a/webui.py +++ b/webui.py @@ -3,6 +3,17 @@ import sys, os if sys.platform == "darwin": os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1" +import logging + +logging.getLogger("numba").setLevel(logging.WARNING) +logging.getLogger("markdown_it").setLevel(logging.WARNING) +logging.getLogger("urllib3").setLevel(logging.WARNING) +logging.getLogger("matplotlib").setLevel(logging.WARNING) + +logging.basicConfig(level=logging.INFO, format="| %(name)s | %(levelname)s | %(message)s") + +logger = logging.getLogger(__name__) + import torch import argparse import commons @@ -67,8 +78,12 @@ if __name__ == "__main__": parser.add_argument("-m", "--model", default="./logs/as/G_8000.pth", help="path of your model") parser.add_argument("-c", "--config", default="./configs/config.json", help="path of your config file") parser.add_argument("--share", default=False, help="make link public") + parser.add_argument("-d", "--debug", action="store_true", help="enable DEBUG-LEVEL log") args = parser.parse_args() + if args.debug: + logger.info("Enable DEBUG-LEVEL log") + logging.basicConfig(level=logging.DEBUG) hps = utils.get_hparams_from_file(args.config) device = ( @@ -92,8 +107,7 @@ if __name__ == "__main__": speaker_ids = hps.data.spk2id speakers = list(speaker_ids.keys()) - app = gr.Blocks() - with app: + with gr.Blocks() as app: with gr.Row(): with gr.Column(): text = gr.TextArea(label="Text", placeholder="Input Text Here", From a4f0b57ae856c732d65435181bbc63edf6286537 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=BA=90=E6=96=87=E9=9B=A8?= <41315874+fumiama@users.noreply.github.com> Date: Mon, 4 Sep 2023 01:27:16 +0800 Subject: [PATCH 2/3] feat: add mps to chinese bert --- text/chinese_bert.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/text/chinese_bert.py b/text/chinese_bert.py index cb84ce0..45e5e28 100644 --- a/text/chinese_bert.py +++ b/text/chinese_bert.py @@ -1,7 +1,16 @@ import torch +import sys from transformers import AutoTokenizer, AutoModelForMaskedLM -device = torch.device("cuda" if torch.cuda.is_available() else "cpu") +device = torch.device( + "cuda" + if torch.cuda.is_available() + else ( + "mps" + if sys.platform == "darwin" and torch.backends.mps.is_available() + else "cpu" + ) + ) tokenizer = AutoTokenizer.from_pretrained("./bert/chinese-roberta-wwm-ext-large") model = AutoModelForMaskedLM.from_pretrained("./bert/chinese-roberta-wwm-ext-large").to(device) From 781b16b18133d4f3050274859dd59e14ef3b687c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=BA=90=E6=96=87=E9=9B=A8?= <41315874+fumiama@users.noreply.github.com> Date: Mon, 4 Sep 2023 02:04:41 +0800 Subject: [PATCH 3/3] fix: padding --- webui.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/webui.py b/webui.py index ae387bb..2e6576f 100644 --- a/webui.py +++ b/webui.py @@ -103,7 +103,7 @@ if __name__ == "__main__": **hps.model).to(device) _ = net_g.eval() - _ = utils.load_checkpoint(args.model, net_g, None,skip_optimizer=True) + _ = utils.load_checkpoint(args.model, net_g, None, skip_optimizer=True) speaker_ids = hps.data.spk2id speakers = list(speaker_ids.keys())