formatting

This commit is contained in:
51616 2024-12-20 10:18:30 +00:00
parent e0e010301e
commit a38779abd7
6 changed files with 68 additions and 71 deletions

View file

@ -2,8 +2,8 @@ import dataclasses
import os
import sys
from dataclasses import dataclass, field
from typing import Any, Dict, List, Literal, NewType, Optional, Tuple
from enum import Enum, auto
from typing import Any, Dict, List, Literal, NewType, Optional, Tuple
from transformers import MODEL_FOR_CAUSAL_LM_MAPPING, HfArgumentParser
@ -40,7 +40,7 @@ class LoRAArguments:
default=0.05,
metadata={"help": ("LoRA dropout.")},
)
target_modules: Optional[List[str]] = field(
target_modules: Optional[list[str]] = field(
default=None,
metadata={"help": ("LoRA target modules.")},
)

View file

@ -1,8 +1,8 @@
import pandas as pd
from typing import Literal, List, Dict, Any, Tuple, Iterator, Callable, Optional, Union
from transformers import PreTrainedTokenizerBase
from typing import Any, Callable, Dict, Iterator, List, Literal, Optional, Tuple, Union
import pandas as pd
from training_utils import TRAINING_TASK
from transformers import PreTrainedTokenizerBase
IGNORE_INDEX = -100
@ -10,7 +10,7 @@ IGNORE_INDEX = -100
def get_sft_prompt_formatting_fn(
sft_mode: TRAINING_TASK,
tokenizer: PreTrainedTokenizerBase,
) -> Callable[[Dict[str, Any]], Dict[str, Any]]:
) -> Callable[[dict[str, Any]], dict[str, Any]]:
"""
Get a function that formats examples for supervised fine-tuning.
@ -70,9 +70,9 @@ def get_sft_prompt_formatting_fn(
def convert_ctx_prompt_response_to_messages(
example: Dict[str, Any],
example: dict[str, Any],
add_ctx_to_chat: bool = True,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
Convert context/prompt/response format to chat messages format.
@ -106,7 +106,7 @@ def convert_ctx_prompt_response_to_messages(
return example
def get_preprocessing_fn(ds_name: str) -> Callable[[Dict[str, Any]], Dict[str, Any]]:
def get_preprocessing_fn(ds_name: str) -> Callable[[dict[str, Any]], dict[str, Any]]:
"""
Get preprocessing function for a specific dataset.
@ -230,9 +230,9 @@ def get_preprocessing_fn(ds_name: str) -> Callable[[Dict[str, Any]], Dict[str, A
# taken from https://github.com/huggingface/trl/issues/632#issuecomment-1972630547
def get_assistant_start_end_indices(
messages: List[Dict[str, str]],
messages: list[dict[str, str]],
conversation_text: str,
) -> List[Tuple[int, int]]:
) -> list[tuple[int, int]]:
"""
Get the start and end indices of assistant messages in conversation text.
@ -263,8 +263,8 @@ def get_assistant_start_end_indices(
def get_masked_labels(
conversation_ids: Dict[str, List[Any]],
assistant_ranges: List[Tuple[int, int]],
conversation_ids: dict[str, list[Any]],
assistant_ranges: list[tuple[int, int]],
) -> Iterator[int]:
"""
Generate masked labels for conversation, masking non-assistant tokens.
@ -290,11 +290,11 @@ def get_masked_labels(
def tokenize_chat_messages(
example: Dict[str, Any],
example: dict[str, Any],
tokenizer: PreTrainedTokenizerBase,
mask_assistant_inputs: bool = True,
tokenizer_kwargs: Optional[Dict[str, Any]] = None,
) -> Dict[str, List[int]]:
tokenizer_kwargs: Optional[dict[str, Any]] = None,
) -> dict[str, list[int]]:
"""
Tokenize chat messages and optionally mask non-assistant tokens.

View file

@ -1,29 +1,27 @@
import logging
import numpy as np
import torch
from datasets import load_dataset
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
HfArgumentParser,
TrainingArguments,
DataCollatorForSeq2Seq,
EvalPrediction,
)
from configs import CtxTrainingArguments, LoRAArguments, ModelArguments, ExperimentSetup
from utils import log_num_train_params
from model_loading import get_model_and_tokenizer, get_lora_config
from modeling_utils import ModulatedPretrainedModel
from configs import CtxTrainingArguments, ExperimentSetup, LoRAArguments, ModelArguments
from data_utils import (
convert_ctx_prompt_response_to_messages,
get_preprocessing_fn,
get_sft_prompt_formatting_fn,
tokenize_chat_messages,
)
from datasets import load_dataset
from model_loading import get_lora_config, get_model_and_tokenizer
from modeling_utils import ModulatedPretrainedModel
from training_utils import TRAINING_TASK, train_model
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
DataCollatorForSeq2Seq,
EvalPrediction,
HfArgumentParser,
TrainingArguments,
)
from utils import log_num_train_params
logger = logging.getLogger(__name__)
@ -72,17 +70,6 @@ def main():
"use_reentrant": False
} # manually add this argument in the code
# "meta-llama/Llama-3.1-8B-Instruct",
# model = AutoModelForCausalLM.from_pretrained(
# base_model_name,
# torch_dtype=torch.bfloat16,
# attn_implementation="flash_attention_2",
# )
# tokenizer = AutoTokenizer.from_pretrained(base_model_name)
# tokenizer.pad_token_id = tokenizer.eos_token_id
# tokenizer.padding_side = "right"
model_name = model_args.model_name_or_path
model, tokenizer = get_model_and_tokenizer(
@ -135,12 +122,6 @@ def main():
"val": tokenized_ds["eval"],
}
# train_ds = dataset["train"].map(tokenize, batched=True)
# eval_ds = {
# "train": dataset["train"].select(range(100)).map(tokenize, batched=True),
# "val": dataset["eval"].map(tokenize, batched=True),
# }
# DataCollatorForSeq2Seq also pads the `labels`
# useful when we're computing the labels manually
# or masking the loss only on completion

View file

@ -2,13 +2,10 @@ import logging
import os
import torch
from peft import PeftModel, LoraConfig, VeraConfig, PeftConfig
from peft import LoraConfig, PeftConfig, PeftModel, VeraConfig
from peft import get_peft_config as _get_peft_config
from peft.utils import PeftType
from transformers import AutoModelForCausalLM, AutoTokenizer, AutoModel
from hyper_llm_modulator.utils.pooling import get_pooling_fn
from hyper_llm_modulator.utils.preprocessing import add_full_stop, apply_sfr_template
from transformers import AutoModel, AutoModelForCausalLM, AutoTokenizer
logger = logging.getLogger()

View file

@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Optional, Union, Tuple, Any
from typing import Any, Optional, Tuple, Union
import torch
from torch import nn
@ -46,7 +46,7 @@ class ModulatedPretrainedModel(nn.Module):
ctx_ids: Optional[torch.LongTensor] = None,
ctx_attention_mask: Optional[torch.LongTensor] = None,
**model_inputs_kwargs: dict[str, Any],
) -> Union[Tuple, ModelOutput]:
) -> Union[tuple, ModelOutput]:
"""Forward pass of the modulated model.
Args:

View file

@ -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,7 +152,9 @@ 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(
@ -170,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}")
@ -183,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(
@ -198,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(
@ -224,9 +241,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()
@ -288,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(
@ -305,9 +324,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()