doc-to-lora/icae_v2/modeling_icae_multi_span.py
2024-12-20 10:18:30 +00:00

334 lines
12 KiB
Python

# 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.TrainingArguments):
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