(attn_implementation=None)
| 874 | |
| 875 | |
| 876 | def train(attn_implementation=None): |
| 877 | global local_rank |
| 878 | |
| 879 | parser = transformers.HfArgumentParser( |
| 880 | (ModelArguments, DataArguments, TrainingArguments)) |
| 881 | model_args, data_args, training_args = parser.parse_args_into_dataclasses() |
| 882 | local_rank = training_args.local_rank |
| 883 | compute_dtype = (torch.float16 if training_args.fp16 else (torch.bfloat16 if training_args.bf16 else torch.float32)) |
| 884 | |
| 885 | bnb_model_from_pretrained_args = {} |
| 886 | if training_args.bits in [4, 8]: |
| 887 | from transformers import BitsAndBytesConfig |
| 888 | bnb_model_from_pretrained_args.update(dict( |
| 889 | device_map={"": training_args.device}, |
| 890 | load_in_4bit=training_args.bits == 4, |
| 891 | load_in_8bit=training_args.bits == 8, |
| 892 | quantization_config=BitsAndBytesConfig( |
| 893 | load_in_4bit=training_args.bits == 4, |
| 894 | load_in_8bit=training_args.bits == 8, |
| 895 | llm_int8_skip_modules=["mm_projector"], |
| 896 | llm_int8_threshold=6.0, |
| 897 | llm_int8_has_fp16_weight=False, |
| 898 | bnb_4bit_compute_dtype=compute_dtype, |
| 899 | bnb_4bit_use_double_quant=training_args.double_quant, |
| 900 | bnb_4bit_quant_type=training_args.quant_type # {'fp4', 'nf4'} |
| 901 | ) |
| 902 | )) |
| 903 | |
| 904 | if model_args.vision_tower is not None: |
| 905 | if 'mpt' in model_args.model_name_or_path: |
| 906 | config = transformers.AutoConfig.from_pretrained(model_args.model_name_or_path, trust_remote_code=True) |
| 907 | config.attn_config['attn_impl'] = training_args.mpt_attn_impl |
| 908 | model = LlavaMptForCausalLM.from_pretrained( |
| 909 | model_args.model_name_or_path, |
| 910 | config=config, |
| 911 | cache_dir=training_args.cache_dir, |
| 912 | **bnb_model_from_pretrained_args |
| 913 | ) |
| 914 | else: |
| 915 | print (model_args.model_name_or_path) |
| 916 | model = LlavaLlamaForCausalLM.from_pretrained( |
| 917 | model_args.model_name_or_path, |
| 918 | cache_dir=training_args.cache_dir, |
| 919 | attn_implementation=attn_implementation, |
| 920 | use_depth = model_args.use_depth, |
| 921 | gt_depth = data_args.gt_depth, |
| 922 | torch_dtype=(torch.bfloat16 if training_args.bf16 else None), |
| 923 | **bnb_model_from_pretrained_args |
| 924 | ) |
| 925 | depth_state = torch.load(model_args.depth_path) |
| 926 | |
| 927 | model_dict = model.state_dict() |
| 928 | pretrained_dict = {'depth.'+key: value for key, value in depth_state.items() if ('depth.'+key) in model_dict.keys() } |
| 929 | |
| 930 | model_dict.update(pretrained_dict) |
| 931 | model.load_state_dict(model_dict) |
| 932 | |
| 933 |
no test coverage detected