MCPcopy Create free account
hub / github.com/OpenBMB/ToolBench / train

Function train

toolbench/train/train_lora.py:91–163  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

89
90
91def train():
92 parser = transformers.HfArgumentParser(
93 (ModelArguments, DataArguments, TrainingArguments, LoraArguments)
94 )
95 (
96 model_args,
97 data_args,
98 training_args,
99 lora_args,
100 ) = parser.parse_args_into_dataclasses()
101
102 if training_args.source_model_max_length < training_args.model_max_length:
103 condense_ratio = int(training_args.model_max_length/training_args.source_model_max_length)
104 # ratio = N means the sequence length is expanded by N, remember to change the model_max_length to 8192 (2048 * ratio) for ratio = 4
105 replace_llama_with_condense(ratio=condense_ratio)
106
107 world_size = int(os.environ.get("WORLD_SIZE", 1))
108 ddp = world_size != 1
109 device_map = {"": int(os.environ.get("LOCAL_RANK") or 0)} if ddp else None
110 model = transformers.AutoModelForCausalLM.from_pretrained(
111 model_args.model_name_or_path,
112 cache_dir=training_args.cache_dir,
113 device_map=device_map
114 )
115 lora_config = LoraConfig(
116 r=lora_args.lora_r,
117 lora_alpha=lora_args.lora_alpha,
118 target_modules=lora_args.lora_target_modules,
119 lora_dropout=lora_args.lora_dropout,
120 bias=lora_args.lora_bias,
121 task_type="CAUSAL_LM",
122 )
123 model = get_peft_model(model, lora_config)
124 if training_args.deepspeed is not None and training_args.local_rank == 0:
125 model.print_trainable_parameters()
126
127 if training_args.gradient_checkpointing:
128 logging.warning(
129 "gradient checkpointing with lora makes requires_grad "
130 "incorrect and needs a monkey patch in Trainer or the "
131 "wrapped model's forward. ref: "
132 "https://github.com/lm-sys/FastChat/pull/138#issuecomment-1509172198"
133 )
134 model.enable_input_require_grads()
135
136 tokenizer = transformers.AutoTokenizer.from_pretrained(
137 model_args.model_name_or_path,
138 cache_dir=training_args.cache_dir,
139 model_max_length=training_args.model_max_length,
140 padding_side="right",
141 use_fast=False,
142 )
143 tokenizer.pad_token = tokenizer.unk_token
144
145 data_module = make_supervised_data_module(tokenizer=tokenizer, data_args=data_args)
146 trainer = Trainer(
147 model=model, tokenizer=tokenizer, args=training_args, **data_module
148 )

Callers 1

train_lora.pyFile · 0.70

Calls 3

Tested by

no test coverage detected