| 22 | |
| 23 | |
| 24 | def parse_args(args): |
| 25 | parser = argparse.ArgumentParser(description="LISA Model Training") |
| 26 | parser.add_argument("--local_rank", default=0, type=int, help="node rank") |
| 27 | parser.add_argument( |
| 28 | "--version", default="liuhaotian/llava-llama-2-13b-chat-lightning-preview" |
| 29 | ) |
| 30 | parser.add_argument("--vis_save_path", default="./vis_output", type=str) |
| 31 | parser.add_argument( |
| 32 | "--precision", |
| 33 | default="bf16", |
| 34 | type=str, |
| 35 | choices=["fp32", "bf16", "fp16"], |
| 36 | help="precision for inference", |
| 37 | ) |
| 38 | parser.add_argument("--image_size", default=1024, type=int, help="image size") |
| 39 | parser.add_argument("--model_max_length", default=512, type=int) |
| 40 | parser.add_argument("--lora_r", default=8, type=int) |
| 41 | parser.add_argument( |
| 42 | "--vision-tower", default="openai/clip-vit-large-patch14", type=str |
| 43 | ) |
| 44 | parser.add_argument("--load_in_8bit", action="store_true", default=False) |
| 45 | parser.add_argument("--load_in_4bit", action="store_true", default=False) |
| 46 | |
| 47 | parser.add_argument( |
| 48 | "--dataset", default="sem_seg||refer_seg||vqa||reason_seg", type=str |
| 49 | ) |
| 50 | parser.add_argument("--sample_rates", default="9,3,3,1", type=str) |
| 51 | parser.add_argument( |
| 52 | "--sem_seg_data", |
| 53 | default="ade20k||cocostuff||pascal_part||paco_lvis||mapillary", |
| 54 | type=str, |
| 55 | ) |
| 56 | parser.add_argument( |
| 57 | "--refer_seg_data", default="refclef||refcoco||refcoco+||refcocog", type=str |
| 58 | ) |
| 59 | parser.add_argument("--vqa_data", default="llava_instruct_150k", type=str) |
| 60 | parser.add_argument("--reason_seg_data", default="ReasonSeg|train", type=str) |
| 61 | parser.add_argument("--val_dataset", default="ReasonSeg|val", type=str) |
| 62 | parser.add_argument("--dataset_dir", default="./dataset", type=str) |
| 63 | parser.add_argument("--log_base_dir", default="./runs", type=str) |
| 64 | parser.add_argument("--exp_name", default="lisa", type=str) |
| 65 | parser.add_argument("--epochs", default=10, type=int) |
| 66 | parser.add_argument("--steps_per_epoch", default=500, type=int) |
| 67 | parser.add_argument( |
| 68 | "--batch_size", default=2, type=int, help="batch size per device per step" |
| 69 | ) |
| 70 | parser.add_argument( |
| 71 | "--grad_accumulation_steps", |
| 72 | default=10, |
| 73 | type=int, |
| 74 | ) |
| 75 | parser.add_argument("--val_batch_size", default=1, type=int) |
| 76 | parser.add_argument("--workers", default=4, type=int) |
| 77 | parser.add_argument("--lr", default=0.0003, type=float) |
| 78 | parser.add_argument("--ce_loss_weight", default=1.0, type=float) |
| 79 | parser.add_argument("--dice_loss_weight", default=0.5, type=float) |
| 80 | parser.add_argument("--bce_loss_weight", default=2.0, type=float) |
| 81 | parser.add_argument("--lora_alpha", default=16, type=int) |