mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
can now use seq2seq for generation-based eval
This commit is contained in:
parent
92135a94af
commit
8c71161071
4 changed files with 173 additions and 39 deletions
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue