import json import logging from dataclasses import fields from enum import Enum import numpy as np from transformers import ( GenerationConfig, Trainer, Seq2SeqTrainer, Seq2SeqTrainingArguments, ) from transformers.trainer_utils import get_last_checkpoint TRAINING_TASK = Enum("TRAINING_TASK", ["CAUSAL_LM", "COMPLETION"]) logger = logging.getLogger() def save_generated_text(samples, output_dir, split): with open(f"{output_dir}/{split}_generated_text.jsonl", "w") as f: for sample in samples: f.write(json.dumps(sample) + "\n") def decode_test_result(test_dataset, test_result, tokenizer): for sample, pred_toks, labels in zip( test_dataset, test_result.predictions, test_result.label_ids ): start_idx = np.argmax(labels != -100, axis=0) input_toks = sample["input_ids"][:start_idx] gen_toks = pred_toks[start_idx:] label_toks = labels[start_idx:] # labels are padded with -100, so we need to replace them with the pad token id label_toks = np.where(label_toks == -100, tokenizer.pad_token_id, label_toks) input_text = tokenizer.decode(input_toks, skip_special_tokens=True) gen_text = tokenizer.decode(gen_toks, skip_special_tokens=True) label_text = tokenizer.decode(label_toks, skip_special_tokens=True) yield {"input": input_text, "generated": gen_text, "label": label_text} def eval_generation(eval_trainer, tokenizer, dataset, split, gen_kwargs): eval_result = eval_trainer.predict( dataset, metric_key_prefix=split, **gen_kwargs, ) eval_trainer.log_metrics("eval" if split == "val" else split, eval_result.metrics) eval_trainer.save_metrics("eval" if split == "val" else split, eval_result.metrics) save_generated_text( decode_test_result(dataset, eval_result, tokenizer), split=split, output_dir=eval_trainer.args.output_dir, ) def train_model( model, tokenizer, training_args, train_dataset=None, val_dataset=None, test_dataset=None, train_collator=None, generation_collator=None, compute_metrics=None, compute_generation_based_metrics=None, ): # last_checkpoint = None # if ( # os.path.isdir(training_args.output_dir) # and not training_args.overwrite_output_dir # ): # last_checkpoint = get_last_checkpoint(training_args.output_dir) # if last_checkpoint is None and len(os.listdir(training_args.output_dir)) > 0: # raise ValueError( # f"Output directory ({training_args.output_dir})" # " already exists and is not empty. " # "Use --overwrite_output_dir to overcome." # ) # elif ( # last_checkpoint is not None and training_args.resume_from_checkpoint is None # ): # print( # f"Checkpoint detected, resuming training at {last_checkpoint}. " # "To avoid this behavior, change " # "the `--output_dir` or add `--overwrite_output_dir` to train from scratch." # ) # if ( # max( # training_args.per_device_train_batch_size, # training_args.per_device_eval_batch_size, # ) # == 1 # ): # data_collator = None # # print training_args at local_rank 0 # local_rank = int(os.getenv("LOCAL_RANK", "0")) # if local_rank == 0: # print(training_args) trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, eval_dataset=val_dataset, data_collator=train_collator, compute_metrics=compute_metrics, ) checkpoint = None # if training_args.resume_from_checkpoint is not None: # checkpoint = training_args.resume_from_checkpoint # elif last_checkpoint is not None: # checkpoint = last_checkpoint # print(f"Loaded from the checkpoint: {checkpoint}") # TODO: save the best model based on eval loss? train_result = trainer.train(resume_from_checkpoint=checkpoint) trainer.log_metrics("train", train_result.metrics) metrics = trainer.evaluate() trainer.log_metrics("eval", metrics) trainer.save_metrics("eval", metrics) trainer.save_model() ############## Evaluation # TODO: generalize gen_kwargs for validation max_new_tokens = 100 # max_input_len=2**13 # for input truncation gen_kwargs = dict(do_sample=False, max_new_tokens=max_new_tokens) # pad_token_id=tokenizer.pad_token_id, # eos_token_id=? eval_trainer_args = {} # Copy only necessary attributes from training_args to eval_trainer_args seq2seq_training_args_fields = {f.name for f in fields(Seq2SeqTrainingArguments)} for attr, value in training_args.to_dict().items(): if attr in seq2seq_training_args_fields: eval_trainer_args[attr] = value eval_trainer_args["eval_strategy"] = "no" eval_trainer_args["overwrite_output_dir"] = True # NOTE: could also set kv_cache implementation here eval_trainer_args = Seq2SeqTrainingArguments( **eval_trainer_args, predict_with_generate=True, generation_config=GenerationConfig(**gen_kwargs), ) # Seq2SeqTrainer is actually just the same as Trainer # (although it uses a different data collator, i.e., explicit prompt/answer separation) # it just allows `predict_with_generate` # allowing us to compute metrics on the generated outputs # no clue why they call this seq2seq... logger.info("=" * 80 + "\n" + "Evaluating model..." + "\n" + "=" * 80) model.eval() eval_trainer = Seq2SeqTrainer( model=model, args=eval_trainer_args, # TODO: use a different collator for test, e.g., more max_len truncation # w/ left padding? # removing label part from input_ids data_collator=generation_collator, compute_metrics=compute_generation_based_metrics, ) if val_dataset is not None: if isinstance(val_dataset, dict): val_dataset = val_dataset["val"] eval_generation(eval_trainer, tokenizer, val_dataset, "val", gen_kwargs) if test_dataset is not None: eval_generation(eval_trainer, tokenizer, test_dataset, "test", gen_kwargs)