-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
110 lines (89 loc) · 3.53 KB
/
Copy pathmain.py
File metadata and controls
110 lines (89 loc) · 3.53 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
import os
from itertools import count
from pathlib import Path
import torch
from torch.utils.data import DataLoader
from tokenizers import Tokenizer
from config.arguments import ArgumentParser
from data.dataset_manager import DatasetManager
from data.dataset import ValidationDataset
from training.distributed_setup import DistributedTrainingSetup
from training.model_setup import ModelSetup
from training.training_loop import TrainingLoop
from training.checkpoint_manager import CheckpointManager
from training.validation import ValidationManager
from utils.utils import is_main_process
from model import Bert
def main():
"""Main training function."""
# Parse arguments and setup configuration
args = ArgumentParser.parse_and_setup()
# Setup tokenizer and distributed training environment
tokenizer = Tokenizer.from_file(str(args.tokenizer_path))
DistributedTrainingSetup.setup(args, tokenizer)
# Prepare model, optimizer, scheduler, and EMA
model, ema, optimizer, scheduler, global_step, start_epoch = ModelSetup.prepare_model_and_optimizer(args)
# Initialize dataloaders (will be populated during training)
train_dataloader, valid_dataloader = None, None
torch.cuda.empty_cache()
# Setup development dataloaders for evaluation
dev_paths = [
'../babyLM_2025_data/dev/bnc_spoken_100M-correct_tokenized.bin',
'../babyLM_2025_data/dev/childes_100M-correct_tokenized.bin',
'../babyLM_2025_data/dev/gutenberg_100M-correct_tokenized.bin',
'../babyLM_2025_data/dev/open_subtitles_100M-correct_tokenized.bin',
'../babyLM_2025_data/dev/simple_wiki_100M-correct_tokenized.bin',
'../babyLM_2025_data/dev/switchboard_100M-correct_tokenized.bin'
]
dev_dataloaders = []
for dev_path in dev_paths:
try:
dev_data = ValidationDataset(dev_path, tokenizer, args)
dev_dataloader = DataLoader(
dev_data,
shuffle=False,
batch_size=16,
generator=torch.Generator().manual_seed(42),
drop_last=True,
pin_memory=True
)
dev_dataloaders.append(dev_dataloader)
except Exception as e:
print(f"Warning: Could not load dev dataset {dev_path}: {e}")
# Initialize frequency masking distribution (if used)
token_freq_distribution = None
# Main training loop over epochs
for epoch in count(start=start_epoch):
# Load datasets for current epoch
train_dataloader, valid_dataloader, token_freq_distribution = DatasetManager.load_datasets(
args,
tokenizer,
epoch,
global_step,
train_dataloader,
valid_dataloader,
token_freq_distribution
)
# Run training epoch
global_step = TrainingLoop.training_epoch_vector(
model,
ema,
train_dataloader,
valid_dataloader,
optimizer,
scheduler,
global_step,
epoch,
args,
tokenizer,
dev_dataloaders=dev_dataloaders,
)
# Check if training is complete
if global_step >= args.max_steps:
break
# Save final checkpoint
CheckpointManager.save_checkpoint(model, ema, optimizer, scheduler, global_step, args)
if is_main_process():
print("Training completed successfully!")
if __name__ == "__main__":
main()