Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -8,5 +8,6 @@ wheels/

# Virtual environments
.venv
model_checkpoints/
checkpoints/
data/
model_checkpoints/
8 changes: 8 additions & 0 deletions check.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
from transformers import PreTrainedTokenizerFast, AutoTokenizer
from src.utils.paths import TOKENIZER_DIR
from pathlib import Path


tok = AutoTokenizer.from_pretrained("tokenizer_checkpoint")

# tok = AutoTokenizer.from_pretrained(TOKENIZER_DIR)
2 changes: 1 addition & 1 deletion configs/datashard.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
from pathlib import Path


from src.paths import TOKENIZER_DIR, DATA_DIR
from src.utils.paths import TOKENIZER_DIR, DATA_DIR


@dataclass
Expand Down
44 changes: 44 additions & 0 deletions configs/ft_config.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,44 @@
from pathlib import Path
import torch
from trl import SFTConfig
from src.utils.paths import FINETUNE_OUT_DIR

output_dir = FINETUNE_OUT_DIR / "alpaca_model"


sft_config = SFTConfig(
output_dir=Path(output_dir),
assistant_only_loss=True,
max_length=512,

per_device_train_batch_size=16,
per_device_eval_batch_size=8,

gradient_accumulation_steps=1,

learning_rate=5e-5,
weight_decay=0.01,
lr_scheduler_type="cosine",
warmup_steps=0.03,
num_train_epochs=3,

logging_steps=10,
eval_strategy="steps",
eval_steps=100,
save_strategy="steps",
save_steps=100,
save_total_limit=2,
gradient_checkpointing=False,
metric_for_best_model="eval_loss",
greater_is_better=False,
load_best_model_at_end=True,

dataloader_num_workers=4,
dataset_kwargs={"num_proc": 4},

bf16=torch.cuda.is_bf16_supported(),
fp16=not torch.cuda.is_bf16_supported(),
report_to="none"
)


65 changes: 30 additions & 35 deletions configs/model.py
Original file line number Diff line number Diff line change
@@ -1,39 +1,34 @@
from dataclasses import dataclass
from transformers import PretrainedConfig


@dataclass
class ModelConfig:
"""Hyperparameters for the BetterGPT model architecture.
class BetterGPTConfig(PretrainedConfig):
model_type = "better_gpt"

Validated at construction time: emb_dim must be divisible by head_count,
and the resulting head_dim must be even (required by RoPE).
"""
vocab_size: int = 8192
emb_dim: int = 512
num_blocks: int = 8
head_count: int = 8
seq_length: int = 512
ffn_multiple: int = 128
def __init__(
self,
vocab_size=8192,
emb_dim=512,
num_blocks=8,
head_count=8,
seq_length=512,
ffn_multiple=128,
tie_word_embeddings=True,
rmsnorm_eps=1e-6,
rope_base=10000,
**kwargs
):
assert emb_dim % head_count == 0, "emb_dim must be divisible by head_count"
self.vocab_size = vocab_size
self.emb_dim = emb_dim
self.num_blocks = num_blocks
self.head_count = head_count
self.rmsnorm_eps = rmsnorm_eps
self.seq_length = seq_length
self.ffn_multiple = ffn_multiple
self.rope_base = rope_base
self.head_dim = self.emb_dim // self.head_count
self.architectures = ["BetterGPTModel"]

# Initialize the Hugging Face base config
super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)

def __post_init__(self):
if self.vocab_size <= 0:
raise ValueError(f"vocab_size must be > 0, got {self.vocab_size}")
if self.emb_dim <= 0:
raise ValueError(f"emb_dim must be > 0, got {self.emb_dim}")
if self.head_count <= 0:
raise ValueError(f"head_count must be > 0, got {self.head_count}")
if self.emb_dim % self.head_count != 0:
raise ValueError(
f"emb_dim ({self.emb_dim}) must be divisible by head_count ({self.head_count})"
)
head_dim = self.emb_dim // self.head_count
if head_dim % 2 != 0:
raise ValueError(
f"head_dim ({head_dim}) must be even for RoPE; adjust emb_dim or head_count"
)
if self.num_blocks <= 0:
raise ValueError(f"num_blocks must be > 0, got {self.num_blocks}")
if self.seq_length <= 0:
raise ValueError(f"seq_length must be > 0, got {self.seq_length}")
if self.ffn_multiple <= 0:
raise ValueError(f"ffn_multiple must be > 0, got {self.ffn_multiple}")
17 changes: 17 additions & 0 deletions configs/template.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@

chat_template = """
{%- for message in messages -%}
{% if message['role'] == 'system' %}
<|im_start|><|system|>
{{ message['content'] }}<|im_end|>
{% elif message['role'] == 'user' %}
<|im_start|><|user|>
{{ message['content'] }}<|im_end|>
{% elif message['role'] == 'assistant' %}
<|im_start|><|assistant|>
{% generation %}{{ message['content'] }}{{ eos_token }}{% endgeneration %}
{% endif %}
{% endfor %}
{% if add_generation_prompt %}<|im_start|><|assistant|>
{% endif %}
"""
2 changes: 1 addition & 1 deletion configs/training.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
import torch
from dataclasses import dataclass

from src.paths import CHECKPOINT_DIR
from src.utils.paths import CHECKPOINT_DIR

checkpoint_dir = CHECKPOINT_DIR

