mirror of
https://github.com/wassname/TTS.git
synced 2026-10-02 12:00:11 +08:00
config refactor #4 WIP
This commit is contained in:
1 parent
97bd5f9734
commit
dc50f5f0b0
6 files changed
+229
-267
No files matched your search
+18
-29
@@ -1,6 +1,10 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
import datetime
|
||||
import glob
|
||||
import importlib
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
@@ -67,6 +71,20 @@ def count_parameters(model):
|
||||
return sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
|
||||
|
||||
def to_camel(text):
|
||||
text = text.capitalize()
|
||||
text = re.sub(r'(?!^)_([a-zA-Z])', lambda m: m.group(1).upper(), text)
|
||||
text = text.replace('Tts', 'TTS')
|
||||
return text
|
||||
|
||||
|
||||
def find_module(module_path: str, module_name: str) -> object:
|
||||
module_name = module_name.lower()
|
||||
module = importlib.import_module(module_path+'.'+module_name)
|
||||
class_name = to_camel(module_name)
|
||||
return getattr(module, class_name)
|
||||
|
||||
|
||||
def get_user_data_dir(appname):
|
||||
if sys.platform == "win32":
|
||||
import winreg # pylint: disable=import-outside-toplevel
|
||||
@@ -139,32 +157,3 @@ class KeepAverage:
|
||||
for key, value in value_dict.items():
|
||||
self.update_value(key, value)
|
||||
|
||||
|
||||
def check_argument(name,
|
||||
c,
|
||||
prerequest=None,
|
||||
enum_list=None,
|
||||
max_val=None,
|
||||
min_val=None,
|
||||
restricted=False,
|
||||
alternative=None,
|
||||
allow_none=False):
|
||||
if isinstance(prerequest, List()):
|
||||
if any([f not in c.keys() for f in prerequest]):
|
||||
return
|
||||
else:
|
||||
if prerequest not in c.keys():
|
||||
return
|
||||
if alternative in c.keys() and c[alternative] is not None:
|
||||
return
|
||||
if allow_none and c[name] is None:
|
||||
return
|
||||
if restricted:
|
||||
assert name in c.keys(), f" [!] {name} not defined in config.json"
|
||||
if name in c.keys():
|
||||
if max_val:
|
||||
assert c[name] <= max_val, f" [!] {name} is larger than max value {max_val}"
|
||||
if min_val:
|
||||
assert c[name] >= min_val, f" [!] {name} is smaller than min value {min_val}"
|
||||
if enum_list:
|
||||
assert c[name].lower() in enum_list, f' [!] {name} is not a valid value'
|
||||
Reference in new issue
Block a user