Skip to content

Commit 90aee78

Browse files
author
Harikrishnan V
committed
finetune alpaca
1 parent 2a05a72 commit 90aee78

29 files changed

Lines changed: 539 additions & 297 deletions

configs/datashard.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,7 @@
33
from pathlib import Path
44

55

6-
from src.paths import TOKENIZER_DIR, DATA_DIR
6+
from src.utils.paths import TOKENIZER_DIR, DATA_DIR
77

88

99
@dataclass

configs/ft_config.py

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
from pathlib import Path
2+
import torch
3+
from trl import SFTConfig
4+
from src.utils.paths import FINETUNE_OUT_DIR
5+
6+
output_dir = FINETUNE_OUT_DIR / "alpaca_model"
7+
8+
9+
sft_config = SFTConfig(
10+
output_dir=Path(output_dir),
11+
assistant_only_loss=True,
12+
max_length=512,
13+
14+
per_device_train_batch_size=16,
15+
per_device_eval_batch_size=8,
16+
17+
gradient_accumulation_steps=1,
18+
19+
learning_rate=5e-5,
20+
weight_decay=0.01,
21+
lr_scheduler_type="cosine",
22+
warmup_steps=0.03,
23+
num_train_epochs=3,
24+
25+
logging_steps=10,
26+
eval_strategy="steps",
27+
eval_steps=100,
28+
save_strategy="steps",
29+
save_steps=100,
30+
save_total_limit=2,
31+
gradient_checkpointing=False,
32+
metric_for_best_model="eval_loss",
33+
greater_is_better=False,
34+
load_best_model_at_end=True,
35+
36+
dataloader_num_workers=4,
37+
dataset_kwargs={"num_proc": 4},
38+
39+
bf16=torch.cuda.is_bf16_supported(),
40+
fp16=not torch.cuda.is_bf16_supported(),
41+
report_to="none"
42+
)
43+
44+

configs/model.py

Lines changed: 24 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -1,39 +1,28 @@
1-
from dataclasses import dataclass
1+
from transformers import PretrainedConfig
22

33

4-
@dataclass
5-
class ModelConfig:
6-
"""Hyperparameters for the BetterGPT model architecture.
4+
class BetterGPTConfig(PretrainedConfig):
5+
model_type = "better_gpt"
76

8-
Validated at construction time: emb_dim must be divisible by head_count,
9-
and the resulting head_dim must be even (required by RoPE).
10-
"""
11-
vocab_size: int = 8192
12-
emb_dim: int = 512
13-
num_blocks: int = 8
14-
head_count: int = 8
15-
seq_length: int = 512
16-
ffn_multiple: int = 128
7+
def __init__(
8+
self,
9+
vocab_size=8192,
10+
emb_dim=512,
11+
num_blocks=8,
12+
head_count=8,
13+
seq_length=512,
14+
ffn_multiple=128,
15+
tie_word_embeddings=True,
16+
rmsnorm_eps=1e-6,
17+
**kwargs
18+
):
19+
assert emb_dim % head_count == 0, "emb_dim must be divisible by head_count"
20+
self.vocab_size = vocab_size
21+
self.emb_dim = emb_dim
22+
self.num_blocks = num_blocks
23+
self.head_count = head_count
24+
self.seq_length = seq_length
25+
self.ffn_multiple = ffn_multiple
1726

18-
def __post_init__(self):
19-
if self.vocab_size <= 0:
20-
raise ValueError(f"vocab_size must be > 0, got {self.vocab_size}")
21-
if self.emb_dim <= 0:
22-
raise ValueError(f"emb_dim must be > 0, got {self.emb_dim}")
23-
if self.head_count <= 0:
24-
raise ValueError(f"head_count must be > 0, got {self.head_count}")
25-
if self.emb_dim % self.head_count != 0:
26-
raise ValueError(
27-
f"emb_dim ({self.emb_dim}) must be divisible by head_count ({self.head_count})"
28-
)
29-
head_dim = self.emb_dim // self.head_count
30-
if head_dim % 2 != 0:
31-
raise ValueError(
32-
f"head_dim ({head_dim}) must be even for RoPE; adjust emb_dim or head_count"
33-
)
34-
if self.num_blocks <= 0:
35-
raise ValueError(f"num_blocks must be > 0, got {self.num_blocks}")
36-
if self.seq_length <= 0:
37-
raise ValueError(f"seq_length must be > 0, got {self.seq_length}")
38-
if self.ffn_multiple <= 0:
39-
raise ValueError(f"ffn_multiple must be > 0, got {self.ffn_multiple}")
27+
# Initialize the Hugging Face base config
28+
super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)

configs/training.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
import torch
22
from dataclasses import dataclass
33

4-
from src.paths import CHECKPOINT_DIR
4+
from src.utils.paths import CHECKPOINT_DIR
55

66
checkpoint_dir = CHECKPOINT_DIR
77

scripts/create_data_shards.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,9 @@
11
from tokenizers import Tokenizer
22
from transformers import PreTrainedTokenizerFast, AutoTokenizer
33

4-
from src.paths import DATA_DIR,TOKENIZER_DIR
4+
from src.utils.paths import DATA_DIR,TOKENIZER_DIR
55
from src.data_preparation.make_shards import ShardDataset
6-
from src.logger import get_logger
6+
from src.utils.logger import get_logger
77

88

99
logger = get_logger(__name__)

scripts/sample_pretrain.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,8 +4,8 @@
44

55
from configs.model import ModelConfig
66
from configs.training import TrainingConfig
7-
from src.models.model import BetterGPT
8-
from src.logger import get_logger
7+
from src.models.base_model import BetterGPT
8+
from src.utils.logger import get_logger
99

1010
logger = get_logger("sample")
1111

scripts/train_model.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,12 +4,12 @@
44

55
from transformers import get_cosine_schedule_with_warmup
66

7-
from src.models.model import BetterGPT
8-
from configs.model import ModelConfig
7+
from src.models.base_model import BetterGPT
8+
from configs.model import BetterGPTConfig as ModelConfig
99
from configs.training import TrainingConfig
1010
from src.data_preparation.data_loader import train_loader, val_loader
1111
from src.pretraining.trainer import training
12-
from src.logger import get_logger
12+
from src.utils.logger import get_logger
1313

1414
logger = get_logger("train")
1515

scripts/train_tokenizer.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,8 @@
1-
import logging
21
from datasets import load_dataset
3-
from src.paths import TOKENIZER_DIR
2+
from src.utils.paths import TOKENIZER_DIR
43
from src.tokenizer import TokenizerTrainer
54
from configs.tokenizer_config import TokenizerConfig
6-
from src.logger import get_logger
5+
from src.utils.logger import get_logger
76

87

98
logger = get_logger(__name__)

src/data_preparation/data_loader.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -3,8 +3,8 @@
33
from torch.utils.data import DataLoader
44
from configs.model import ModelConfig
55
from configs.training import TrainingConfig
6-
from src.paths import DATA_DIR
7-
from src.logger import get_logger
6+
from src.utils.paths import DATA_DIR
7+
from src.utils.logger import get_logger
88

99
logger = get_logger(__name__)
1010

src/data_preparation/make_shards.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
from tokenizers import Tokenizer
66
from transformers import AutoTokenizer
77

8-
from src.logger import get_logger
8+
from src.utils.logger import get_logger
99
from configs.datashard import DatasetConfig
1010
from configs.tokenizer_config import TokenizerConfig
1111

0 commit comments

Comments
 (0)