Expand Down
4 changes: 2 additions & 2 deletions scripts/create_data_shards.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
from tokenizers import Tokenizer
from transformers import PreTrainedTokenizerFast, AutoTokenizer

from src.paths import DATA_DIR,TOKENIZER_DIR
from src.utils.paths import DATA_DIR,TOKENIZER_DIR
from src.data_preparation.make_shards import ShardDataset
from src.logger import get_logger
from src.utils.logger import get_logger


logger = get_logger(__name__)
Expand Down
90 changes: 90 additions & 0 deletions scripts/finetune_alpaca.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,90 @@
import torch

from trl import SFTTrainer
from transformers import EarlyStoppingCallback
from transformers import PreTrainedTokenizerFast,EarlyStoppingCallback
from datasets import load_dataset

from configs.ft_config import sft_config
from src.utils.model_loader import load_model
from src.utils.ft_tokenizer_helper import prepare_tokenizer
from src.utils.prepare_alpaca_dataset import convert_alpaca_format
from src.utils.paths import CHECKPOINT_DIR, TOKENIZER_DIR
from src.utils.data_splitter import split_dataset
from configs.tokenizer_config import TokenizerConfig
from src.utils.logger import get_logger


logger = get_logger(__name__)

# model_path = CHECKPOINT_DIR
# tokenizer_path = TOKENIZER_DIR
model_path = "./checkpoints/model.pt"
tokenizer_path = "./tokenizer_checkpoint"

def main():
try:
logger.info("Starting fine-tuning process")
dataset = load_dataset("yahma/alpaca-cleaned", split="train")
logger.info("Alpaca Dataset loaded successfully")
except Exception as e:
logger.error(f"Error occurred while loading dataset or tokenizer: {e}")
raise
try:
dataset = dataset.map(convert_alpaca_format, batched=True)
logger.info("Alpaca Dataset converted to chat format successfully")
except Exception as e:
logger.error(f"Error occurred while converting dataset to chat format: {e}")
raise
try:
train_dataset, eval_dataset = split_dataset(dataset)
logger.info("Alpaca Dataset split into training and validation sets successfully")
except Exception as e:
logger.error(f"Error occurred while splitting dataset: {e}")
raise
try:
tokenizer = PreTrainedTokenizerFast.from_pretrained(tokenizer_path)
logger.info("Tokenizer loaded successfully")
except Exception as e:
logger.error(f"Error occurred while loading tokenizer: {e}")
raise
try:
tokenizer = prepare_tokenizer(tokenizer)
logger.info("Tokenizer prepared successfully")
except Exception as e:
logger.error(f"Error occurred while preparing tokenizer: {e}")
raise
try:
model = load_model(model_path, tokenizer)
logger.info("Model retied and loaded successfully")
except Exception as e:
logger.error(f"Error occurred while loading model: {e}")
raise
return model, train_dataset, eval_dataset, tokenizer


if __name__ == "__main__":
model, train_dataset, eval_dataset, tokenizer = main()

model = load_model(model_path,tokenizer)

num_params = sum(p.numel() for p in model.parameters())
logger.info(f"{num_params/1e6:.2f}M parameters")

trainer = SFTTrainer(
model=model,
args=sft_config,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
processing_class=tokenizer,
callbacks=[EarlyStoppingCallback(early_stopping_patience=3)]
)

trainer.train()
logger.info("Fine-tuning completed successfully. Saving model and tokenizer...")

trainer.save_model(model_path)
tokenizer.save_pretrained(tokenizer_path)



4 changes: 2 additions & 2 deletions scripts/sample_pretrain.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,8 @@

from configs.model import ModelConfig
from configs.training import TrainingConfig
from src.models.model import BetterGPT
from src.logger import get_logger
from src.models.base_model import BetterGPT
from src.utils.logger import get_logger

logger = get_logger("sample")

Expand Down
6 changes: 3 additions & 3 deletions scripts/train_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,12 +4,12 @@

from transformers import get_cosine_schedule_with_warmup

from src.models.model import BetterGPT
from configs.model import ModelConfig
from src.models.base_model import BetterGPT
from configs.model import BetterGPTConfig as ModelConfig
from configs.training import TrainingConfig
from src.data_preparation.data_loader import train_loader, val_loader
from src.pretraining.trainer import training
from src.logger import get_logger
from src.utils.logger import get_logger

logger = get_logger("train")

Expand Down
5 changes: 2 additions & 3 deletions scripts/train_tokenizer.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,8 @@
import logging
from datasets import load_dataset
from src.paths import TOKENIZER_DIR
from src.utils.paths import TOKENIZER_DIR
from src.tokenizer import TokenizerTrainer
from configs.tokenizer_config import TokenizerConfig
from src.logger import get_logger
from src.utils.logger import get_logger


logger = get_logger(__name__)
Expand Down
4 changes: 2 additions & 2 deletions src/data_preparation/data_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,8 @@
from torch.utils.data import DataLoader
from configs.model import ModelConfig
from configs.training import TrainingConfig
from src.paths import DATA_DIR
from src.logger import get_logger
from src.utils.paths import DATA_DIR
from src.utils.logger import get_logger

logger = get_logger(__name__)

Expand Down
2 changes: 1 addition & 1 deletion src/data_preparation/make_shards.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
from tokenizers import Tokenizer
from transformers import AutoTokenizer

from src.logger import get_logger
from src.utils.logger import get_logger
from configs.datashard import DatasetConfig
from configs.tokenizer_config import TokenizerConfig

Expand Down
Loading
Loading