CoolFace
Apppublic

fffiloni/miniGPT4-Video-Zero

sourceHugging Faceupdated 1y agoView on Hugging Face
22likes
train.py129 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 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 
fffiloni/miniGPT4-Video-Zero · CoolFace