MLVLab/Human_Object_Interaction
1
1# ------------------------------------------------------------------------2# HOTR official code : engine/arg_parser.py3# Copyright (c) Kakao Brain, Inc. and its affiliates. All Rights Reserved4# Modified arguments are represented with *5# ------------------------------------------------------------------------6# Modified from DETR (https://github.com/facebookresearch/detr)7# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved8# ------------------------------------------------------------------------9import argparse10import hotr.util.misc as utils11 12def get_args_parser():13 parser = argparse.ArgumentParser('Set transformer detector', add_help=False)14 parser.add_argument('--lr', default=1e-4, type=float)15 parser.add_argument('--lr_backbone', default=1e-5, type=float)16 parser.add_argument('--batch_size', default=2, type=int)17 parser.add_argument('--weight_decay', default=1e-4, type=float)18 parser.add_argument('--epochs', default=100, type=int)19 parser.add_argument('--lr_drop', default=80, type=int)20 parser.add_argument('--clip_max_norm', default=0.1, type=float,21 help='gradient clipping max norm')22 23 # DETR Model parameters24 parser.add_argument('--frozen_weights', type=str, default=None,25 help="Path to the pretrained model. If set, only the mask head will be trained")26 parser.add_argument('--pretrain_interaction_tf', type=str, default=None,27 help="Path to the pretrained model. If set, only the mask head will be trained")28 29 # DETR Backbone30 parser.add_argument('--backbone', default='resnet50', type=str,31 help="Name of the convolutional backbone to use")32 parser.add_argument('--dilation', action='store_true',33 help="If true, we replace stride with dilation in the last convolutional block (DC5)")34 parser.add_argument('--position_embedding', default='sine', type=str, choices=('sine', 'learned'),35 help="Type of positional embedding to use on top of the image features")36 37 # DETR Transformer (= Encoder, Instance Decoder)38 parser.add_argument('--enc_layers', default=6, type=int,39 help="Number of encoding layers in the transformer")40 parser.add_argument('--dec_layers', default=6, type=int,41 help="Number of decoding layers in the transformer")42 parser.add_argument('--dim_feedforward', default=2048, type=int,43 help="Intermediate size of the feedforward layers in the transformer blocks")44 parser.add_argument('--hidden_dim', default=256, type=int,45 help="Size of the embeddings (dimension of the transformer)")46 parser.add_argument('--dropout', default=0.1, type=float,47 help="Dropout applied in the transformer")48 parser.add_argument('--nheads', default=8, type=int,49 help="Number of attention heads inside the transformer's attentions")50 parser.add_argument('--num_queries', default=100, type=int,51 help="Number of query slots")52 parser.add_argument('--pre_norm', action='store_true')53 parser.add_argument('--decoder_form', default=2, type=int,54 help="1-decoder or 2-decoder")55 # Segmentation56 parser.add_argument('--masks', action='store_true',57 help="Train segmentation head if the flag is provided")58 59 # Loss Option60 parser.add_argument('--no_aux_loss', dest='aux_loss', action='store_false',61 help="Disables auxiliary decoding losses (loss at each layer)")62 63 # Loss coefficients (DETR)64 parser.add_argument('--mask_loss_coef', default=1, type=float)65 parser.add_argument('--dice_loss_coef', default=1, type=float)66 parser.add_argument('--bbox_loss_coef', default=5, type=float)67 parser.add_argument('--giou_loss_coef', default=2, type=float)68 parser.add_argument('--eos_coef', default=0.1, type=float,69 help="Relative classification weight of the no-object class")70 71 # Matcher (DETR)72 parser.add_argument('--set_cost_class', default=1, type=float,73 help="Class coefficient in the matching cost")74 parser.add_argument('--set_cost_bbox', default=5, type=float,75 help="L1 box coefficient in the matching cost")76 parser.add_argument('--set_cost_giou', default=2, type=float,77 help="giou box coefficient in the matching cost")78 79 # * HOI Detection80 parser.add_argument('--HOIDet', action='store_true',81 help="Train HOI Detection head if the flag is provided")82 parser.add_argument('--share_enc', action='store_true',83 help="Share the Encoder in DETR for HOI Detection if the flag is provided")84 parser.add_argument('--pretrained_dec', action='store_true',85 help="Use Pre-trained Decoder in DETR for Interaction Decoder if the flag is provided") 86 parser.add_argument('--hoi_enc_layers', default=1, type=int,87 help="Number of decoding layers in HOI transformer")88 parser.add_argument('--hoi_dec_layers', default=1, type=int,89 help="Number of decoding layers in HOI transformer")90 parser.add_argument('--hoi_nheads', default=8, type=int,91 help="Number of decoding layers in HOI transformer")92 parser.add_argument('--hoi_dim_feedforward', default=2048, type=int,93 help="Number of decoding layers in HOI transformer")94 # parser.add_argument('--hoi_mode', type=str, default=None, help='[inst | pair | all]')95 parser.add_argument('--num_hoi_queries', default=100, type=int,96 help="Number of Queries for Interaction Decoder")97 parser.add_argument('--hoi_aux_loss', action='store_true')98 99 100 # * HOTR Matcher101 parser.add_argument('--set_cost_idx', default=1, type=float,102 help="IDX coefficient in the matching cost")103 parser.add_argument('--set_cost_act', default=1, type=float,104 help="Action coefficient in the matching cost")105 parser.add_argument('--set_cost_tgt', default=1, type=float,106 help="Target coefficient in the matching cost")107 108 # * HOTR Loss coefficients109 parser.add_argument('--temperature', default=0.05, type=float, help="temperature")110 parser.add_argument('--hoi_consistency_loss_coef', default=1, type=float)111 parser.add_argument('--hoi_idx_loss_coef', default=1, type=float)112 parser.add_argument('--hoi_idx_consistency_loss_coef', default=1, type=float)113 parser.add_argument('--hoi_act_loss_coef', default=1, type=float)114 parser.add_argument('--hoi_act_consistency_loss_coef', default=1, type=float)115 parser.add_argument('--hoi_tgt_loss_coef', default=1, type=float)116 parser.add_argument('--hoi_tgt_consistency_loss_coef', default=1, type=float)117 parser.add_argument('--hoi_eos_coef', default=0.1, type=float, help="Relative classification weight of the no-object class")118 119 parser.add_argument('--ramp_down_epoch',default=10000,type=int)120 parser.add_argument('--ramp_up_epoch',default=0,type=int)121 #consistency122 parser.add_argument('--use_consis',action='store_true',help='use consistency regularization')123 parser.add_argument('--share_dec_param',action='store_true',help = 'share decoder parameters of all stages')124 parser.add_argument("--augpath_name", type=utils.arg_as_list,default=[],125 help='choose which augmented inference paths to use. (p2:x->HO->I,p3:x->HI->O,p4:x->OI->H)') 126 parser.add_argument('--stop_grad_stage',action='store_true',help='Do not back propogate loss to previous stage')127 parser.add_argument('--path_id', default=0, type=int)128 129 # * dataset parameters130 parser.add_argument('--dataset_file', help='[coco | vcoco]')131 parser.add_argument('--data_path', type=str)132 parser.add_argument('--object_threshold', type=float, default=0, help='Threshold for object confidence')133 134 # machine parameters135 parser.add_argument('--output_dir', default='',136 help='path where to save, empty for no saving')137 parser.add_argument('--custom_path', default='',138 help="Data path for custom inference. Only required for custom_main.py")139 parser.add_argument('--device', default='cuda',140 help='device to use for training / testing')141 parser.add_argument('--seed', default=42, type=int)142 parser.add_argument('--resume', default='', help='resume from checkpoint')143 parser.add_argument('--start_epoch', default=0, type=int, metavar='N',144 help='start epoch')145 parser.add_argument('--num_workers', default=2, type=int)146 147 # mode148 parser.add_argument('--eval', action='store_true', help="Only evaluate results if the flag is provided")149 parser.add_argument('--validate', action='store_true', help="Validate after every epoch")150 151 # distributed training parameters152 parser.add_argument('--world_size', default=1, type=int,153 help='number of distributed processes')154 parser.add_argument('--dist_url', default='env://', help='url used to set up distributed training')155 156 # * WanDB157 parser.add_argument('--wandb', action='store_true')158 parser.add_argument('--project_name', default='hotr_cpc')159 parser.add_argument('--group_name', default='mlv')160 parser.add_argument('--run_name', default='run_000001')161 return parser162 