File size: 2,552 Bytes
701cf7d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Pretrain Ares from scratch - iterative scaling
"""
import sys
sys.path.append("src")
import os
import torch
from ares.config import get_config
from ares.model.model import AresForCausalLM
from ares.tokenizer.tokenizer import AresTokenizer
from ares.training.dataset_pipeline import DataPipeline
from ares.training.trainer import AresTrainer
from torch.utils.data import Dataset, DataLoader

class TokenizedDataset(Dataset):
    def __init__(self, token_chunks):
        self.chunks = token_chunks
    def __len__(self): return len(self.chunks)
    def __getitem__(self, idx):
        chunk = self.chunks[idx]
        # input = chunk, label = chunk (shift inside model)
        return torch.tensor(chunk, dtype=torch.long), torch.tensor(chunk, dtype=torch.long)

def main():
    import argparse
    parser = argparse.ArgumentParser()
    parser.add_argument("--config", type=str, default="tiny", choices=["tiny","small","medium","billion","large"])
    parser.add_argument("--tokenizer", type=str, default="data/tokenizer.json")
    parser.add_argument("--steps", type=int, default=500)
    parser.add_argument("--max_seq", type=int, default=512)
    args = parser.parse_args()

    config = get_config(args.config)
    print(f"[Ares] Config {args.config}: ~{config.num_parameters_approx/1e6:.1f}M params")

    tokenizer = AresTokenizer(vocab_file=args.tokenizer if os.path.exists(args.tokenizer) else None, vocab_size=config.vocab_size)
    print(f"[Ares] Tokenizer vocab {len(tokenizer.vocab)}")

    model = AresForCausalLM(config)
    print(f"[Ares] Actual params {model.count_parameters()/1e6:.2f}M")

    pipeline = DataPipeline(tokenizer, max_seq_len=args.max_seq)
    texts = pipeline.load_hf_dataset("allenai/c4", num_samples=20000)
    # Limit chunks for demo
    chunks = []
    for chunk in pipeline.tokenized_stream(texts):
        chunks.append(chunk)
        if len(chunks) >= 5000:
            break
    print(f"[Data] Got {len(chunks)} chunks")

    dataset = TokenizedDataset(chunks)

    # DataLoader that yields tensors
    def collate(batch):
        input_ids = torch.stack([b[0] for b in batch])
        labels = torch.stack([b[1] for b in batch])
        return input_ids, labels

    loader = DataLoader(dataset, batch_size=2, shuffle=True, collate_fn=collate)

    trainer = AresTrainer(model, tokenizer, config, device="cuda" if torch.cuda.is_available() else "cpu")
    trainer.train(loader, total_steps=args.steps, warmup_steps=50, lr=3e-4, save_dir=f"checkpoints/{args.config}")

if __name__ == "__main__":
    main()