()
| 3 | |
| 4 | |
| 5 | def get_argument_parser(): |
| 6 | |
| 7 | parser = argparse.ArgumentParser() |
| 8 | |
| 9 | parser.add_argument('--data_path', default='data/', type=str, help='path to datasets') |
| 10 | parser.add_argument('--dataset', default='f30k', help='dataset coco or f30k') |
| 11 | |
| 12 | parser.add_argument('--margin', default=0.2, type=float, help='Rank loss margin.') |
| 13 | parser.add_argument('--num_epochs', default=30, type=int, help='Number of training epochs.') |
| 14 | parser.add_argument('--batch_size', default=128, type=int, help='Size of a training mini-batch.') |
| 15 | parser.add_argument('--embed_size', default=512, type=int, help='Dimensionality of the joint embedding.') |
| 16 | |
| 17 | parser.add_argument('--grad_clip', default=2., type=float, help='Gradient clipping threshold.') |
| 18 | parser.add_argument('--learning_rate', default=2e-4, type=float, help='Initial learning rate.') |
| 19 | |
| 20 | parser.add_argument('--workers', default=8, type=int, help='Number of data loader workers.') |
| 21 | parser.add_argument('--log_step', default=200, type=int, help='Number of steps to logger.info and record the log.') |
| 22 | parser.add_argument('--val_step', default=500, type=int, help='Number of steps to run validation.') |
| 23 | |
| 24 | parser.add_argument('--logger_name', default='runs/test', help='Path to save Tensorboard log.') |
| 25 | |
| 26 | parser.add_argument('--max_violation', action='store_true', help='Use max instead of sum in the rank loss.') |
| 27 | parser.add_argument('--vse_mean_warmup_epochs', type=int, default=1, help='The number of warmup epochs using mean vse loss') |
| 28 | parser.add_argument('--embedding_warmup_epochs', type=int, default=0, help='The number of epochs for warming up the embedding layer') |
| 29 | |
| 30 | parser.add_argument('--f30k_img_path', type=str, default='/home/fzr/data/flickr30k-images', help='the path of f30k images') |
| 31 | parser.add_argument('--coco_img_path', type=str, default='/home/fzr/data/coco/', help='the path of coco images') |
| 32 | |
| 33 | # vision transformer |
| 34 | parser.add_argument('--img_res', type=int, default=224, help='the image resolution for ViT input') |
| 35 | parser.add_argument('--vit_type', type=str, default='vit', help='the type of vit model') |
| 36 | |
| 37 | # use DDP for training |
| 38 | parser.add_argument('--multi_gpu', type=int, default=0, help='whether use multi-gpu for training') |
| 39 | parser.add_argument('--world_size', type=int, default=1, help='number of distributed processes') |
| 40 | parser.add_argument("--rank", type=int, default=0, help='the parameter for rank') |
| 41 | parser.add_argument("--local_rank", type=int, default=0, help='the parameter for local rank') |
| 42 | |
| 43 | parser.add_argument('--dist_backend', type=str, default='nccl', help='the backend for ddp') |
| 44 | parser.add_argument('--dist_url', type=str, default='env://', help='url used to set up distributed training') |
| 45 | parser.add_argument('--seed', type=int, default=0, help='fix the seed for reproducibility') |
| 46 | |
| 47 | # others |
| 48 | parser.add_argument('--size_augment', type=int, default=1, help='whether use the Size Augmentation') |
| 49 | parser.add_argument('--loss', type=str, default='vse', help='the objectve function for optimization') |
| 50 | parser.add_argument('--eval', type=int, default=1, help='whether evaluation after training process') |
| 51 | |
| 52 | parser.add_argument('--save_results', type=int, default=1, help='whether save the evaluation results') |
| 53 | parser.add_argument('--evaluate_cxc', type=int, default=0, help='the special evaluation for MS-COCO') |
| 54 | parser.add_argument('--gpu-id', type=int, default=0, help='the gpu-id for runing') |
| 55 | |
| 56 | parser.add_argument('--bert_path', type=str, default='../weights_models/bert-base-uncased') |
| 57 | |
| 58 | # optimizer |
| 59 | parser.add_argument("--lr_schedules", default=[9, 15, 20, 25], type=int, nargs="+", help='epoch schedules for lr decay') |
| 60 | parser.add_argument("--decay_rate", default=0.3, type=float, help='lr decay_rate for optimizer') |
| 61 | |
| 62 | parser.add_argument('--shard_size', type=int, default=256, help='the shard_size for cross-attention') |
nothing calls this directly
no outgoing calls
no test coverage detected