doc-to-lora/icae_v2/modeling_icae_multi_span.py
2025-05-27 21:18:15 +09:00

334 lines
12 KiB
Python

# ICAE that supports multi span concat
import math
from dataclasses import dataclass, field
import torch
import torch.nn as nn
import transformers
from peft import get_peft_model
from safetensors.torch import load_file
from transformers import 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: str | None = 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: torch.LongTensor | None = 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