a7medbm/MiniGPT4-video
0
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 wandb34import torch.distributed as dist35 36def parse_args():37 parser = argparse.ArgumentParser(description="Training",add_help=False)38 39 parser.add_argument("--cfg-path", required=True, help="path to configuration file.")40 parser.add_argument(41 "--options",42 nargs="+"43 )44 parser.add_argument("--job_name",default="minigpt_spatial_coco_control",type=str)45 # distributed training parameters46 parser.add_argument('--world_size', default=1, type=int,47 help='number of distributed processes')48 parser.add_argument('--local_rank', default=-1, type=int)49 parser.add_argument('--dist_on_itp', action='store_true')50 parser.add_argument('--dist_url', default='env://',51 help='url used to set up distributed training')52 53 # args = parser.parse_args()54 55 56 57 58 return parser59 60 61def setup_seeds(config):62 seed = config.run_cfg.seed + get_rank()63 64 random.seed(seed)65 np.random.seed(seed)66 torch.manual_seed(seed)67 68 cudnn.benchmark = False69 cudnn.deterministic = True70 71 72def get_runner_class(cfg):73 """74 Get runner class from config. Default to epoch-based runner.75 """76 runner_cls = registry.get_runner_class(cfg.run_cfg.get("runner", "runner_base"))77 78 return runner_cls79 80 81def main():82 # allow auto-dl completes on main process without timeout when using NCCL backend.83 # os.environ["NCCL_BLOCKING_WAIT"] = "1"84 85 # set before init_distributed_mode() to ensure the same job_id shared across all ranks.86 87 print("start!!!")88 job_id = now()89 args = parse_args().parse_args()90 91 92 print("0000")93 cfg = Config(args)94 95 if 'LOCAL_RANK' not in os.environ:96 print("not in the os")97 os.environ['LOCAL_RANK'] = str(args.local_rank)98 print("111")99 local_rank = int(os.environ.get('LOCAL_RANK', 0))100 torch.cuda.set_device(local_rank)101 102 print("local rank",local_rank)103 104 dist.init_process_group(backend='nccl', init_method='env://')105 106 num_nodes = dist.get_world_size()107 print(f"Number of nodes: {num_nodes}")108 109 110 init_distributed_mode(cfg.run_cfg)111 112 setup_seeds(cfg)113 114 # set after in115 # it_distributed_mode() to only log on master.116 setup_logger()117 118 119 wandb.login()120 # print(wandb.run)121 122 123 cfg.pretty_print()124 125 task = tasks.setup_task(cfg)126 datasets = task.build_datasets(cfg)127 model = task.build_model(cfg)128 if cfg.run_cfg.rank == 0:129 print("project name", args.job_name)130 131 wandb.init(project="minigpt4-spatial",name=args.job_name)132 133 wandb.config = {"learning_rate": 0.0001, "epochs": 100, "batch_size": 8}134 wandb.watch(model)135 136 # print('+++++++++++++++++')137 # print(type(model))138 # print('+++++++++++++++++')139 # print(model)140 # print('+++++++++++++++++')141 # print(model.super().device)142 # print('+++++++++++++++++')143 # print(model.device)144 145 runner = get_runner_class(cfg)(146 cfg=cfg, job_id=job_id, task=task, model=model, datasets=datasets147 )148 runner.train()149 150 151if __name__ == "__main__":152 main()153 