fffiloni/miniGPT4-Video-Zero
22
1"""2 Copyright (c) 2022, salesforce.com, inc.3 All rights reserved.4 SPDX-License-Identifier: BSD-3-Clause5 For full license text, see the LICENSE_Lavis file in the repo root or https://opensource.org/licenses/BSD-3-Clause6"""7 8import argparse9import os10import random11 12import numpy as np13import torch14import torch.backends.cudnn as cudnn15 16import minigpt4.tasks as tasks17from minigpt4.common.config import Config18from minigpt4.common.dist_utils import get_rank, init_distributed_mode19from minigpt4.common.logger import setup_logger20from minigpt4.common.optims import (21 LinearWarmupCosineLRScheduler,22 LinearWarmupStepLRScheduler,23)24from minigpt4.common.registry import registry25from minigpt4.common.utils import now26 27# imports modules for registration28from minigpt4.datasets.builders import *29from minigpt4.models import *30from minigpt4.processors import *31from minigpt4.runners import *32from minigpt4.tasks import *33import wandb34 35 36def parse_args():37 parser = argparse.ArgumentParser(description="Training")38 39 parser.add_argument("--cfg-path",default="train_configs_llama2/224_v2_llama2_video.yaml", required=False, help="path to configuration file.")40 parser.add_argument(41 "--options",42 nargs="+",43 help="override some settings in the used config, the key-value pair "44 "in xxx=yyy format will be merged into config file (deprecate), "45 "change to --cfg-options instead.",46 )47 parser.add_argument("--job_name",default="test",type=str)48 args = parser.parse_args()49 50 return args51 52 53def setup_seeds(config):54 seed = config.run_cfg.seed + get_rank()55 56 random.seed(seed)57 np.random.seed(seed)58 torch.manual_seed(seed)59 60 cudnn.benchmark = False61 cudnn.deterministic = True62 63 64def get_runner_class(cfg):65 """66 Get runner class from config. Default to epoch-based runner.67 """68 runner_cls = registry.get_runner_class(cfg.run_cfg.get("runner", "runner_base"))69 70 return runner_cls71 72 73def setup_environ_flags(rank):74 """Set environment flags for debugging purposes"""75 os.environ["TORCH_SHOW_CPP_STACKTRACES"] = str(1)76 os.environ["NCCL_ASYNC_ERROR_HANDLING"] = str(1)77 os.environ["TORCH_DISTRIBUTED_DEBUG"] = "DETAIL"78 if rank == 0:79 print(f"--> Running with torch dist debug set to detail")80 81 82def main():83 # allow auto-dl completes on main process without timeout when using NCCL backend.84 # os.environ["NCCL_BLOCKING_WAIT"] = "1"85 86 # set before init_distributed_mode() to ensure the same job_id shared across all ranks.87 setup_environ_flags(get_rank())88 job_id = now()89 args = parse_args()90 cfg = Config(args)91 init_distributed_mode(cfg.run_cfg)92 setup_seeds(cfg)93 94 # set after in95 # it_distributed_mode() to only log on master.96 setup_logger()97 wandb.login()98 # print(wandb.run)99 cfg.pretty_print()100 101 task = tasks.setup_task(cfg)102 datasets = task.build_datasets(cfg)103 model = task.build_model(cfg)104 if not hasattr(cfg.run_cfg, 'rank') or cfg.run_cfg.rank == 0:105 print("project name", args.job_name)106 107 wandb.init(project="minigpt4-spatial",name=args.job_name)108 109 wandb.config = {"learning_rate": 0.0001, "epochs": 100, "batch_size": 8}110 wandb.watch(model)111 112 # print('+++++++++++++++++')113 # print(type(model))114 # print('+++++++++++++++++')115 # print(model)116 # print('+++++++++++++++++')117 # print(model.super().device)118 # print('+++++++++++++++++')119 # print(model.device)120 121 runner = get_runner_class(cfg)(122 cfg=cfg, job_id=job_id, task=task, model=model, datasets=datasets123 )124 runner.train()125 126 127if __name__ == "__main__":128 main()129 