mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
334 lines
12 KiB
Python
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
|