| import sys |
|
|
| sys.path.append('.') |
|
|
| import yaml |
| import argparse |
|
|
| from easydict import EasyDict |
| from scripts.utils.others import setup_seed |
| from scripts.utils.module_loader import * |
|
|
|
|
| def run(config): |
| |
| model = load_model(config.model) |
|
|
| |
| |
| |
|
|
| |
| data_module = load_dataset(config.dataset) |
|
|
| |
| trainer = load_trainer(config) |
|
|
| |
| trainer.fit(model=model, datamodule=data_module) |
|
|
| |
| if model.save_path is not None: |
| if config.model.kwargs.get("use_lora", False): |
| |
| config.model.kwargs.lora_config_path = model.save_path |
| model = load_model(config.model) |
|
|
| else: |
| model.load_checkpoint(model.save_path, load_prev_scheduler=model.load_prev_scheduler) |
|
|
| trainer.test(model=model, datamodule=data_module) |
|
|
|
|
| def get_args(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument('-c', '--config', help="running configurations", type=str, required=True) |
| return parser.parse_args() |
|
|
|
|
| def main(args): |
| with open(args.config, 'r', encoding='utf-8') as r: |
| config = EasyDict(yaml.safe_load(r)) |
|
|
| if config.setting.seed: |
| setup_seed(config.setting.seed) |
|
|
| |
| for k, v in config.setting.os_environ.items(): |
| if v is not None and k not in os.environ: |
| os.environ[k] = str(v) |
|
|
| elif k in os.environ: |
| |
| config.setting.os_environ[k] = os.environ[k] |
|
|
| |
| if config.setting.os_environ.NODE_RANK != 0: |
| config.Trainer.logger = False |
|
|
| run(config) |
|
|
|
|
| if __name__ == '__main__': |
| main(get_args()) |