Team Ai
Apppublic

MLVLab/Human_Object_Interaction

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
1likes
main.py243 linesDownload Raw Back to root
1# ------------------------------------------------------------------------2# HOTR official code : main.py3# Copyright (c) Kakao Brain, Inc. and its affiliates. All Rights Reserved4# ------------------------------------------------------------------------5# Modified from DETR (https://github.com/facebookresearch/detr)6# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved7# ------------------------------------------------------------------------8import argparse9import datetime10import json11import random12import time13import multiprocessing14from pathlib import Path15 16import numpy as np17import torch18from torch.utils.data import DataLoader, DistributedSampler19 20import hotr.data.datasets as datasets21import hotr.util.misc as utils22from hotr.engine.arg_parser import get_args_parser23from hotr.data.datasets import build_dataset, get_coco_api_from_dataset24from hotr.engine.trainer import train_one_epoch25from hotr.engine import hoi_evaluator, hoi_accumulator26from hotr.models import build_model27import wandb28 29from hotr.util.logger import print_params, print_args30 31def save_ckpt(args, model_without_ddp, optimizer, lr_scheduler, epoch, filename):32    # save_ckpt: function for saving checkpoints33    output_dir = Path(args.output_dir)34    if args.output_dir:35        checkpoint_path = output_dir / f'{filename}.pth'36        utils.save_on_master({37            'model': model_without_ddp.state_dict(),38            'optimizer': optimizer.state_dict(),39            'lr_scheduler': lr_scheduler.state_dict(),40            'epoch': epoch,41            'args': args,42        }, checkpoint_path)43 44def main(args):45    utils.init_distributed_mode(args)46 47    if args.frozen_weights is not None:48        print("Freeze weights for detector")49 50    if not torch.cuda.is_available():51        args.device = 'cpu'52    device = torch.device(args.device)53 54    # fix the seed for reproducibility55    seed = args.seed + utils.get_rank()56    torch.manual_seed(seed)57    np.random.seed(seed)58    random.seed(seed)59 60    # Data Setup61    dataset_train = build_dataset(image_set='train', args=args)62    dataset_val = build_dataset(image_set='val' if not args.eval else 'test', args=args)63    assert dataset_train.num_action() == dataset_val.num_action(), "Number of actions should be the same between splits"64    args.num_classes = dataset_train.num_category()65    args.num_actions = dataset_train.num_action()66    args.action_names = dataset_train.get_actions()67    if args.share_enc: args.hoi_enc_layers = args.enc_layers68    if args.pretrained_dec: args.hoi_dec_layers = args.dec_layers69    if args.dataset_file == 'vcoco':70        # Save V-COCO dataset statistics71        args.valid_ids = np.array(dataset_train.get_object_label_idx()).nonzero()[0]72        args.invalid_ids = np.argwhere(np.array(dataset_train.get_object_label_idx()) == 0).squeeze(1)73        args.human_actions = dataset_train.get_human_action()74        args.object_actions = dataset_train.get_object_action()75        args.num_human_act = dataset_train.num_human_act()76    elif args.dataset_file == 'hico-det':77        args.valid_obj_ids = dataset_train.get_valid_obj_ids()78    print_args(args)79 80    if args.distributed:81        sampler_train = DistributedSampler(dataset_train, shuffle=True)82        sampler_val = DistributedSampler(dataset_val, shuffle=False)83    else:84        sampler_train = torch.utils.data.RandomSampler(dataset_train)85        sampler_val = torch.utils.data.SequentialSampler(dataset_val)86 87    batch_sampler_train = torch.utils.data.BatchSampler(88        sampler_train, args.batch_size, drop_last=True)89 90    data_loader_train = DataLoader(dataset_train, batch_sampler=batch_sampler_train,91                                  collate_fn=utils.collate_fn, num_workers=args.num_workers)92    data_loader_val = DataLoader(dataset_val, args.batch_size, sampler=sampler_val,93                                 drop_last=False, collate_fn=utils.collate_fn, num_workers=args.num_workers)94 95    # Model Setup96    model, criterion, postprocessors = build_model(args)97    # import pdb;pdb.set_trace()98    model.to(device)99    100    model_without_ddp = model101    if args.distributed:102        model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[args.gpu])103        model_without_ddp = model.module104    n_parameters = print_params(model)105 106    param_dicts = [107        {"params": [p for n, p in model_without_ddp.named_parameters() if "backbone" not in n and p.requires_grad]},108        {109            "params": [p for n, p in model_without_ddp.named_parameters() if "backbone" in n and p.requires_grad],110            "lr": args.lr_backbone,111        },112    ]113    optimizer = torch.optim.AdamW(param_dicts, lr=args.lr, weight_decay=args.weight_decay)114    lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, args.lr_drop)115    # lr_scheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, [1,100])116    117 118    # Weight Setup119    if args.frozen_weights is not None:120        if args.frozen_weights.startswith('https'):121            checkpoint = torch.hub.load_state_dict_from_url(122                args.frozen_weights, map_location='cpu', check_hash=True)123        else:124            checkpoint = torch.load(args.frozen_weights, map_location='cpu')125        model_without_ddp.detr.load_state_dict(checkpoint['model'])126    127    if args.resume:128        if args.resume.startswith('https'):129            checkpoint = torch.hub.load_state_dict_from_url(130                args.resume, map_location='cpu', check_hash=True)131        else:132            checkpoint = torch.load(args.resume, map_location='cpu')133        model_without_ddp.load_state_dict(checkpoint['model'])134        if not args.eval and 'optimizer' in checkpoint and 'lr_scheduler' in checkpoint and 'epoch' in checkpoint:135            optimizer.load_state_dict(checkpoint['optimizer'])136            # lr_scheduler.load_state_dict(checkpoint['lr_scheduler'])137            args.start_epoch = checkpoint['epoch'] + 1138    # import pdb;pdb.set_trace()139    if args.eval:140        # test only mode141        if args.HOIDet:142            if args.dataset_file == 'vcoco':143                total_res = hoi_evaluator(args, model, criterion, postprocessors, data_loader_val, device)144                sc1, sc2 = hoi_accumulator(args, total_res, True, False)145            elif args.dataset_file == 'hico-det':146                test_stats = hoi_evaluator(args, model, None, postprocessors, data_loader_val, device)147                print(f'| mAP (full)\t\t: {test_stats["mAP"]:.2f}')148                print(f'| mAP (rare)\t\t: {test_stats["mAP rare"]:.2f}')149                print(f'| mAP (non-rare)\t: {test_stats["mAP non-rare"]:.2f}')150            else: raise ValueError(f'dataset {args.dataset_file} is not supported.')151            return152        else:153            test_stats, coco_evaluator = evaluate_coco(model, criterion, postprocessors,154                                                  data_loader_val, base_ds, device, args.output_dir)155            if args.output_dir:156                utils.save_on_master(coco_evaluator.coco_eval["bbox"].eval, output_dir / "eval.pth")157            return158 159    # stats160    scenario1, scenario2 = 0, 0161    best_mAP, best_rare, best_non_rare = 0, 0, 0162 163    # add argparse164    if args.wandb and utils.get_rank() == 0:165        wandb.init(166            project=args.project_name,167            group=args.group_name,168            name=args.run_name,169            config=args170        )171        wandb.watch(model)172 173    # Training starts here!174    # lr_scheduler.step()175    start_time = time.time()176    for epoch in range(args.start_epoch, args.epochs):177        if args.distributed:178            sampler_train.set_epoch(epoch)179        train_stats = train_one_epoch(180            model, criterion, data_loader_train, optimizer, device, epoch, args.epochs, args.ramp_up_epoch,args.ramp_down_epoch,args.hoi_consistency_loss_coef,181            args.clip_max_norm, dataset_file=args.dataset_file, log=args.wandb)182        lr_scheduler.step()183 184        # Validation185        if args.validate:186            print('-'*100)187            if args.dataset_file == 'vcoco':188                total_res = hoi_evaluator(args, model, criterion, postprocessors, data_loader_val, device)189                if utils.get_rank() == 0:190                    sc1, sc2 = hoi_accumulator(args, total_res, False, args.wandb)191                    if sc1 > scenario1:192                        scenario1 = sc1193                        scenario2 = sc2194                        save_ckpt(args, model_without_ddp, optimizer, lr_scheduler, epoch, filename='best')195                    print(f'| Scenario #1 mAP : {sc1:.2f} ({scenario1:.2f})')196                    print(f'| Scenario #2 mAP : {sc2:.2f} ({scenario2:.2f})')197            elif args.dataset_file == 'hico-det':198                test_stats = hoi_evaluator(args, model, None, postprocessors, data_loader_val, device)199                if utils.get_rank() == 0:200                    if test_stats['mAP'] > best_mAP:201                        best_mAP = test_stats['mAP']202                        best_rare = test_stats['mAP rare']203                        best_non_rare = test_stats['mAP non-rare']204                        save_ckpt(args, model_without_ddp, optimizer, lr_scheduler, epoch, filename='best')205                    print(f'| mAP (full)\t\t: {test_stats["mAP"]:.2f} ({best_mAP:.2f})')206                    print(f'| mAP (rare)\t\t: {test_stats["mAP rare"]:.2f} ({best_rare:.2f})')207                    print(f'| mAP (non-rare)\t: {test_stats["mAP non-rare"]:.2f} ({best_non_rare:.2f})')208                    if args.wandb and utils.get_rank() == 0:209                        wandb.log({210                            'mAP': test_stats['mAP'],211                            'mAP rare': test_stats['mAP rare'],212                            'mAP non-rare': test_stats['mAP non-rare']213                        })214            print('-'*100)215        216        save_ckpt(args, model_without_ddp, optimizer, lr_scheduler, epoch, filename='checkpoint')217        if (epoch + 1) % args.lr_drop == 0 :218            save_ckpt(args, model_without_ddp, optimizer, lr_scheduler, epoch, filename='checkpoint_'+str(epoch))219        # if (epoch + 1) % args.pseudo_epoch == 0 :220        #     save_ckpt(args, model_without_ddp, optimizer, lr_scheduler, epoch, filename='checkpoint_pseudo_'+str(epoch))221    total_time = time.time() - start_time222    total_time_str = str(datetime.timedelta(seconds=int(total_time)))223    print('Training time {}'.format(total_time_str))224    if args.dataset_file == 'vcoco':225        print(f'| Scenario #1 mAP : {scenario1:.2f}')226        print(f'| Scenario #2 mAP : {scenario2:.2f}')227    elif args.dataset_file == 'hico-det':228        print(f'| mAP (full)\t\t: {best_mAP:.2f}')229        print(f'| mAP (rare)\t\t: {best_rare:.2f}')230        print(f'| mAP (non-rare)\t: {best_non_rare:.2f}')231 232 233if __name__ == '__main__':234    parser = argparse.ArgumentParser(235        'End-to-End Human Object Interaction training and evaluation script',236        parents=[get_args_parser()]237    )238    args = parser.parse_args()239    if args.output_dir:240        args.output_dir += f"/{args.group_name}/{args.run_name}/"241        Path(args.output_dir).mkdir(parents=True, exist_ok=True)242    main(args)243