can now use seq2seq for generation-based eval

This commit is contained in:
51616 2024-12-23 08:34:33 +00:00
parent 92135a94af
commit 8c71161071
4 changed files with 173 additions and 39 deletions

View file

@ -1,3 +1,4 @@
from collections import defaultdict
from copy import copy
from functools import partial
from importlib.resources import read_binary
@ -29,6 +30,7 @@ from transformers import (
HfArgumentParser,
TrainingArguments,
)
from rouge_score import rouge_scorer
from utils import (
extract_cli_args,
get_run_name,
@ -50,22 +52,21 @@ from configs import (
logger = logging.getLogger()
def compute_per_token_acc(shifted_logits, shifted_labels):
indices = np.where(shifted_labels != -100)
acc = (shifted_logits.argmax(-1) == shifted_labels)[indices].mean()
return {"per_token_acc": acc, "num_valid_tokens": indices[0].size}
def compute_per_token_acc(shift_logits, shift_labels, valid_masks):
indices = np.where(valid_masks)
acc = (shift_logits.argmax(-1) == shift_labels)[indices].mean()
return {"per_token_acc": acc}
def compute_prefix_matching(shifted_logits, shifted_labels):
masks = np.where(shifted_labels != -100, 1, 0)
lengths = np.sum(masks, axis=1)
def compute_prefix_matching(shift_logits, shift_labels, valid_masks):
lengths = np.sum(valid_masks, axis=1)
is_wrong = (shifted_logits.argmax(-1) != shifted_labels) * masks
is_correct = (shifted_logits.argmax(-1) == shifted_labels) * masks
is_wrong = (shift_logits.argmax(-1) != shift_labels) * valid_masks
is_correct = (shift_logits.argmax(-1) == shift_labels) * valid_masks
# NOTE: not reliable for multi-turn conversations
# ie, all tokens in the following user's turn will be correct
# still monotonically correlate with perf though
wrong_pos = np.argmax(is_wrong, axis=1) - np.argmax(masks, axis=1)
wrong_pos = np.argmax(is_wrong, axis=1) - np.argmax(valid_masks, axis=1)
perf = wrong_pos / lengths
# if all tokens are correct, set to 1
@ -73,9 +74,9 @@ def compute_prefix_matching(shifted_logits, shifted_labels):
return {"prefix_matching": perf.mean()}
def compute_entropy(shifted_logits, shifted_labels):
indices = np.where(shifted_labels != -100)
logits = shifted_logits[indices]
def compute_entropy(shift_logits, shift_labels, valid_masks):
indices = np.where(valid_masks)
logits = shift_logits[indices]
probs = torch.softmax(torch.tensor(logits), dim=-1)
entropy = -torch.sum(probs * torch.log(probs), dim=-1)
return {"entropy": entropy.mean()}
@ -91,13 +92,54 @@ def compute_metrics(eval_pred: EvalPrediction) -> dict:
"""
# compute per token accuracy
logits, labels = eval_pred.predictions, eval_pred.label_ids
shifted_logits = logits[..., :-1, :]
shifted_labels = labels[..., 1:]
shift_logits = logits[..., :-1, :]
shift_labels = labels[..., 1:]
valid_masks = np.where(shift_labels != -100, 1, 0)
per_token_acc = compute_per_token_acc(shifted_logits, shifted_labels)
prefix_matching = compute_prefix_matching(shifted_logits, shifted_labels)
entropy = compute_entropy(shifted_logits, shifted_labels)
return dict(**per_token_acc, **prefix_matching, **entropy)
per_token_acc = compute_per_token_acc(shift_logits, shift_labels, valid_masks)
prefix_matching = compute_prefix_matching(shift_logits, shift_labels, valid_masks)
entropy = compute_entropy(shift_logits, shift_labels, valid_masks)
return dict(
**per_token_acc,
**prefix_matching,
**entropy,
num_valid_tokens=valid_masks.sum(),
num_samples=valid_masks.shape[0],
)
def compute_rouge(pred_texts, label_texts):
out = defaultdict(list)
scorer = rouge_scorer.RougeScorer(["rouge1", "rougeL"], use_stemmer=False)
for pred_text, label_text in zip(pred_texts, label_texts):
scores = scorer.score(pred_text, label_text)
for k, v in scores.items():
out[f"{k}.f1"].append(v.fmeasure)
for k in out:
out[k] = np.mean(out[k])
return out
def compute_generation_based_metrics(
eval_pred: EvalPrediction,
tokenizer: AutoTokenizer,
) -> dict:
pred_toks, labels = eval_pred.predictions, eval_pred.label_ids
start_indices = np.argmax(labels != -100, axis=1)
gen_toks = [x[start_indices[i] :] for i, x in enumerate(pred_toks)]
label_toks = [x[start_indices[i] :] for i, x in enumerate(labels)]
# print(pred_toks[indices])
# print(labels[indices])
# breakpoint()
gen_text = tokenizer.batch_decode(gen_toks, skip_special_tokens=True)
label_text = tokenizer.batch_decode(label_toks, skip_special_tokens=True)
rouge = compute_rouge(gen_text, label_text)
# acc = (pred_toks[indices] == labels[indices]).mean()
# "gen_text": gen_text, "label_text": label_text
return {**rouge}
def main(output_dir: str):
@ -163,8 +205,11 @@ def main(output_dir: str):
logger.info("Loading dataset...")
train_file = "data/raw_datasets/context_numbers/train.jsonl"
eval_file = "data/raw_datasets/context_numbers/val.jsonl"
ds = load_dataset("json", data_files={"train": train_file, "eval": eval_file})
val_file = "data/raw_datasets/context_numbers/val.jsonl"
test_file = "data/raw_datasets/context_numbers/test.jsonl"
ds = load_dataset(
"json", data_files={"train": train_file, "val": val_file, "test": test_file}
)
# preprocessing
ds = ds.map(get_preprocessing_fn("context_numbers"))
add_ctx_to_chat = not isinstance(model, ModulatedPretrainedModel)
@ -206,14 +251,15 @@ def main(output_dir: str):
validate_columns(tokenized_ds)
train_ds = tokenized_ds["train"]
eval_ds = {
val_ds = {
"train": tokenized_ds["train"].select(range(100)),
"val": tokenized_ds["eval"],
"val": tokenized_ds["val"],
}
test_ds = tokenized_ds["test"]
logger.debug(f"train_ds: {train_ds}")
logger.debug(f"eval_ds: {eval_ds}")
logger.debug(f"val_ds: {val_ds}")
logger.debug(f"test_ds: {test_ds}")
# TODO: change to a faster collator? e.g.,
# https://huggingface.co/blog/packing-with-FA2
# data_collator = DataCollatorForSeq2Seq(tokenizer, model, pad_to_multiple_of=8)
@ -258,13 +304,18 @@ def main(output_dir: str):
# might improve/decrease training speed w/ longer inputs
# TODO: add wandb notes somewhere
# wandb.init(project="ctx_to_lora", name=run_name, notes=args.notes)
# TODO: different collator for generation-based eval
train_model(
model,
train_ds,
eval_ds,
tokenizer,
training_args,
train_ds,
val_ds,
test_ds,
partial(collator, tokenizer=tokenizer),
compute_metrics,
partial(compute_generation_based_metrics, tokenizer=tokenizer),
)

View file

@ -32,6 +32,7 @@ def get_model_and_tokenizer(
dtype,
)
tokenizer = get_tokenizer(model_name_or_path, tokenizer_kwargs, peft_config, train)
model.config.pad_token_id = tokenizer.pad_token_id
return model, tokenizer

View file

@ -325,6 +325,14 @@ class ModulatedPretrainedModel(nn.Module):
self.register_module("hypernet", hypernet)
self.register_module("ctx_encoder", ctx_encoder)
# Delegate to base_model
@property
def config(self):
return self.base_model.config
def get_input_embeddings(self):
return self.base_model.get_input_embeddings()
def state_dict(self, *args, **kwargs):
state_dict = super().state_dict(*args, **kwargs)
# remove non-trainable and base_model's params
@ -337,9 +345,6 @@ class ModulatedPretrainedModel(nn.Module):
# NOTE: might have to set `strict=False` as we don't save all the params
return super().load_state_dict(state_dict, *args, **kwargs)
def get_input_embeddings(self):
return self.base_model.get_input_embeddings()
def get_ctx_features(
self,
examples: dict[str, Any],
@ -388,6 +393,33 @@ class ModulatedPretrainedModel(nn.Module):
return model_outputs
def generate(
self,
ctx_features: Optional[Float[Tensor, "bs ctx_length feature_dim"]] = None,
ctx_attn_mask: Optional[Integer[Tensor, "bs ctx_length"]] = None,
**model_inputs_kwargs: dict[str, Any],
):
if ctx_features is None:
logger.warning(
"No context ids provided, using the base model for the forward pass"
)
model_outputs = self.base_model(**model_inputs_kwargs)
# model_outputs.generated_loras = None
return model_outputs
generated_loras = self.hypernet.generate_loras(ctx_features, ctx_attn_mask)
# apply lora hook to the base model
# self.apply_generated_loras(generated_loras)
with apply_generated_loras(
self.base_model,
generated_loras,
self.hypernet.layer_indices,
self.training,
):
model_outputs = self.base_model.generate(**model_inputs_kwargs)
return model_outputs
@contextmanager
def apply_generated_loras(

View file

@ -1,6 +1,12 @@
from enum import Enum
from transformers import Seq2SeqTrainer, Trainer
import numpy as np
from transformers import (
GenerationConfig,
Trainer,
Seq2SeqTrainer,
Seq2SeqTrainingArguments,
)
from transformers.trainer_utils import get_last_checkpoint
@ -9,11 +15,14 @@ TRAINING_TASK = Enum("TRAINING_TASK", ["CAUSAL_LM", "COMPLETION"])
def train_model(
model,
train_dataset,
eval_dataset,
tokenizer,
training_args,
train_dataset=None,
val_dataset=None,
test_dataset=None,
data_collator=None,
compute_metrics=None,
compute_generation_based_metrics=None,
):
# last_checkpoint = None
@ -51,16 +60,11 @@ def train_model(
# if local_rank == 0:
# print(training_args)
# 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...
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
eval_dataset=val_dataset,
data_collator=data_collator,
compute_metrics=compute_metrics,
)
@ -81,3 +85,49 @@ def train_model(
trainer.log_metrics("eval", metrics)
trainer.save_metrics("eval", metrics)
trainer.save_model()
############## Evaluation
# NOTE: could also set kv_cache implementation here
eval_trainer_args = Seq2SeqTrainingArguments(
predict_with_generate=True,
generation_max_length=2**13,
generation_config=GenerationConfig(
do_sample=False,
# just a placeholder, will be overridden with `max_new_tokens`
max_length=2**13,
max_new_tokens=100,
# pad_token_id=tokenizer.pad_token_id,
# eos_token_id=?
),
**training_args.to_dict(),
# compute_metrics=compute_metrics,
)
# 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...
model.eval()
eval_trainer = Seq2SeqTrainer(
model=model,
args=eval_trainer_args,
# train_dataset=train_dataset,
eval_dataset=val_dataset,
# TODO: use a different collator for test, e.g., more max_len
data_collator=data_collator,
compute_metrics=compute_generation_based_metrics,
)
if val_dataset is not None:
eval_result = eval_trainer.evaluate()
eval_trainer.log_metrics("eval", eval_result)
eval_trainer.save_metrics("eval", eval_result)
if test_dataset is not None:
test_result = eval_trainer.predict(test_dataset)
print(test_result)
# eval_trainer.log_metrics("test", test_result)
# eval_trainer.save_metrics("test", test_result)