diff --git a/hyperlora/dev_v2.jsonl b/hyperlora/dev_v2.jsonl new file mode 100644 index 0000000..e071b78 --- /dev/null +++ b/hyperlora/dev_v2.jsonl @@ -0,0 +1,6 @@ +{"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."} \ No newline at end of file diff --git a/hyperlora/fine_tuned_inference.py b/hyperlora/fine_tuned_inference.py new file mode 100644 index 0000000..4ec8f47 --- /dev/null +++ b/hyperlora/fine_tuned_inference.py @@ -0,0 +1,127 @@ +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") diff --git a/hyperlora/fine_tuned_inference_script.sh b/hyperlora/fine_tuned_inference_script.sh new file mode 100644 index 0000000..9de96f7 --- /dev/null +++ b/hyperlora/fine_tuned_inference_script.sh @@ -0,0 +1,16 @@ +#!/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 diff --git a/hyperlora/instruction_finetune.py b/hyperlora/instruction_finetune.py index c638a88..718a3e4 100644 --- a/hyperlora/instruction_finetune.py +++ b/hyperlora/instruction_finetune.py @@ -1,4 +1,3 @@ -import transformers from datasets import load_dataset from modeling_icae_multi_span import ( ICAE, @@ -12,21 +11,64 @@ 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 lora_config = LoraConfig( - r=model_args.lora_r, lora_alpha=32, lora_dropout=0.05, bias="none", task_type="CAUSAL_LM" + 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 @@ -36,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...") @@ -46,17 +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( - instruct_ft_tokenize_function, batched=True, fn_kwargs={"model": model, "mem": MEM_TOKENS} + 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() diff --git a/hyperlora/modeling_icae_multi_span.py b/hyperlora/modeling_icae_multi_span.py new file mode 100644 index 0000000..e1b3822 --- /dev/null +++ b/hyperlora/modeling_icae_multi_span.py @@ -0,0 +1,312 @@ +# 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 diff --git a/hyperlora/training_utils.py b/hyperlora/training_utils.py new file mode 100644 index 0000000..254e633 --- /dev/null +++ b/hyperlora/training_utils.py @@ -0,0 +1,233 @@ +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 diff --git a/icae_v2/fine_tuned_inference.py b/icae_v2/fine_tuned_inference.py index 4ec8f47..a494c02 100644 --- a/icae_v2/fine_tuned_inference.py +++ b/icae_v2/fine_tuned_inference.py @@ -22,7 +22,11 @@ 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" + 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 @@ -31,7 +35,9 @@ 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.load_state_dict( + state_dict, strict=False +) # only load lora and memory token embeddings model = model.to(device) @@ -73,7 +79,9 @@ with torch.no_grad(): # 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] + [model.ft_token_id] + + tokenized_prompt["input_ids"] + + [733, 28748, 16289, 28793] ) prompt_right_ids = torch.LongTensor([prompt_right_ids]).to(device) @@ -94,7 +102,9 @@ with torch.no_grad(): 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 + 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] diff --git a/icae_v2/instruction_finetune.py b/icae_v2/instruction_finetune.py index c638a88..8c096ef 100644 --- a/icae_v2/instruction_finetune.py +++ b/icae_v2/instruction_finetune.py @@ -15,7 +15,9 @@ from training_utils import ( def main(): - parser = transformers.HfArgumentParser((ModelArguments, DataArguments, TrainingArguments)) + parser = transformers.HfArgumentParser( + (ModelArguments, DataArguments, TrainingArguments) + ) model_args, data_args, training_args = parser.parse_args_into_dataclasses() print(model_args) @@ -26,7 +28,11 @@ def main(): } # 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" + 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 @@ -50,10 +56,14 @@ def main(): 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} + instruct_ft_tokenize_function, + batched=True, + fn_kwargs={"model": model, "mem": MEM_TOKENS}, ) eval_dataset = eval_dataset.map( - instruct_ft_tokenize_function, batched=True, fn_kwargs={"model": model, "mem": MEM_TOKENS} + instruct_ft_tokenize_function, + batched=True, + fn_kwargs={"model": model, "mem": MEM_TOKENS}, ) train_model(model, train_dataset, eval_dataset, training_args) diff --git a/icae_v2/modeling_icae_multi_span.py b/icae_v2/modeling_icae_multi_span.py index 5fa01a9..0edc37f 100644 --- a/icae_v2/modeling_icae_multi_span.py +++ b/icae_v2/modeling_icae_multi_span.py @@ -119,7 +119,9 @@ 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, ) @@ -150,11 +152,16 @@ 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( - self.vocab_size, self.vocab_size + self.mem_size, dtype=torch.long, device=device + self.vocab_size, + self.vocab_size + self.mem_size, + dtype=torch.long, + device=device, ).unsqueeze( 0 ) # mem tokens @@ -167,8 +174,13 @@ 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}") @@ -180,7 +192,9 @@ 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( @@ -195,16 +209,22 @@ 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( @@ -285,7 +305,9 @@ 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( diff --git a/icae_v2/pretrain.py b/icae_v2/pretrain.py index eaf3ecb..292b33e 100644 --- a/icae_v2/pretrain.py +++ b/icae_v2/pretrain.py @@ -15,7 +15,9 @@ from training_utils import ( def main(): - parser = transformers.HfArgumentParser((ModelArguments, DataArguments, TrainingArguments)) + parser = transformers.HfArgumentParser( + (ModelArguments, DataArguments, TrainingArguments) + ) model_args, data_args, training_args = parser.parse_args_into_dataclasses() print(model_args) @@ -26,7 +28,11 @@ def main(): } # 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" + 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 @@ -57,10 +63,16 @@ def main(): pretrain_tokenize_function, batched=True, batch_size=64, - fn_kwargs={"model": model, "mem": MEM_TOKENS, "lm_ratio": training_args.lm_ratio}, + fn_kwargs={ + "model": model, + "mem": MEM_TOKENS, + "lm_ratio": training_args.lm_ratio, + }, ) eval_dataset = eval_dataset.map( - pretrain_tokenize_function, batched=True, fn_kwargs={"model": model, "mem": MEM_TOKENS} + pretrain_tokenize_function, + batched=True, + fn_kwargs={"model": model, "mem": MEM_TOKENS}, ) # don't add lm in the dev set. data_collator = DataCollatorForDynamicPadding(model.pad_token_id) diff --git a/icae_v2/pretrained_inference.py b/icae_v2/pretrained_inference.py index e744321..7f5b443 100644 --- a/icae_v2/pretrained_inference.py +++ b/icae_v2/pretrained_inference.py @@ -22,7 +22,11 @@ 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" + 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 @@ -31,7 +35,9 @@ 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.load_state_dict( + state_dict, strict=False +) # only load lora and memory token embeddings model = model.to(device) @@ -47,7 +53,11 @@ with torch.no_grad(): for line in tqdm(lines): # Tokenize input text tokenized_text = model.tokenizer( - line, truncation=True, max_length=5120, padding=False, return_attention_mask=False + line, + truncation=True, + max_length=5120, + padding=False, + return_attention_mask=False, ) # Generate compressed outputs input_ids = torch.LongTensor([tokenized_text["input_ids"]]).to(device) @@ -72,7 +82,9 @@ with torch.no_grad(): for i in range(512): 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 + 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 diff --git a/icae_v2/training_utils.py b/icae_v2/training_utils.py index 9b3ca9f..39be345 100644 --- a/icae_v2/training_utils.py +++ b/icae_v2/training_utils.py @@ -12,14 +12,19 @@ device = torch.device("cuda" if torch.cuda.is_available() else "cpu") def train_model(model, train_dataset, eval_dataset, training_args, data_collator=None): last_checkpoint = None - if os.path.isdir(training_args.output_dir) and not training_args.overwrite_output_dir: + 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: + 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 " @@ -27,7 +32,10 @@ def train_model(model, train_dataset, eval_dataset, training_args, data_collator ) if ( - max(training_args.per_device_train_batch_size, training_args.per_device_eval_batch_size) + max( + training_args.per_device_train_batch_size, + training_args.per_device_eval_batch_size, + ) == 1 ): data_collator = None @@ -83,7 +91,10 @@ def text_extraction(input_ids, length, lm_ratio=0.0): 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 :] + return ( + input_ids[random_start : random_start + length], + input_ids[random_start + length :], + ) def pretrain_tokenize_function(examples, model, mem, lm_ratio=0.0): @@ -173,7 +184,9 @@ def instruct_ft_tokenize_function(examples, model, mem): total_mem_length = num_segments * model.mem_size prompt_ids = ( - [mem[0]] * total_mem_length + [model.ft_token_id] + prompt_output["input_ids"][idx] + [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] @@ -196,15 +209,26 @@ class DataCollatorForDynamicPadding: 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] + 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 + 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) + 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} + batch = { + "input_ids": input_ids, + "labels": labels, + "prompt_answer_ids": prompt_answer_ids, + } return batch def dynamic_padding(self, sequences, fill_value=-100): @@ -213,7 +237,9 @@ class DataCollatorForDynamicPadding: 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) + 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