CoolFace
Apppublic

a7medbm/MiniGPT4-video

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
train_multinode.py153 linesDownload Raw Back to root
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