mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
apply icae minor mod
This commit is contained in:
parent
f2cdec08ee
commit
e69a76fc9f
9 changed files with 94 additions and 865 deletions
|
|
@ -1,6 +0,0 @@
|
|||
{"input": "French senior civil servant arrested on suspicion of spying for North Korea\n\nNovember 27, 2018 by Joseph Fitsanakis\n\nA senior civil servant in the upper house of the French parliament has been arrested on suspicion of spying for North Korea, according to prosecutors. The news of the suspected spy\u2019s arrest was first reported on Monday by Quotidien, a daily politics and culture show on the Monaco-based television channel TMC. The show cited \u201ca judicial source in Paris\u201d and said that France\u2019s domestic security and counterintelligence agency, the General Directorate for Internal Security (DGSI), was in charge of the espionage case.\n\nThe senior administrator has been identified as Benoit Quennedey, a civil servant who liaises between the French Senate and the Department of Architecture and Heritage, which operates under France\u2019s Ministry of Culture. Quennedey was reportedly detained on Sunday morning and his office in the French Senate was raided by DGSI officers on the same day. Quotidien said that he was arrested on suspicion of \u201ccollecting and delivering to a foreign power information likely to subvert core national interests\u201d. The report did not provide specific information about the type of information that Quennedey is believed to have passed to North Korea. It did state, however, that a counterintelligence investigation into his activities began in March of this year.\n\nQuennedey is believed to be the president of the Franco-Korean Friendship Association, the French branch of a Spanish-based organization that lobbies in favor of international support for North Korea. Korea Friendship Association branches exist in over 30 countries and are believed to be officially sanctioned by Pyongyang. They operate as something akin to the pre-World War II Comintern (Communist International), a Moscow-sanctioned international pressure group that advocated in favor of Soviet-style communism around the world. French media reported on Monday that Quennedey traveled extensively to the Korean Peninsula in the past decade and has written a French-language book on North Korea. News reports said that the French President Emmanuel Macron had been made aware of Quennedey\u2019s arrest. The senior civil servant faces up to 30 years in prison if found guilty of espionage.\n\n\u25ba Author: Joseph Fitsanakis | Date: 27 November 2018 | Permalink\n\n", "prompt": "Identify the person arrested on suspicion of spying for North Korea.", "answer": "Benoit Quennedey"}
|
||||
{"input": "French senior civil servant arrested on suspicion of spying for North Korea\n\nNovember 27, 2018 by Joseph Fitsanakis\n\nA senior civil servant in the upper house of the French parliament has been arrested on suspicion of spying for North Korea, according to prosecutors. The news of the suspected spy\u2019s arrest was first reported on Monday by Quotidien, a daily politics and culture show on the Monaco-based television channel TMC. The show cited \u201ca judicial source in Paris\u201d and said that France\u2019s domestic security and counterintelligence agency, the General Directorate for Internal Security (DGSI), was in charge of the espionage case.\n\nThe senior administrator has been identified as Benoit Quennedey, a civil servant who liaises between the French Senate and the Department of Architecture and Heritage, which operates under France\u2019s Ministry of Culture. Quennedey was reportedly detained on Sunday morning and his office in the French Senate was raided by DGSI officers on the same day. Quotidien said that he was arrested on suspicion of \u201ccollecting and delivering to a foreign power information likely to subvert core national interests\u201d. The report did not provide specific information about the type of information that Quennedey is believed to have passed to North Korea. It did state, however, that a counterintelligence investigation into his activities began in March of this year.\n\nQuennedey is believed to be the president of the Franco-Korean Friendship Association, the French branch of a Spanish-based organization that lobbies in favor of international support for North Korea. Korea Friendship Association branches exist in over 30 countries and are believed to be officially sanctioned by Pyongyang. They operate as something akin to the pre-World War II Comintern (Communist International), a Moscow-sanctioned international pressure group that advocated in favor of Soviet-style communism around the world. French media reported on Monday that Quennedey traveled extensively to the Korean Peninsula in the past decade and has written a French-language book on North Korea. News reports said that the French President Emmanuel Macron had been made aware of Quennedey\u2019s arrest. The senior civil servant faces up to 30 years in prison if found guilty of espionage.\n\n\u25ba Author: Joseph Fitsanakis | Date: 27 November 2018 | Permalink\n\n", "prompt": "List the actions taken against Benoit Quennedey after his arrest.", "answer": "His office in the French Senate was raided by DGSI officers."}
|
||||
{"input": "French senior civil servant arrested on suspicion of spying for North Korea\n\nNovember 27, 2018 by Joseph Fitsanakis\n\nA senior civil servant in the upper house of the French parliament has been arrested on suspicion of spying for North Korea, according to prosecutors. The news of the suspected spy\u2019s arrest was first reported on Monday by Quotidien, a daily politics and culture show on the Monaco-based television channel TMC. The show cited \u201ca judicial source in Paris\u201d and said that France\u2019s domestic security and counterintelligence agency, the General Directorate for Internal Security (DGSI), was in charge of the espionage case.\n\nThe senior administrator has been identified as Benoit Quennedey, a civil servant who liaises between the French Senate and the Department of Architecture and Heritage, which operates under France\u2019s Ministry of Culture. Quennedey was reportedly detained on Sunday morning and his office in the French Senate was raided by DGSI officers on the same day. Quotidien said that he was arrested on suspicion of \u201ccollecting and delivering to a foreign power information likely to subvert core national interests\u201d. The report did not provide specific information about the type of information that Quennedey is believed to have passed to North Korea. It did state, however, that a counterintelligence investigation into his activities began in March of this year.\n\nQuennedey is believed to be the president of the Franco-Korean Friendship Association, the French branch of a Spanish-based organization that lobbies in favor of international support for North Korea. Korea Friendship Association branches exist in over 30 countries and are believed to be officially sanctioned by Pyongyang. They operate as something akin to the pre-World War II Comintern (Communist International), a Moscow-sanctioned international pressure group that advocated in favor of Soviet-style communism around the world. French media reported on Monday that Quennedey traveled extensively to the Korean Peninsula in the past decade and has written a French-language book on North Korea. News reports said that the French President Emmanuel Macron had been made aware of Quennedey\u2019s arrest. The senior civil servant faces up to 30 years in prison if found guilty of espionage.\n\n\u25ba Author: Joseph Fitsanakis | Date: 27 November 2018 | Permalink\n\n", "prompt": "What is the role of Benoit Quennedey in the French government?", "answer": "He is a senior civil servant who liaises between the French Senate and the Department of Architecture and Heritage."}
|
||||
{"input": "French senior civil servant arrested on suspicion of spying for North Korea\n\nNovember 27, 2018 by Joseph Fitsanakis\n\nA senior civil servant in the upper house of the French parliament has been arrested on suspicion of spying for North Korea, according to prosecutors. The news of the suspected spy\u2019s arrest was first reported on Monday by Quotidien, a daily politics and culture show on the Monaco-based television channel TMC. The show cited \u201ca judicial source in Paris\u201d and said that France\u2019s domestic security and counterintelligence agency, the General Directorate for Internal Security (DGSI), was in charge of the espionage case.\n\nThe senior administrator has been identified as Benoit Quennedey, a civil servant who liaises between the French Senate and the Department of Architecture and Heritage, which operates under France\u2019s Ministry of Culture. Quennedey was reportedly detained on Sunday morning and his office in the French Senate was raided by DGSI officers on the same day. Quotidien said that he was arrested on suspicion of \u201ccollecting and delivering to a foreign power information likely to subvert core national interests\u201d. The report did not provide specific information about the type of information that Quennedey is believed to have passed to North Korea. It did state, however, that a counterintelligence investigation into his activities began in March of this year.\n\nQuennedey is believed to be the president of the Franco-Korean Friendship Association, the French branch of a Spanish-based organization that lobbies in favor of international support for North Korea. Korea Friendship Association branches exist in over 30 countries and are believed to be officially sanctioned by Pyongyang. They operate as something akin to the pre-World War II Comintern (Communist International), a Moscow-sanctioned international pressure group that advocated in favor of Soviet-style communism around the world. French media reported on Monday that Quennedey traveled extensively to the Korean Peninsula in the past decade and has written a French-language book on North Korea. News reports said that the French President Emmanuel Macron had been made aware of Quennedey\u2019s arrest. The senior civil servant faces up to 30 years in prison if found guilty of espionage.\n\n\u25ba Author: Joseph Fitsanakis | Date: 27 November 2018 | Permalink\n\n", "prompt": "What are the charges against Benoit Quennedey?", "answer": "He is suspected of collecting and delivering to a foreign power information likely to subvert core national interests."}
|
||||
{"input": "French senior civil servant arrested on suspicion of spying for North Korea\n\nNovember 27, 2018 by Joseph Fitsanakis\n\nA senior civil servant in the upper house of the French parliament has been arrested on suspicion of spying for North Korea, according to prosecutors. The news of the suspected spy\u2019s arrest was first reported on Monday by Quotidien, a daily politics and culture show on the Monaco-based television channel TMC. The show cited \u201ca judicial source in Paris\u201d and said that France\u2019s domestic security and counterintelligence agency, the General Directorate for Internal Security (DGSI), was in charge of the espionage case.\n\nThe senior administrator has been identified as Benoit Quennedey, a civil servant who liaises between the French Senate and the Department of Architecture and Heritage, which operates under France\u2019s Ministry of Culture. Quennedey was reportedly detained on Sunday morning and his office in the French Senate was raided by DGSI officers on the same day. Quotidien said that he was arrested on suspicion of \u201ccollecting and delivering to a foreign power information likely to subvert core national interests\u201d. The report did not provide specific information about the type of information that Quennedey is believed to have passed to North Korea. It did state, however, that a counterintelligence investigation into his activities began in March of this year.\n\nQuennedey is believed to be the president of the Franco-Korean Friendship Association, the French branch of a Spanish-based organization that lobbies in favor of international support for North Korea. Korea Friendship Association branches exist in over 30 countries and are believed to be officially sanctioned by Pyongyang. They operate as something akin to the pre-World War II Comintern (Communist International), a Moscow-sanctioned international pressure group that advocated in favor of Soviet-style communism around the world. French media reported on Monday that Quennedey traveled extensively to the Korean Peninsula in the past decade and has written a French-language book on North Korea. News reports said that the French President Emmanuel Macron had been made aware of Quennedey\u2019s arrest. The senior civil servant faces up to 30 years in prison if found guilty of espionage.\n\n\u25ba Author: Joseph Fitsanakis | Date: 27 November 2018 | Permalink\n\n", "prompt": "Mention the organization Benoit Quennedey is believed to be the president of.", "answer": "Franco-Korean Friendship Association"}
|
||||
{"input": "French senior civil servant arrested on suspicion of spying for North Korea\n\nNovember 27, 2018 by Joseph Fitsanakis\n\nA senior civil servant in the upper house of the French parliament has been arrested on suspicion of spying for North Korea, according to prosecutors. The news of the suspected spy\u2019s arrest was first reported on Monday by Quotidien, a daily politics and culture show on the Monaco-based television channel TMC. The show cited \u201ca judicial source in Paris\u201d and said that France\u2019s domestic security and counterintelligence agency, the General Directorate for Internal Security (DGSI), was in charge of the espionage case.\n\nThe senior administrator has been identified as Benoit Quennedey, a civil servant who liaises between the French Senate and the Department of Architecture and Heritage, which operates under France\u2019s Ministry of Culture. Quennedey was reportedly detained on Sunday morning and his office in the French Senate was raided by DGSI officers on the same day. Quotidien said that he was arrested on suspicion of \u201ccollecting and delivering to a foreign power information likely to subvert core national interests\u201d. The report did not provide specific information about the type of information that Quennedey is believed to have passed to North Korea. It did state, however, that a counterintelligence investigation into his activities began in March of this year.\n\nQuennedey is believed to be the president of the Franco-Korean Friendship Association, the French branch of a Spanish-based organization that lobbies in favor of international support for North Korea. Korea Friendship Association branches exist in over 30 countries and are believed to be officially sanctioned by Pyongyang. They operate as something akin to the pre-World War II Comintern (Communist International), a Moscow-sanctioned international pressure group that advocated in favor of Soviet-style communism around the world. French media reported on Monday that Quennedey traveled extensively to the Korean Peninsula in the past decade and has written a French-language book on North Korea. News reports said that the French President Emmanuel Macron had been made aware of Quennedey\u2019s arrest. The senior civil servant faces up to 30 years in prison if found guilty of espionage.\n\n\u25ba Author: Joseph Fitsanakis | Date: 27 November 2018 | Permalink\n\n", "prompt": "When did the counterintelligence investigation into Quennedey's activities begin?", "answer": "In March of this year."}
|
||||
|
|
@ -1,127 +0,0 @@
|
|||
import json
|
||||
import sys
|
||||
|
||||
import torch
|
||||
from modeling_icae_multi_span import (
|
||||
ICAE,
|
||||
DataArguments,
|
||||
ModelArguments,
|
||||
TrainingArguments,
|
||||
)
|
||||
from peft import LoraConfig
|
||||
from safetensors.torch import load_file
|
||||
from tqdm import tqdm
|
||||
from transformers import AutoModelForCausalLM, HfArgumentParser
|
||||
|
||||
# Set the computation device
|
||||
device = "cuda"
|
||||
|
||||
# Parse model, data, and training arguments
|
||||
parser = HfArgumentParser((ModelArguments, DataArguments, TrainingArguments))
|
||||
model_args, data_args, training_args = parser.parse_args_into_dataclasses()
|
||||
|
||||
# Define Lora configuration
|
||||
lora_config = LoraConfig(
|
||||
r=512, lora_alpha=32, lora_dropout=model_args.lora_dropout, bias="none", task_type="CAUSAL_LM"
|
||||
)
|
||||
|
||||
# Initialize model and send it to CUDA device
|
||||
model = ICAE(model_args, training_args, lora_config)
|
||||
|
||||
# Load the fine-tuned checkpoint
|
||||
print(f"Loading trained checkpoint from {training_args.output_dir}")
|
||||
state_dict = load_file(training_args.output_dir)
|
||||
model.load_state_dict(state_dict, strict=False) # only load lora and memory token embeddings
|
||||
|
||||
model = model.to(device)
|
||||
|
||||
# Read the data file
|
||||
file_path = "./dev_v2.jsonl"
|
||||
lines = None
|
||||
with open(file_path) as f:
|
||||
lines = f.readlines()
|
||||
|
||||
# Prepare the model for evaluation
|
||||
max_out_length = 512
|
||||
model.eval()
|
||||
|
||||
with torch.no_grad():
|
||||
with open("ft_inference.out", "w") as f:
|
||||
|
||||
for line in tqdm(lines):
|
||||
# Tokenize input text
|
||||
data = json.loads(line)
|
||||
tokenized_input = model.tokenizer(
|
||||
data["input"],
|
||||
truncation=True,
|
||||
max_length=5120,
|
||||
padding=False,
|
||||
return_attention_mask=False,
|
||||
)
|
||||
tokenized_prompt = model.tokenizer(
|
||||
data["prompt"],
|
||||
truncation=False,
|
||||
padding=False,
|
||||
return_attention_mask=False,
|
||||
add_special_tokens=False,
|
||||
)
|
||||
# Generate compressed outputs
|
||||
input_ids = torch.LongTensor([tokenized_input["input_ids"]]).to(device)
|
||||
memory_slots = model._compress(input_ids)
|
||||
|
||||
# decoder input has 3 parts: prefix, memory slots and suffix
|
||||
# the following code is for Mistral tokenizer for example: 733, 16289, 28793 are for the Mistral instruction tempmlate
|
||||
prompt_left_ids = torch.LongTensor([[1, 733, 16289, 28793]]).to(device)
|
||||
prompt_right_ids = (
|
||||
[model.ft_token_id] + tokenized_prompt["input_ids"] + [733, 28748, 16289, 28793]
|
||||
)
|
||||
prompt_right_ids = torch.LongTensor([prompt_right_ids]).to(device)
|
||||
|
||||
prompt_left_embs = model.tokens_to_embeddings(prompt_left_ids)
|
||||
prompt_right_embs = model.tokens_to_embeddings(prompt_right_ids)
|
||||
memory_slots = memory_slots.to(prompt_right_embs)
|
||||
|
||||
# Concatenate and clone input embeddings
|
||||
decoder_input_embeddings = torch.cat(
|
||||
(prompt_left_embs, memory_slots.unsqueeze(0), prompt_right_embs), dim=1
|
||||
)
|
||||
output = decoder_input_embeddings.clone()
|
||||
|
||||
generate_text = []
|
||||
past_key_values = None
|
||||
|
||||
# Generate text output
|
||||
for i in range(max_out_length):
|
||||
with model.icae.disable_adapter(): # no independent decoder; use self.icae
|
||||
out = model.icae(
|
||||
inputs_embeds=output, past_key_values=past_key_values, use_cache=True
|
||||
)
|
||||
# out = decoder(inputs_embeds=output, past_key_values=past_key_values, use_cache=True)
|
||||
logit = out.logits[:, -1, : model.vocab_size - 1]
|
||||
past_key_values = out.past_key_values
|
||||
|
||||
next_token_id = torch.argmax(logit, dim=-1)
|
||||
# print(next_token_id)
|
||||
|
||||
if next_token_id.item() == 2: # eos
|
||||
break
|
||||
|
||||
output = (
|
||||
model.icae.get_base_model()
|
||||
.model.embed_tokens(next_token_id)
|
||||
.unsqueeze(1)
|
||||
.to(device)
|
||||
)
|
||||
generate_text.append(next_token_id.item())
|
||||
|
||||
generated_text = model.tokenizer.decode(generate_text)
|
||||
|
||||
# Structure output data
|
||||
output_ = {
|
||||
"input": data["input"],
|
||||
"prompt": data["prompt"],
|
||||
"output": generated_text,
|
||||
"answer": data["answer"],
|
||||
}
|
||||
|
||||
f.write(json.dumps(output_) + "\n")
|
||||
|
|
@ -1,16 +0,0 @@
|
|||
#!/bin/bash
|
||||
|
||||
# MODEL="mistralai/Mistral-7B-v0.1"
|
||||
BASE_MODEL="mistralai/Mistral-7B-Instruct-v0.2"
|
||||
# MODEL="meta-llama/Llama-2-7b-hf"
|
||||
# MODEL="meta-llama/Llama-2-7b-chat-hf"
|
||||
MODEL_NAME="${MODEL//\//-}"
|
||||
|
||||
maxlen=5120
|
||||
mem=128
|
||||
r=512
|
||||
mean_compression_rate=4
|
||||
|
||||
ICAE_MODEL_PATH=$1 # ICAE model to use; wget "https://huggingface.co/sggetao/icae/resolve/main/mistral_7b_ft_icae.safetensors"
|
||||
|
||||
python fine_tuned_inference.py --mean_compression_rate $mean_compression_rate --model_max_length $maxlen --fixed_mem_size $mem --lora_r $r --output_dir $ICAE_MODEL_PATH --model_name_or_path $BASE_MODEL --bf16 --train False
|
||||
|
|
@ -1,119 +0,0 @@
|
|||
from datasets import load_dataset
|
||||
from modeling_icae_multi_span import (
|
||||
ICAE,
|
||||
DataArguments,
|
||||
ModelArguments,
|
||||
TrainingArguments,
|
||||
)
|
||||
from peft import LoraConfig
|
||||
from training_utils import (
|
||||
DataCollatorForDynamicPadding,
|
||||
instruct_ft_tokenize_function,
|
||||
train_model,
|
||||
)
|
||||
from transformers import HfArgumentParser
|
||||
|
||||
|
||||
def compute_metrics(eval_pred) -> dict:
|
||||
"""
|
||||
Custom metrics function for the trainer
|
||||
Args:
|
||||
eval_pred: tuple of predictions and labels
|
||||
Returns:
|
||||
dictionary containing metric names (str) and values (Any)
|
||||
"""
|
||||
|
||||
# preds, labels = eval_preds
|
||||
|
||||
# # predictions is generated tokens for Seq2SeqTrainer
|
||||
# # decode preds and labels
|
||||
# labels = np.where(labels != -100, labels, tokenizer.pad_token_id)
|
||||
# decoded_preds = tokenizer.batch_decode(preds, skip_special_tokens=True)
|
||||
# decoded_labels = tokenizer.batch_decode(labels, skip_special_tokens=True)
|
||||
|
||||
# compute per token accuracy
|
||||
predictions, labels = eval_pred.predictions, eval_pred.label_ids
|
||||
|
||||
# predictions is logits for Trainer
|
||||
preds = predictions.argmax(-1)
|
||||
acc = (preds == labels).mean()
|
||||
return {"per_token_acc": acc}
|
||||
|
||||
|
||||
def main():
|
||||
parser = HfArgumentParser((ModelArguments, DataArguments, TrainingArguments))
|
||||
model_args, data_args, training_args = parser.parse_args_into_dataclasses()
|
||||
|
||||
print(model_args)
|
||||
print(data_args)
|
||||
|
||||
training_args.eval_on_start = True
|
||||
training_args.eval_strategy = "steps"
|
||||
training_args.eval_steps = 500
|
||||
training_args.save_strategy = "no"
|
||||
# training_args.save_steps = 500
|
||||
training_args.logging_strategy = "steps"
|
||||
training_args.logging_steps = 100
|
||||
|
||||
# seq2seq args for generation evaluation
|
||||
# training_args.predict_with_generate = True
|
||||
# training_args.generation_max_length = 100
|
||||
|
||||
training_args.gradient_checkpointing_kwargs = {
|
||||
"use_reentrant": False
|
||||
} # manually add this argument in the code
|
||||
|
||||
lora_config = LoraConfig(
|
||||
r=model_args.lora_r,
|
||||
lora_alpha=32,
|
||||
lora_dropout=0.05,
|
||||
bias="none",
|
||||
task_type="CAUSAL_LM",
|
||||
)
|
||||
|
||||
# check model_args.mem_size and min_tokens_for_lm
|
||||
assert (
|
||||
training_args.fixed_mem_size & (training_args.fixed_mem_size - 1)
|
||||
) == 0, "training_args.fixed_mem_size must be a power of 2"
|
||||
|
||||
memory_size = training_args.fixed_mem_size
|
||||
|
||||
train_file = "../data/raw_datasets/context_numbers/train.jsonl"
|
||||
eval_file = "../data/raw_datasets/context_numbers/val.jsonl"
|
||||
|
||||
print("Loading dataset...")
|
||||
|
||||
dataset = load_dataset("json", data_files={"train": train_file, "eval": eval_file})
|
||||
|
||||
train_dataset = dataset["train"]
|
||||
eval_dataset = dataset["eval"]
|
||||
|
||||
model = ICAE(model_args, training_args, lora_config).to("cuda")
|
||||
MEM_TOKENS = list(range(model.vocab_size, model.vocab_size + memory_size))
|
||||
|
||||
tokenized_train_ds = train_dataset.map(
|
||||
instruct_ft_tokenize_function,
|
||||
batched=True,
|
||||
fn_kwargs={"model": model, "mem": MEM_TOKENS},
|
||||
)
|
||||
|
||||
tokenized_eval_ds = {"train": train_dataset.select(range(100)), "val": eval_dataset}
|
||||
|
||||
for split in tokenized_eval_ds:
|
||||
tokenized_eval_ds[split] = tokenized_eval_ds[split].map(
|
||||
instruct_ft_tokenize_function,
|
||||
batched=True,
|
||||
fn_kwargs={"model": model, "mem": MEM_TOKENS},
|
||||
)
|
||||
|
||||
train_model(
|
||||
model,
|
||||
tokenized_train_ds,
|
||||
tokenized_eval_ds,
|
||||
training_args,
|
||||
compute_metrics=compute_metrics,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -1,312 +0,0 @@
|
|||
# ICAE that supports multi span concat
|
||||
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import transformers
|
||||
from peft import get_peft_model
|
||||
from safetensors.torch import load_file
|
||||
from torch.nn.functional import gelu
|
||||
from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelArguments:
|
||||
model_name_or_path: str = field(default="mistralai/Mistral-7B-v0.1")
|
||||
lora_r: int = field(default=128, metadata={"help": "lora rank"})
|
||||
lora_dropout: float = field(default=0.05, metadata={"help": "lora dropout"})
|
||||
train: bool = field(
|
||||
default=True,
|
||||
metadata={
|
||||
"help": "if true, the model ckpt will be initialized for training; else, it's for inference"
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DataArguments:
|
||||
data_path: str = field(default=None, metadata={"help": "Path to the training data."})
|
||||
debug_data: bool = field(
|
||||
default=False,
|
||||
metadata={"help": "Enable debug dataset to quickly verify the training process"},
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class TrainingArguments(transformers.Seq2SeqTrainingArguments):
|
||||
cache_dir: Optional[str] = field(default=None)
|
||||
optim: str = field(default="adamw_torch")
|
||||
model_max_length: int = field(
|
||||
default=28000,
|
||||
metadata={
|
||||
"help": "Maximum sequence length. Sequences will be right padded (and possibly truncated)."
|
||||
},
|
||||
)
|
||||
fixed_mem_size: int = field(
|
||||
default=128,
|
||||
metadata={"help": "Enalbing the fixed mem size."},
|
||||
)
|
||||
mean_compression_rate: int = field(
|
||||
default=4,
|
||||
metadata={"help": "Mean compression rate; default=4"},
|
||||
)
|
||||
min_tokens_for_lm: int = field(
|
||||
default=64,
|
||||
metadata={"help": "Minimum tokens for lm objective learning"},
|
||||
)
|
||||
leave_tokens_for_lm: int = field(
|
||||
default=8,
|
||||
metadata={"help": "Leave some tokens without loss for lm objective"},
|
||||
)
|
||||
lm_ratio: float = field(
|
||||
default=0.0,
|
||||
metadata={"help": "Ratio for LM training."},
|
||||
)
|
||||
add_special_token_for_lm: bool = field(
|
||||
default=False,
|
||||
metadata={
|
||||
"help": "Add a special token for the prompt of language modeling; default: False"
|
||||
},
|
||||
)
|
||||
restore_from: str = field(
|
||||
default="",
|
||||
metadata={"help": "The checkpoint that should be restored from for fine-tuning"},
|
||||
)
|
||||
|
||||
|
||||
def print_trainable_parameters(model):
|
||||
trainable_parameters = 0
|
||||
all_param = 0
|
||||
for _, param in model.named_parameters():
|
||||
all_param += param.numel()
|
||||
if param.requires_grad:
|
||||
trainable_parameters += param.numel()
|
||||
print(
|
||||
f"trainable params: {trainable_parameters} || all params: {all_param} || trainable%: {100 * trainable_parameters / all_param}"
|
||||
)
|
||||
# for name, param in model.named_parameters():
|
||||
# if param.requires_grad:
|
||||
# print(name, param.shape)
|
||||
|
||||
|
||||
def freeze_model(model):
|
||||
for _, param in model.named_parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
|
||||
class ICAE(torch.nn.Module):
|
||||
def __init__(self, model_args, training_args, lora_config):
|
||||
super().__init__()
|
||||
self.model_args = model_args
|
||||
self.training_args = training_args
|
||||
self.model_name = model_args.model_name_or_path
|
||||
self.icae = AutoModelForCausalLM.from_pretrained(
|
||||
self.model_name,
|
||||
torch_dtype=torch.float16 if training_args.bf16 is False else torch.bfloat16,
|
||||
use_flash_attention_2=True,
|
||||
resume_download=True,
|
||||
)
|
||||
|
||||
self.training = self.model_args.train
|
||||
|
||||
if self.training: # indepedent model for gradient checkpointing
|
||||
self.decoder = AutoModelForCausalLM.from_pretrained(
|
||||
self.model_name,
|
||||
torch_dtype=torch.float16 if training_args.bf16 is False else torch.bfloat16,
|
||||
use_flash_attention_2=True,
|
||||
resume_download=True,
|
||||
)
|
||||
|
||||
self.vocab_size = self.icae.config.vocab_size + 1 # [PAD] token
|
||||
self.pad_token_id = self.vocab_size - 1
|
||||
self.mean_compression_rate = training_args.mean_compression_rate
|
||||
|
||||
# tunable
|
||||
self.mem_size = self.training_args.fixed_mem_size
|
||||
self.vocab_size_with_mem = (
|
||||
self.vocab_size + self.mem_size
|
||||
) # so, the mem tokens are in the range [self.vocab_size, self.vocab_size + self.mem_size)
|
||||
|
||||
# special tokens in addition to mem and length tokens
|
||||
self.ae_token_id = self.vocab_size_with_mem + 0
|
||||
self.lm_token_id = self.vocab_size_with_mem + 1
|
||||
self.ft_token_id = self.vocab_size_with_mem + 2
|
||||
|
||||
self.icae.resize_token_embeddings(self.vocab_size_with_mem + 3)
|
||||
|
||||
# special tokens for Llama-2/Mistral tokenizer
|
||||
self.bos_id = 1
|
||||
self.eos_id = 2
|
||||
|
||||
self.dim = self.icae.config.hidden_size
|
||||
self.icae = get_peft_model(self.icae, lora_config)
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
self.memory_token_embed = nn.Embedding(self.mem_size + 3, self.dim, padding_idx=None)
|
||||
self.loss_fct = nn.CrossEntropyLoss(ignore_index=-100)
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(self.model_name, use_fast=False)
|
||||
self.append_sequence = torch.arange(
|
||||
self.vocab_size, self.vocab_size + self.mem_size, dtype=torch.long, device=device
|
||||
).unsqueeze(
|
||||
0
|
||||
) # mem tokens
|
||||
|
||||
if self.training:
|
||||
self.init()
|
||||
|
||||
def init(self):
|
||||
print("Freezing the decoder...")
|
||||
freeze_model(self.decoder)
|
||||
self.decoder.eval()
|
||||
print_trainable_parameters(self)
|
||||
if self.training_args.restore_from is not None and self.training_args.restore_from != "":
|
||||
print(f"Loading from the pretrained checkpoint: {self.training_args.restore_from}...")
|
||||
state_dict = load_file(self.training_args.restore_from)
|
||||
self.load_state_dict(state_dict)
|
||||
print(f"Finished loading from {self.training_args.restore_from}")
|
||||
print("Enabling gradient checkpointing...")
|
||||
# self.icae.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
|
||||
self.decoder.gradient_checkpointing_enable(
|
||||
gradient_checkpointing_kwargs={"use_reentrant": False}
|
||||
)
|
||||
|
||||
def compute_num_segments(self, total_length):
|
||||
assert total_length > 0
|
||||
num_segments = math.ceil(total_length / (self.mem_size * self.mean_compression_rate))
|
||||
return num_segments
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.LongTensor = None,
|
||||
prompt_answer_ids: torch.LongTensor = None,
|
||||
labels: Optional[torch.LongTensor] = None,
|
||||
):
|
||||
# encoder part
|
||||
batch_size = input_ids.size(0)
|
||||
total_length = input_ids.size(1)
|
||||
num_segments = self.compute_num_segments(total_length)
|
||||
segment_length = math.ceil(total_length / num_segments)
|
||||
|
||||
prompt_answer_embs = self.icae.get_base_model().model.embed_tokens(prompt_answer_ids)
|
||||
max_compressed_length = num_segments * self.mem_size
|
||||
compress_outputs = torch.zeros((max_compressed_length, self.dim)).to(prompt_answer_embs)
|
||||
|
||||
for segment_idx in range(num_segments):
|
||||
|
||||
start_idx = segment_idx * segment_length
|
||||
end_idx = min((segment_idx + 1) * segment_length, total_length)
|
||||
segment_input_ids = input_ids[:, start_idx:end_idx]
|
||||
segment_input_ids = torch.cat([segment_input_ids, self.append_sequence], dim=1)
|
||||
mem_flag = segment_input_ids >= self.vocab_size
|
||||
|
||||
segment_input_embedding = self.icae.get_base_model().model.embed_tokens(
|
||||
segment_input_ids
|
||||
)
|
||||
segment_input_embedding[mem_flag] = self.memory_token_embed(
|
||||
segment_input_ids[mem_flag] - self.vocab_size
|
||||
).to(segment_input_embedding)
|
||||
|
||||
# compress the current segment
|
||||
segment_compress_outputs = self.icae(
|
||||
inputs_embeds=segment_input_embedding, output_hidden_states=True
|
||||
)
|
||||
segment_compress_outputs = segment_compress_outputs.hidden_states[-1]
|
||||
|
||||
# collect memory tokens
|
||||
compress_outputs[segment_idx * self.mem_size : self.mem_size * (segment_idx + 1)] = (
|
||||
segment_compress_outputs[mem_flag]
|
||||
)
|
||||
|
||||
del segment_input_ids, segment_input_embedding
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# decoder part
|
||||
decoder_mem_flag = (prompt_answer_ids >= self.vocab_size) & (
|
||||
prompt_answer_ids < self.vocab_size + self.mem_size
|
||||
) # only mem tokens
|
||||
|
||||
prompt_answer_embs[decoder_mem_flag] = compress_outputs # replace memory slots
|
||||
special_prompt = prompt_answer_ids >= self.vocab_size_with_mem
|
||||
prompt_answer_embs[special_prompt] = self.memory_token_embed(
|
||||
prompt_answer_ids[special_prompt] - self.vocab_size
|
||||
).to(
|
||||
prompt_answer_embs
|
||||
) # replace special token's embedding from self.memory_token_embed
|
||||
|
||||
if self.training: # has an independent se.f.decoder
|
||||
decoder_outputs = self.decoder(
|
||||
inputs_embeds=prompt_answer_embs, output_hidden_states=True
|
||||
)
|
||||
else:
|
||||
with self.icae.disable_adapter(): # no independent decoder; use self.icae
|
||||
decoder_outputs = self.icae(
|
||||
inputs_embeds=prompt_answer_embs, output_hidden_states=True
|
||||
)
|
||||
|
||||
logits = decoder_outputs.logits
|
||||
effective_logits = logits[:, :-1, :].reshape(-1, logits.size(-1))
|
||||
target_ids = labels[:, 1:].reshape(-1)
|
||||
loss = self.loss_fct(effective_logits, target_ids)
|
||||
return {"loss": loss, "logits": logits}
|
||||
|
||||
def tokens_to_embeddings(
|
||||
self, token_ids
|
||||
): # input_tokens can be either normal tokens and special tokens
|
||||
embeddings = self.icae.get_base_model().model.embed_tokens(token_ids)
|
||||
special_flags = token_ids >= self.vocab_size
|
||||
embeddings[special_flags] = self.memory_token_embed(
|
||||
token_ids[special_flags] - self.vocab_size
|
||||
).to(
|
||||
embeddings
|
||||
) # replace special token's embedding from self.memory_token_embed
|
||||
return embeddings
|
||||
|
||||
def _compress(
|
||||
self, input_ids: torch.LongTensor = None
|
||||
): # for inference; compress a fixed length of input into memory slots
|
||||
|
||||
batch_size = input_ids.size(0)
|
||||
total_length = input_ids.size(1)
|
||||
num_segments = self.compute_num_segments(total_length)
|
||||
segment_length = math.ceil(total_length / num_segments)
|
||||
|
||||
max_compressed_length = num_segments * self.mem_size
|
||||
compress_outputs = torch.zeros((max_compressed_length, self.dim))
|
||||
|
||||
for segment_idx in range(num_segments):
|
||||
start_idx = segment_idx * segment_length
|
||||
end_idx = min((segment_idx + 1) * segment_length, total_length)
|
||||
segment_input_ids = input_ids[:, start_idx:end_idx]
|
||||
segment_input_ids = torch.cat([segment_input_ids, self.append_sequence], dim=1)
|
||||
mem_flag = segment_input_ids >= self.vocab_size
|
||||
|
||||
segment_input_embedding = self.icae.get_base_model().model.embed_tokens(
|
||||
segment_input_ids
|
||||
)
|
||||
segment_input_embedding[mem_flag] = self.memory_token_embed(
|
||||
segment_input_ids[mem_flag] - self.vocab_size
|
||||
).to(segment_input_embedding)
|
||||
|
||||
# compress the current segment
|
||||
segment_compress_outputs = self.icae(
|
||||
inputs_embeds=segment_input_embedding, output_hidden_states=True
|
||||
)
|
||||
segment_compress_outputs = segment_compress_outputs.hidden_states[-1]
|
||||
|
||||
# collect memory tokens
|
||||
compress_outputs[segment_idx * self.mem_size : self.mem_size * (segment_idx + 1)] = (
|
||||
segment_compress_outputs[mem_flag]
|
||||
)
|
||||
|
||||
del segment_input_ids, segment_input_embedding
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return compress_outputs
|
||||
|
|
@ -1,233 +0,0 @@
|
|||
import math
|
||||
import os
|
||||
import random
|
||||
|
||||
import torch
|
||||
from transformers import Trainer, Seq2SeqTrainer
|
||||
from transformers.trainer_utils import get_last_checkpoint
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
|
||||
def train_model(
|
||||
model,
|
||||
train_dataset,
|
||||
eval_dataset,
|
||||
training_args,
|
||||
data_collator=None,
|
||||
compute_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)
|
||||
|
||||
# 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,
|
||||
data_collator=data_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}")
|
||||
|
||||
train_result = trainer.train(resume_from_checkpoint=checkpoint)
|
||||
trainer.save_model()
|
||||
trainer.log_metrics("train", train_result.metrics)
|
||||
metrics = trainer.evaluate()
|
||||
trainer.log_metrics("eval", metrics)
|
||||
trainer.save_metrics("eval", metrics)
|
||||
|
||||
|
||||
def text_extraction(input_ids, length, lm_ratio=0.0):
|
||||
|
||||
input_len = len(input_ids)
|
||||
assert input_len >= 1, f"Error: invalid input length ({input_len})"
|
||||
|
||||
# ae
|
||||
if random.random() >= lm_ratio:
|
||||
if input_len <= length: # if shorter, keep the complete text
|
||||
return input_ids, []
|
||||
else:
|
||||
last_start = input_len - length
|
||||
random_start = random.randint(0, last_start)
|
||||
return input_ids[random_start : random_start + length], []
|
||||
|
||||
# lm
|
||||
if input_len <= length:
|
||||
r = random.randint(0, input_len - 1)
|
||||
return input_ids[: r + 1], input_ids[r + 1 :]
|
||||
else:
|
||||
last_start = input_len - length
|
||||
random_start = random.randint(0, last_start)
|
||||
return input_ids[random_start : random_start + length], input_ids[random_start + length :]
|
||||
|
||||
|
||||
def pretrain_tokenize_function(examples, model, mem, lm_ratio=0.0):
|
||||
text_output = model.tokenizer(
|
||||
examples["text"], truncation=False, padding=False, return_attention_mask=False
|
||||
)
|
||||
text_output["prompt_answer_ids"] = []
|
||||
text_output["labels"] = []
|
||||
|
||||
max_len = model.training_args.model_max_length # heuristic
|
||||
|
||||
for idx in range(len(text_output["input_ids"])):
|
||||
|
||||
ae = True
|
||||
a, b = text_extraction(text_output["input_ids"][idx], max_len, lm_ratio=lm_ratio)
|
||||
length_a = len(a)
|
||||
num_segments = model.compute_num_segments(length_a)
|
||||
total_mem_length = num_segments * model.mem_size
|
||||
|
||||
if (
|
||||
len(b) > model.training_args.min_tokens_for_lm
|
||||
): # avoid too few tokens for lm, which is a waste of computing
|
||||
ae = False
|
||||
b = b[:max_len]
|
||||
|
||||
text_output["input_ids"][idx] = a
|
||||
|
||||
# decoder part: note that in v2, we add mem_tokens to the prompt_ids
|
||||
# for easy implementation; which is different from v1 implementation
|
||||
# where mem tokens are not in the prompt_ids
|
||||
if ae: # autoencoding objective
|
||||
prompt_ids = [mem[0]] * total_mem_length + [model.ae_token_id]
|
||||
answer_ids = a + [model.eos_id] # if ae, eos token
|
||||
else: # lm objective
|
||||
prompt_ids = [mem[0]] * total_mem_length
|
||||
if model.training_args.add_special_token_for_lm:
|
||||
prompt_ids += [model.lm_token_id]
|
||||
answer_ids = b # if lm, no eos token
|
||||
|
||||
text_output["prompt_answer_ids"].append(prompt_ids + answer_ids)
|
||||
if ae:
|
||||
labels = [-100] * len(prompt_ids) + answer_ids
|
||||
else:
|
||||
labels = (
|
||||
[-100] * len(prompt_ids)
|
||||
+ [-100] * model.training_args.leave_tokens_for_lm
|
||||
+ answer_ids[model.training_args.leave_tokens_for_lm :]
|
||||
) # no loss for leave_tokens_for_lm
|
||||
text_output["labels"].append(labels)
|
||||
assert len(text_output["prompt_answer_ids"][-1]) == len(labels)
|
||||
|
||||
return text_output
|
||||
|
||||
|
||||
def instruct_ft_tokenize_function(examples, model, mem):
|
||||
text_output = model.tokenizer(
|
||||
examples["input"],
|
||||
max_length=5120,
|
||||
truncation=True,
|
||||
padding=False,
|
||||
return_attention_mask=False,
|
||||
add_special_tokens=False,
|
||||
)
|
||||
prompt_output = model.tokenizer(
|
||||
examples["prompt"],
|
||||
truncation=False,
|
||||
padding=False,
|
||||
return_attention_mask=False,
|
||||
add_special_tokens=False,
|
||||
)
|
||||
label_output = model.tokenizer(
|
||||
examples["answer"],
|
||||
truncation=False,
|
||||
padding=False,
|
||||
return_attention_mask=False,
|
||||
add_special_tokens=False,
|
||||
)
|
||||
text_output["prompt_answer_ids"] = []
|
||||
text_output["labels"] = []
|
||||
|
||||
max_len = model.training_args.model_max_length # heuristic
|
||||
|
||||
for idx in range(len(text_output["input_ids"])):
|
||||
|
||||
length = len(text_output["input_ids"][idx])
|
||||
num_segments = model.compute_num_segments(length)
|
||||
total_mem_length = num_segments * model.mem_size
|
||||
|
||||
prompt_ids = (
|
||||
[mem[0]] * total_mem_length + [model.ft_token_id] + prompt_output["input_ids"][idx]
|
||||
)
|
||||
prompt_ids = (
|
||||
[1, 733, 16289, 28793] + prompt_ids + [733, 28748, 16289, 28793]
|
||||
) # special formats for prompt in Mistral
|
||||
answer_ids = label_output["input_ids"][idx] + [model.eos_id]
|
||||
|
||||
text_output["prompt_answer_ids"].append(prompt_ids + answer_ids)
|
||||
|
||||
labels = [-100] * len(prompt_ids) + answer_ids
|
||||
text_output["labels"].append(labels)
|
||||
|
||||
assert len(text_output["prompt_answer_ids"][-1]) == len(labels)
|
||||
|
||||
return text_output
|
||||
|
||||
|
||||
class DataCollatorForDynamicPadding:
|
||||
def __init__(self, pad_token_id, pad_to_multiple_of=None):
|
||||
self.pad_token_id = pad_token_id
|
||||
self.pad_to_multiple_of = pad_to_multiple_of
|
||||
|
||||
def __call__(self, examples):
|
||||
input_ids = [torch.tensor(example["input_ids"], dtype=torch.long) for example in examples]
|
||||
labels = [torch.tensor(example["labels"], dtype=torch.long) for example in examples]
|
||||
prompt_answer_ids = [
|
||||
torch.tensor(example["prompt_answer_ids"], dtype=torch.long) for example in examples
|
||||
]
|
||||
input_ids = self.dynamic_padding(input_ids, fill_value=self.pad_token_id)
|
||||
prompt_answer_ids = self.dynamic_padding(prompt_answer_ids, fill_value=self.pad_token_id)
|
||||
labels = self.dynamic_padding(labels)
|
||||
batch = {"input_ids": input_ids, "labels": labels, "prompt_answer_ids": prompt_answer_ids}
|
||||
return batch
|
||||
|
||||
def dynamic_padding(self, sequences, fill_value=-100):
|
||||
max_length = max(len(x) for x in sequences)
|
||||
if self.pad_to_multiple_of:
|
||||
max_length = (
|
||||
(max_length - 1) // self.pad_to_multiple_of + 1
|
||||
) * self.pad_to_multiple_of
|
||||
padded_sequences = torch.full((len(sequences), max_length), fill_value, dtype=torch.long)
|
||||
for i, seq in enumerate(sequences):
|
||||
padded_sequences[i, : len(seq)] = seq
|
||||
return padded_sequences
|
||||
|
|
@ -1,4 +1,3 @@
|
|||
import transformers
|
||||
from datasets import load_dataset
|
||||
from modeling_icae_multi_span import (
|
||||
ICAE,
|
||||
|
|
@ -12,17 +11,54 @@ from training_utils import (
|
|||
instruct_ft_tokenize_function,
|
||||
train_model,
|
||||
)
|
||||
from transformers import HfArgumentParser
|
||||
|
||||
|
||||
def compute_metrics(eval_pred) -> dict:
|
||||
"""
|
||||
Custom metrics function for the trainer
|
||||
Args:
|
||||
eval_pred: tuple of predictions and labels
|
||||
Returns:
|
||||
dictionary containing metric names (str) and values (Any)
|
||||
"""
|
||||
|
||||
# preds, labels = eval_preds
|
||||
|
||||
# # predictions is generated tokens for Seq2SeqTrainer
|
||||
# # decode preds and labels
|
||||
# labels = np.where(labels != -100, labels, tokenizer.pad_token_id)
|
||||
# decoded_preds = tokenizer.batch_decode(preds, skip_special_tokens=True)
|
||||
# decoded_labels = tokenizer.batch_decode(labels, skip_special_tokens=True)
|
||||
|
||||
# compute per token accuracy
|
||||
predictions, labels = eval_pred.predictions, eval_pred.label_ids
|
||||
|
||||
# predictions is logits for Trainer
|
||||
preds = predictions.argmax(-1)
|
||||
acc = (preds == labels).mean()
|
||||
return {"per_token_acc": acc}
|
||||
|
||||
|
||||
def main():
|
||||
parser = transformers.HfArgumentParser(
|
||||
(ModelArguments, DataArguments, TrainingArguments)
|
||||
)
|
||||
parser = HfArgumentParser((ModelArguments, DataArguments, TrainingArguments))
|
||||
model_args, data_args, training_args = parser.parse_args_into_dataclasses()
|
||||
|
||||
print(model_args)
|
||||
print(data_args)
|
||||
|
||||
training_args.eval_on_start = True
|
||||
training_args.eval_strategy = "steps"
|
||||
training_args.eval_steps = 500
|
||||
training_args.save_strategy = "no"
|
||||
# training_args.save_steps = 500
|
||||
training_args.logging_strategy = "steps"
|
||||
training_args.logging_steps = 100
|
||||
|
||||
# seq2seq args for generation evaluation
|
||||
# training_args.predict_with_generate = True
|
||||
# training_args.generation_max_length = 100
|
||||
|
||||
training_args.gradient_checkpointing_kwargs = {
|
||||
"use_reentrant": False
|
||||
} # manually add this argument in the code
|
||||
|
|
@ -42,8 +78,8 @@ def main():
|
|||
|
||||
memory_size = training_args.fixed_mem_size
|
||||
|
||||
train_file = "/path/to/train/file"
|
||||
eval_file = "/path/to/dev/file"
|
||||
train_file = "../data/raw_datasets/context_numbers/train.jsonl"
|
||||
eval_file = "../data/raw_datasets/context_numbers/val.jsonl"
|
||||
|
||||
print("Loading dataset...")
|
||||
|
||||
|
|
@ -52,21 +88,32 @@ def main():
|
|||
train_dataset = dataset["train"]
|
||||
eval_dataset = dataset["eval"]
|
||||
|
||||
model = ICAE(model_args, training_args, lora_config)
|
||||
model = ICAE(model_args, training_args, lora_config).to("cuda")
|
||||
MEM_TOKENS = list(range(model.vocab_size, model.vocab_size + memory_size))
|
||||
|
||||
train_dataset = train_dataset.map(
|
||||
instruct_ft_tokenize_function,
|
||||
batched=True,
|
||||
fn_kwargs={"model": model, "mem": MEM_TOKENS},
|
||||
)
|
||||
eval_dataset = eval_dataset.map(
|
||||
tokenized_train_ds = train_dataset.map(
|
||||
instruct_ft_tokenize_function,
|
||||
batched=True,
|
||||
fn_kwargs={"model": model, "mem": MEM_TOKENS},
|
||||
)
|
||||
|
||||
train_model(model, train_dataset, eval_dataset, training_args)
|
||||
tokenized_eval_ds = {"train": train_dataset.select(range(100)), "val": eval_dataset}
|
||||
|
||||
for split in tokenized_eval_ds:
|
||||
tokenized_eval_ds[split] = tokenized_eval_ds[split].map(
|
||||
instruct_ft_tokenize_function,
|
||||
batched=True,
|
||||
fn_kwargs={"model": model, "mem": MEM_TOKENS},
|
||||
)
|
||||
|
||||
train_model(
|
||||
model,
|
||||
tokenized_train_ds,
|
||||
tokenized_eval_ds,
|
||||
training_args,
|
||||
compute_metrics=compute_metrics,
|
||||
)
|
||||
|
||||
|
||||
main()
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
|
|||
|
|
@ -119,9 +119,7 @@ class ICAE(torch.nn.Module):
|
|||
if self.training: # indepedent model for gradient checkpointing
|
||||
self.decoder = AutoModelForCausalLM.from_pretrained(
|
||||
self.model_name,
|
||||
torch_dtype=torch.float16
|
||||
if training_args.bf16 is False
|
||||
else torch.bfloat16,
|
||||
torch_dtype=torch.float16 if training_args.bf16 is False else torch.bfloat16,
|
||||
use_flash_attention_2=True,
|
||||
resume_download=True,
|
||||
)
|
||||
|
|
@ -152,9 +150,7 @@ class ICAE(torch.nn.Module):
|
|||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
self.memory_token_embed = nn.Embedding(
|
||||
self.mem_size + 3, self.dim, padding_idx=None
|
||||
)
|
||||
self.memory_token_embed = nn.Embedding(self.mem_size + 3, self.dim, padding_idx=None)
|
||||
self.loss_fct = nn.CrossEntropyLoss(ignore_index=-100)
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(self.model_name, use_fast=False)
|
||||
self.append_sequence = torch.arange(
|
||||
|
|
@ -174,13 +170,8 @@ class ICAE(torch.nn.Module):
|
|||
freeze_model(self.decoder)
|
||||
self.decoder.eval()
|
||||
print_trainable_parameters(self)
|
||||
if (
|
||||
self.training_args.restore_from is not None
|
||||
and self.training_args.restore_from != ""
|
||||
):
|
||||
print(
|
||||
f"Loading from the pretrained checkpoint: {self.training_args.restore_from}..."
|
||||
)
|
||||
if self.training_args.restore_from is not None and self.training_args.restore_from != "":
|
||||
print(f"Loading from the pretrained checkpoint: {self.training_args.restore_from}...")
|
||||
state_dict = load_file(self.training_args.restore_from)
|
||||
self.load_state_dict(state_dict)
|
||||
print(f"Finished loading from {self.training_args.restore_from}")
|
||||
|
|
@ -192,9 +183,7 @@ class ICAE(torch.nn.Module):
|
|||
|
||||
def compute_num_segments(self, total_length):
|
||||
assert total_length > 0
|
||||
num_segments = math.ceil(
|
||||
total_length / (self.mem_size * self.mean_compression_rate)
|
||||
)
|
||||
num_segments = math.ceil(total_length / (self.mem_size * self.mean_compression_rate))
|
||||
return num_segments
|
||||
|
||||
def forward(
|
||||
|
|
@ -209,22 +198,16 @@ class ICAE(torch.nn.Module):
|
|||
num_segments = self.compute_num_segments(total_length)
|
||||
segment_length = math.ceil(total_length / num_segments)
|
||||
|
||||
prompt_answer_embs = self.icae.get_base_model().model.embed_tokens(
|
||||
prompt_answer_ids
|
||||
)
|
||||
prompt_answer_embs = self.icae.get_base_model().model.embed_tokens(prompt_answer_ids)
|
||||
max_compressed_length = num_segments * self.mem_size
|
||||
compress_outputs = torch.zeros((max_compressed_length, self.dim)).to(
|
||||
prompt_answer_embs
|
||||
)
|
||||
compress_outputs = torch.zeros((max_compressed_length, self.dim)).to(prompt_answer_embs)
|
||||
|
||||
for segment_idx in range(num_segments):
|
||||
|
||||
start_idx = segment_idx * segment_length
|
||||
end_idx = min((segment_idx + 1) * segment_length, total_length)
|
||||
segment_input_ids = input_ids[:, start_idx:end_idx]
|
||||
segment_input_ids = torch.cat(
|
||||
[segment_input_ids, self.append_sequence], dim=1
|
||||
)
|
||||
segment_input_ids = torch.cat([segment_input_ids, self.append_sequence], dim=1)
|
||||
mem_flag = segment_input_ids >= self.vocab_size
|
||||
|
||||
segment_input_embedding = self.icae.get_base_model().model.embed_tokens(
|
||||
|
|
@ -241,9 +224,9 @@ class ICAE(torch.nn.Module):
|
|||
segment_compress_outputs = segment_compress_outputs.hidden_states[-1]
|
||||
|
||||
# collect memory tokens
|
||||
compress_outputs[
|
||||
segment_idx * self.mem_size : self.mem_size * (segment_idx + 1)
|
||||
] = segment_compress_outputs[mem_flag]
|
||||
compress_outputs[segment_idx * self.mem_size : self.mem_size * (segment_idx + 1)] = (
|
||||
segment_compress_outputs[mem_flag]
|
||||
)
|
||||
|
||||
del segment_input_ids, segment_input_embedding
|
||||
torch.cuda.empty_cache()
|
||||
|
|
@ -305,9 +288,7 @@ class ICAE(torch.nn.Module):
|
|||
start_idx = segment_idx * segment_length
|
||||
end_idx = min((segment_idx + 1) * segment_length, total_length)
|
||||
segment_input_ids = input_ids[:, start_idx:end_idx]
|
||||
segment_input_ids = torch.cat(
|
||||
[segment_input_ids, self.append_sequence], dim=1
|
||||
)
|
||||
segment_input_ids = torch.cat([segment_input_ids, self.append_sequence], dim=1)
|
||||
mem_flag = segment_input_ids >= self.vocab_size
|
||||
|
||||
segment_input_embedding = self.icae.get_base_model().model.embed_tokens(
|
||||
|
|
@ -324,9 +305,9 @@ class ICAE(torch.nn.Module):
|
|||
segment_compress_outputs = segment_compress_outputs.hidden_states[-1]
|
||||
|
||||
# collect memory tokens
|
||||
compress_outputs[
|
||||
segment_idx * self.mem_size : self.mem_size * (segment_idx + 1)
|
||||
] = segment_compress_outputs[mem_flag]
|
||||
compress_outputs[segment_idx * self.mem_size : self.mem_size * (segment_idx + 1)] = (
|
||||
segment_compress_outputs[mem_flag]
|
||||
)
|
||||
|
||||
del segment_input_ids, segment_input_embedding
|
||||
torch.cuda.empty_cache()
|
||||
|
|
|
|||
|
|
@ -3,13 +3,20 @@ import os
|
|||
import random
|
||||
|
||||
import torch
|
||||
from transformers import Trainer
|
||||
from transformers import Seq2SeqTrainer, Trainer
|
||||
from transformers.trainer_utils import get_last_checkpoint
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
|
||||
def train_model(model, train_dataset, eval_dataset, training_args, data_collator=None):
|
||||
def train_model(
|
||||
model,
|
||||
train_dataset,
|
||||
eval_dataset,
|
||||
training_args,
|
||||
data_collator=None,
|
||||
compute_metrics=None,
|
||||
):
|
||||
|
||||
last_checkpoint = None
|
||||
if (
|
||||
|
|
@ -19,7 +26,8 @@ def train_model(model, train_dataset, eval_dataset, training_args, data_collator
|
|||
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. "
|
||||
f"Output directory ({training_args.output_dir})"
|
||||
" already exists and is not empty. "
|
||||
"Use --overwrite_output_dir to overcome."
|
||||
)
|
||||
elif (
|
||||
|
|
@ -45,12 +53,18 @@ def train_model(model, train_dataset, eval_dataset, training_args, data_collator
|
|||
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,
|
||||
data_collator=data_collator,
|
||||
compute_metrics=compute_metrics,
|
||||
)
|
||||
|
||||
checkpoint = None
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue