MCPcopy Create free account
hub / github.com/cpystan/SD-VLM / train

Function train

llava/train/train.py:876–1093  ·  view source on GitHub ↗
(attn_implementation=None)

Source from the content-addressed store, hash-verified

874
875
876def 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

Callers 3

train_mem.pyFile · 0.90
train_xformers.pyFile · 0.90
train.pyFile · 0.85

Calls 13

LLaVATrainerClass · 0.90
find_all_linear_namesFunction · 0.85
rank0_printFunction · 0.85
updateMethod · 0.45
get_modelMethod · 0.45

Tested by

no test coverage detected