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

Function train

toolbench/train/train.py:251–291  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

249
250
251def train():
252 global local_rank
253
254 parser = transformers.HfArgumentParser(
255 (ModelArguments, DataArguments, TrainingArguments)
256 )
257 model_args, data_args, training_args = parser.parse_args_into_dataclasses()
258 if training_args.source_model_max_length < training_args.model_max_length:
259 condense_ratio = int(training_args.model_max_length/training_args.source_model_max_length)
260 # ratio = N means the sequence length is expanded by N, remember to change the model_max_length to 8192 (2048 * ratio) for ratio = 4
261 replace_llama_with_condense(ratio=condense_ratio)
262 local_rank = training_args.local_rank
263 tokenizer = transformers.AutoTokenizer.from_pretrained(
264 model_args.model_name_or_path,
265 cache_dir=training_args.cache_dir,
266 model_max_length=training_args.model_max_length,
267 padding_side="right",
268 use_fast=False,
269 )
270 tokenizer.pad_token = tokenizer.unk_token
271
272 data_module = make_supervised_data_module(tokenizer=tokenizer, data_args=data_args)
273 world_size = int(os.environ.get("WORLD_SIZE", 1))
274 ddp = world_size != 1
275 device_map = {"": int(os.environ.get("LOCAL_RANK") or 0)} if ddp else None
276 model = transformers.AutoModelForCausalLM.from_pretrained(
277 model_args.model_name_or_path,
278 cache_dir=training_args.cache_dir,
279 device_map=device_map
280 )
281 model.config.use_cache = False
282 trainer = Trainer(
283 model=model, tokenizer=tokenizer, args=training_args, **data_module
284 )
285
286 if list(pathlib.Path(training_args.output_dir).glob("checkpoint-*")):
287 trainer.train(resume_from_checkpoint=True)
288 else:
289 trainer.train()
290 trainer.save_state()
291 safe_save_model_for_hf_trainer(trainer=trainer, output_dir=training_args.output_dir)
292
293
294if __name__ == "__main__":

Callers 2

train_mem.pyFile · 0.90
train.pyFile · 0.70

Tested by

no test coverage detected