mirror of
https://github.com/ostris/ai-toolkit.git
synced 2026-01-26 16:39:47 +00:00
Allow name to be passed through command line
This commit is contained in:
10
run.py
10
run.py
@@ -36,6 +36,14 @@ def main():
|
||||
action='store_true',
|
||||
help='Continue running additional jobs even if a job fails'
|
||||
)
|
||||
|
||||
# flag to continue if failed job
|
||||
parser.add_argument(
|
||||
'-n', '--name',
|
||||
type=str,
|
||||
default=None,
|
||||
help='Name to replace [name] tag in config file, useful for shared config file'
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
config_file_list = args.config_file_list
|
||||
@@ -49,7 +57,7 @@ def main():
|
||||
|
||||
for config_file in config_file_list:
|
||||
try:
|
||||
job = get_job(config_file)
|
||||
job = get_job(config_file, args.name)
|
||||
job.run()
|
||||
job.cleanup()
|
||||
jobs_completed += 1
|
||||
|
||||
@@ -15,22 +15,24 @@ def get_cwd_abs_path(path):
|
||||
return path
|
||||
|
||||
|
||||
def preprocess_config(config: OrderedDict):
|
||||
def preprocess_config(config: OrderedDict, name: str = None):
|
||||
if "job" not in config:
|
||||
raise ValueError("config file must have a job key")
|
||||
if "config" not in config:
|
||||
raise ValueError("config file must have a config section")
|
||||
if "name" not in config["config"]:
|
||||
if "name" not in config["config"] and name is None:
|
||||
raise ValueError("config file must have a config.name key")
|
||||
# we need to replace tags. For now just [name]
|
||||
name = config["config"]["name"]
|
||||
if name is not None:
|
||||
config["config"]["name"] = name
|
||||
else:
|
||||
name = config["config"]["name"]
|
||||
config_string = json.dumps(config)
|
||||
config_string = config_string.replace("[name]", name)
|
||||
config = json.loads(config_string, object_pairs_hook=OrderedDict)
|
||||
return config
|
||||
|
||||
|
||||
|
||||
# Fixes issue where yaml doesnt load exponents correctly
|
||||
fixed_loader = yaml.SafeLoader
|
||||
fixed_loader.add_implicit_resolver(
|
||||
@@ -44,7 +46,8 @@ fixed_loader.add_implicit_resolver(
|
||||
|\\.(?:nan|NaN|NAN))$''', re.X),
|
||||
list(u'-+0123456789.'))
|
||||
|
||||
def get_config(config_file_path):
|
||||
|
||||
def get_config(config_file_path, name=None):
|
||||
# first check if it is in the config folder
|
||||
config_path = os.path.join(TOOLKIT_ROOT, 'config', config_file_path)
|
||||
# see if it is in the config folder with any of the possible extensions if it doesnt have one
|
||||
@@ -75,4 +78,4 @@ def get_config(config_file_path):
|
||||
else:
|
||||
raise ValueError(f"Config file {config_file_path} must be a json or yaml file")
|
||||
|
||||
return preprocess_config(config)
|
||||
return preprocess_config(config, name)
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
from toolkit.config import get_config
|
||||
|
||||
|
||||
def get_job(config_path):
|
||||
config = get_config(config_path)
|
||||
def get_job(config_path, name=None):
|
||||
config = get_config(config_path, name)
|
||||
if not config['job']:
|
||||
raise ValueError('config file is invalid. Missing "job" key')
|
||||
|
||||
|
||||
Reference in New Issue
Block a user