kgrabko commited on
Commit
46e36df
·
verified ·
1 Parent(s): 4c19f28

Upload train_236b_heavy_mixed_val_data.py

Browse files
Files changed (1) hide show
  1. train_236b_heavy_mixed_val_data.py +154 -0
train_236b_heavy_mixed_val_data.py ADDED
@@ -0,0 +1,154 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # ==============================================================================
2
+ # COPYRIGHT (C) 2025 KONSTANTIN VLADIMIROVICH GRABKO. ALL RIGHTS RESERVED.
3
+ # PATENT PENDING | CMS MANHATTAN JIRACK TECHNOLOGY
4
+ #
5
+ # This software is licensed under the Commercial License Agreement V.1.2.
6
+ # Any use, modification, or distribution of this code requires compliance with
7
+ # the terms found in the LICENSE.md file in the root directory.
8
+ #
9
+ # NO PATENTING RIGHTS: Users are strictly prohibited from filing patent claims
10
+ # based on the BRE or SWA architectures disclosed herein.
11
+ # Contact: grabko@cmsmanhattan.com | +1 (516) 777-0945
12
+ # ==============================================================================
13
+ # COPYRIGHT (C) 2025 KONSTANTIN VLADIMIROVICH GRABKO. ALL RIGHTS RESERVED.
14
+ # PATENT PENDING | CMS MANHATTAN JIRACK TECHNOLOGY | VERSION 236B MIXED
15
+ # Optimized for Extreme Depth (192 Layers) & Hybrid Knowledge
16
+ # ==============================================================================
17
+
18
+ import torch
19
+ import torch.nn as nn
20
+ import os
21
+ import random
22
+ import json
23
+ from torch.utils.data import DataLoader, IterableDataset
24
+ from transformers import AutoTokenizer
25
+ from datasets import load_dataset
26
+ from accelerate import Accelerator
27
+ import sys
28
+
29
+ # Импорт вашей архитектуры 236B
30
+ from JiRackTernaryPyTorch_236b import JiRackTernary236B, JiRackTernaryConfig
31
+
32
+ # --- КОНФИГУРАЦИЯ CMS MANHATTAN ---
33
+ MODEL_ID = "./models/jirack_236b_init"
34
+ CULTURAL_DATA_FILE = "cultural_finetune.jsonl"
35
+ GENERAL_DATA_LINK = "monology/pile-uncopyrighted" # Ссылка на The Pile
36
+ CHECKPOINT_DIR = "checkpoints_jirack_236b_mixed"
37
+
38
+ MIX_RATIO = 0.35 # 35% Культурный код / 65% The Pile
39
+ BATCH_SIZE = 1
40
+ GRAD_ACCUM_STEPS = 48 # Баланс между скоростью и стабильностью для 236B
41
+ LEARNING_RATE = 3.5e-6 # Специфический LR для 192 слоев
42
+ BLOCK_SIZE = 2048 # 2k контекст
43
+
44
+ # --- МИКСЕР ДАННЫХ ДЛЯ 236B ---
45
+ class CMSDataMixer236B(IterableDataset):
46
+ def __init__(self, tokenizer, client_file, pile_link, mix_ratio=0.35):
47
+ self.tokenizer = tokenizer
48
+ self.mix_ratio = mix_ratio
49
+
50
+ # Стриминг The Pile (Общие знания)
51
+ print(f">>> [MIXER] Connecting to General Knowledge: {pile_link}")
52
+ self.pile_stream = load_dataset(pile_link, split="train", streaming=True)
53
+
54
+ # Загрузка вашего Эволюционного Индекса
55
+ self.cultural_data = []
56
+ if os.path.exists(client_file):
57
+ with open(client_file, 'r', encoding='utf-8') as f:
58
+ for line in f:
59
+ self.cultural_data.append(json.loads(line))
60
+ print(f">>> [MIXER] Loaded {len(self.cultural_data)} client samples.")
61
+ else:
62
+ print(f"⚠️ WARNING: {client_file} not found. Running on Pile only.")
63
+
64
+ def __iter__(self):
65
+ pile_iterator = iter(self.pile_stream)
66
+ while True:
67
+ # Вероятностный выбор источника данных
68
+ if random.random() < self.mix_ratio and self.cultural_data:
69
+ sample = random.choice(self.cultural_data)
70
+ text = f"Question: {sample['question']}\nAnswer: {sample['answer']}"
71
+ else:
72
+ try:
73
+ sample = next(pile_iterator)
74
+ text = sample['text']
75
+ except StopIteration:
76
+ pile_iterator = iter(self.pile_stream)
77
+ continue
78
+
79
+ tokens = self.tokenizer(
80
+ text, truncation=True, max_length=BLOCK_SIZE, padding="max_length", return_tensors="pt"
81
+ )
82
+ yield {
83
+ "input_ids": tokens["input_ids"].squeeze(0),
84
+ "labels": tokens["input_ids"].squeeze(0)
85
+ }
86
+
87
+ # --- ПРОЦЕСС ОБУЧЕНИЯ ---
88
+ def train_236b():
89
+ # Инициализация акселератора (распределение весов 236B по GPU)
90
+ accelerator = Accelerator(gradient_accumulation_steps=GRAD_ACCUM_STEPS)
91
+ device = accelerator.device
92
+
93
+ if accelerator.is_main_process and not os.path.exists(CHECKPOINT_DIR):
94
+ os.makedirs(CHECKPOINT_DIR)
95
+
96
+ # 1. Токенайзер
97
+ tokenizer = AutoTokenizer.from_pretrained("meta-llama/Meta-Llama-3-8B")
98
+ if tokenizer.pad_token is None:
99
+ tokenizer.pad_token = tokenizer.eos_token
100
+
101
+ # 2. Модель 236B (192 слоя)
102
+ config = JiRackTernaryConfig()
103
+ model = JiRackTernary236B(config)
104
+
105
+ # КРИТИЧЕСКИ: Включаем градиентный чекпоинтинг для экономии VRAM
106
+ model.gradient_checkpointing_enable()
107
+
108
+ # 3. Подготовка данных
109
+ dataset = CMSDataMixer236B(tokenizer, CULTURAL_DATA_FILE, GENERAL_DATA_LINK, mix_ratio=MIX_RATIO)
110
+ loader = DataLoader(dataset, batch_size=BATCH_SIZE)
111
+
112
+ # 4. Оптимизатор
113
+ optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=0.01)
114
+
115
+ # Подготовка через accelerator
116
+ model, optimizer, loader = accelerator.prepare(model, optimizer, loader)
117
+
118
+ print(f"\n--- [CMS MANHATTAN] 236B MIXED ENGINE ONLINE ---")
119
+ print(f"Model Depth: 192 Layers | Width: 10240 | Mix: {int(MIX_RATIO*100)}% Client")
120
+
121
+ model.train()
122
+ for step, batch in enumerate(loader):
123
+ with accelerator.accumulate(model):
124
+ outputs = model(**batch)
125
+ loss = outputs.loss
126
+ accelerator.backward(loss)
127
+
128
+ # Защита от взрыва градиентов
129
+ if accelerator.sync_gradients:
130
+ accelerator.clip_grad_norm_(model.parameters(), 1.0)
131
+
132
+ optimizer.step()
133
+ optimizer.zero_grad()
134
+
135
+ if step % 20 == 0 and accelerator.is_main_process:
136
+ print(f"Step {step} | Loss: {loss.item():.4f} | VRAM: {torch.cuda.memory_allocated()/1e9:.1f}GB")
137
+
138
+ # Сохранение состояния
139
+ if step > 0 and step % 500 == 0 and accelerator.is_main_process:
140
+ save_path = os.path.join(CHECKPOINT_DIR, f"step_{step}")
141
+ accelerator.save_state(save_path)
142
+ print(f">>> [CMS] 236B Checkpoint saved: {save_path}")
143
+ torch.cuda.empty_cache()
144
+
145
+ if __name__ == "__main__":
146
+ # Оптимизация аллокатора CUDA
147
+ os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
148
+ try:
149
+ train_236b()
150
+ except KeyboardInterrupt:
151
+ print("\n[!] Остановка. Прогресс сохранен.")
152
+ except Exception as e:
153
+ print(f"FATAL ERROR: {e}")
154
+ sys.exit(1)