LH-Tech-AI commited on
Commit
e908f87
·
verified ·
1 Parent(s): 7435246

Upload 2 files

Browse files
Files changed (2) hide show
  1. finetune.py +143 -0
  2. prepare.py +144 -0
finetune.py ADDED
@@ -0,0 +1,143 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import time
3
+ import math
4
+ import torch
5
+ from model import GPTConfig, GPT
6
+
7
+ import numpy as np
8
+
9
+ # -----------------------------------------------------------------------------
10
+ # KONFIGURATION FÜR SMALLMPRO FINETUNING
11
+ out_dir = '/media/leo/Data/checkpoints/350m_Apex_1.5_Final_NEW_More_Anti_Forgetting'
12
+ init_from = '/media/leo/Data/checkpoints/350m_fineweb' # Pfad zum Pretraining-Ordner
13
+ dataset = 'alpaca_cleaned_mixed_NEW'
14
+
15
+ # Sanfte Hyperparameter gegen Catastrophic Forgetting
16
+ batch_size = 4
17
+ gradient_accumulation_steps = 32
18
+ block_size = 1024
19
+ learning_rate = 2e-5
20
+ max_iters = 3000
21
+ weight_decay = 0.1
22
+ dropout = 0.1
23
+ warmup_iters = 100
24
+ min_lr = 3e-6
25
+ beta1, beta2 = 0.9, 0.95
26
+ device = 'cuda'
27
+ dtype = 'bfloat16'
28
+ compile = True
29
+ save_interval = 500
30
+ # -----------------------------------------------------------------------------
31
+
32
+ os.makedirs(out_dir, exist_ok=True)
33
+ torch.manual_seed(1337)
34
+ device_type = 'cuda' if 'cuda' in device else 'cpu'
35
+ ptdtype = {'float32': torch.float32, 'bfloat16': torch.bfloat16, 'float16': torch.float16}[dtype]
36
+ ctx = torch.amp.autocast(device_type=device_type, dtype=ptdtype)
37
+
38
+ # Daten-Loader (Alpaca Binärdatei)
39
+ data_dir = os.path.join('data', dataset)
40
+ train_data = np.memmap(os.path.join(data_dir, 'train.bin'), dtype=np.uint16, mode='r')
41
+ train_mask = np.memmap(os.path.join(data_dir, 'train_mask.bin'), dtype=np.uint8, mode='r')
42
+
43
+ def get_batch():
44
+ ix = torch.randint(len(train_data) - block_size, (batch_size,))
45
+ x = torch.stack([torch.from_numpy((train_data[i:i+block_size]).astype(np.int64)) for i in ix])
46
+ y = torch.stack([torch.from_numpy((train_data[i+1:i+1+block_size]).astype(np.int64)) for i in ix])
47
+ # Maske laden (entspricht y, also um 1 verschoben)
48
+ m = torch.stack([torch.from_numpy((train_mask[i+1:i+1+block_size]).astype(np.int64)) for i in ix])
49
+
50
+ # WICHTIG: Ersetze in y alle Stellen, wo m == 0 ist, durch -100
51
+ # PyTorch CrossEntropyLoss ignoriert -100 automatisch
52
+ y[m == 0] = -100
53
+
54
+ x, y = x.to(device), y.to(device)
55
+ return x, y
56
+
57
+ # Modell laden
58
+ print(f"📥 Lade Pretraining-Checkpoint aus {init_from}...")
59
+ #ckpt_files = sorted([f for f in os.listdir(init_from) if f.endswith('.pt')])
60
+ #if not ckpt_files:
61
+ # raise FileNotFoundError("Kein Checkpoint im init_from Verzeichnis gefunden!")
62
+
63
+ #ckpt_path = os.path.join(init_from, ckpt_files[-1])
64
+ ckpt_path = os.path.join(init_from, 'base_model_best_42k.pt')
65
+ checkpoint = torch.load(ckpt_path, map_location=device)
66
+ gptconf = GPTConfig(**checkpoint['model_args'])
67
+ model = GPT(gptconf)
68
+ state_dict = checkpoint['model']
69
+
70
+ # Fix für potenzielle 'orig_mod' Prefixe
71
+ unwanted_prefix = '_orig_mod.'
72
+ for k,v in list(state_dict.items()):
73
+ if k.startswith(unwanted_prefix):
74
+ state_dict[k[len(unwanted_prefix):]] = state_dict.pop(k)
75
+
76
+ model.load_state_dict(state_dict)
77
+ model.to(device)
78
+
79
+ if compile:
80
+ print("🚀 Kompiliere Modell...")
81
+ model = torch.compile(model)
82
+
83
+ optimizer = model.configure_optimizers(weight_decay, learning_rate, (beta1, beta2), device_type)
84
+ scaler = torch.cuda.amp.GradScaler(enabled=(dtype == 'float16'))
85
+
86
+ # LR Scheduler
87
+ def get_lr(it):
88
+ if it < warmup_iters: return learning_rate * it / warmup_iters
89
+ if it > max_iters: return min_lr
90
+ decay_ratio = (it - warmup_iters) / (max_iters - warmup_iters)
91
+ coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio))
92
+ return min_lr + coeff * (learning_rate - min_lr)
93
+
94
+ # Trainings-Schleife
95
+ print(f"🛠️ Starte Finetuning: Apex 1.5 lernt Chatten...")
96
+ model.train()
97
+ t0 = time.time()
98
+
99
+ for iter_num in range(max_iters + 1):
100
+ lr = get_lr(iter_num)
101
+ for param_group in optimizer.param_groups:
102
+ param_group['lr'] = lr
103
+
104
+ for micro_step in range(gradient_accumulation_steps):
105
+ X, Y = get_batch()
106
+ with ctx:
107
+ logits, loss = model(X, Y)
108
+ loss = loss / gradient_accumulation_steps
109
+ scaler.scale(loss).backward()
110
+
111
+ scaler.unscale_(optimizer)
112
+ torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
113
+ scaler.step(optimizer)
114
+ scaler.update()
115
+ optimizer.zero_grad(set_to_none=True)
116
+
117
+ if iter_num % 10 == 0:
118
+ dt = time.time() - t0
119
+ print(f"Iter {iter_num}: Loss {loss.item()*gradient_accumulation_steps:.4f}, Zeit {dt*1000:.2f}ms, LR {lr:.2e}")
120
+ t0 = time.time()
121
+
122
+ if iter_num > 0 and iter_num % save_interval == 0:
123
+ checkpoint_name = f'Apex_1.5_iter_{iter_num}.pt'
124
+ save_path = os.path.join(out_dir, checkpoint_name)
125
+ print(f"💾 Speichere Zwischen-Checkpoint: {checkpoint_name}")
126
+ raw_model = model._orig_mod if compile else model
127
+ checkpoint_data = {
128
+ 'model': raw_model.state_dict(),
129
+ 'model_args': checkpoint['model_args'],
130
+ 'iter_num': iter_num,
131
+ 'lr': lr,
132
+ }
133
+ torch.save(checkpoint_data, save_path)
134
+
135
+ # Finales Speichern
136
+ print(f"💾 Finetuning beendet. Speichere Apex 1.5...")
137
+ final_checkpoint = {
138
+ 'model': model.state_dict() if not compile else model._orig_mod.state_dict(),
139
+ 'model_args': checkpoint['model_args'],
140
+ 'config': checkpoint.get('config', {}),
141
+ }
142
+ torch.save(final_checkpoint, os.path.join(out_dir, 'Apex_1.5_Final.pt'))
143
+ print("✅ Apex 1.5 wurde erfolgreich gespeichert!")
prepare.py ADDED
@@ -0,0 +1,144 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import numpy as np
3
+ import tiktoken
4
+ from datasets import load_dataset
5
+ from tqdm import tqdm
6
+
7
+ # --- KONFIGURATION ---
8
+ OUTPUT_DIR = "data/alpaca_cleaned_mixed_NEW"
9
+ # Wichtig: Das muss exakt der Tokenizer sein, den dein Modell verwendet (meist GPT-2)
10
+ TOKENIZER_NAME = "gpt2"
11
+ SEED = 1337
12
+
13
+ # Balance: Wie viel Alpaca vs. FineWeb?
14
+ # Zu viel FineWeb = Modell antwortet nicht im Chat-Stil
15
+ # Zu wenig FineWeb = Modell wird dumm (vergisst Weltwissen)
16
+ FINEWEB_SAMPLES = 520000
17
+
18
+ enc = tiktoken.get_encoding(TOKENIZER_NAME)
19
+ EOS_TOKEN = "<|endoftext|>" # End of Sequence Token
20
+
21
+ def format_prompt_with_mask(instruction, input_text, output):
22
+ """
23
+ Formatiert den Prompt und erstellt die Loss-Maske.
24
+ Format:
25
+ Instruction: ...
26
+ Input: ... (optional)
27
+ Response: ... <|endoftext|>
28
+ """
29
+ # 1. Den "Prompt"-Teil bauen (Frage) -> Wird NICHT trainiert (Maske 0)
30
+ if input_text and input_text.strip():
31
+ prompt_text = f"Instruction:\n{instruction}\n\nInput:\n{input_text}\n\nResponse:\n"
32
+ else:
33
+ prompt_text = f"Instruction:\n{instruction}\n\nResponse:\n"
34
+
35
+ # 2. Den "Completion"-Teil bauen (Antwort) -> Wird TRAINIERT (Maske 1)
36
+ completion_text = f"{output}{EOS_TOKEN}"
37
+
38
+ # 3. Tokenisieren
39
+ # encode_plain verhindert, dass special tokens im normalen Text interpretiert werden
40
+ prompt_ids = enc.encode(prompt_text, allowed_special={'<|endoftext|>'})
41
+ completion_ids = enc.encode(completion_text, allowed_special={'<|endoftext|>'})
42
+
43
+ # 4. Zusammenfügen
44
+ full_ids = prompt_ids + completion_ids
45
+
46
+ # 5. Maske erstellen
47
+ # 0 = Ignorieren (Loss wird hier nicht berechnet)
48
+ # 1 = Trainieren (Modell soll lernen, das vorherzusagen)
49
+ mask = [0] * len(prompt_ids) + [1] * len(completion_ids)
50
+
51
+ return full_ids, mask
52
+
53
+ def main():
54
+ np.random.seed(SEED)
55
+ print(f"🚀 Starte Prepare-Script für SmaLLMPro (350M SFT)...")
56
+ print(f"📚 Tokenizer: {TOKENIZER_NAME}")
57
+
58
+ os.makedirs(OUTPUT_DIR, exist_ok=True)
59
+
60
+ # --- 1. DATENSÄTZE LADEN ---
61
+ print("📥 Lade 'yahma/alpaca-cleaned' (Chat-Instruktionen)...")
62
+ alpaca = load_dataset("yahma/alpaca-cleaned", split='train')
63
+
64
+ print(f"📥 Lade 'HuggingFaceFW/fineweb-edu' (Sample-10BT) für {FINEWEB_SAMPLES} Samples...")
65
+ fineweb = load_dataset("HuggingFaceFW/fineweb-edu", name="sample-10BT", split='train', streaming=True)
66
+
67
+ all_tokens = []
68
+ all_masks = []
69
+
70
+ # --- 2. ALPACA VERARBEITEN (Masking aktiv) ---
71
+ print("⚙️ Verarbeite Alpaca...")
72
+ for ex in tqdm(alpaca, desc="Alpaca"):
73
+ ids, mask = format_prompt_with_mask(ex['instruction'], ex['input'], ex['output'])
74
+ all_tokens.extend(ids)
75
+ all_masks.extend(mask)
76
+
77
+ alpaca_len = len(all_tokens)
78
+ print(f" -> Alpaca Tokens: {alpaca_len:,}")
79
+
80
+ # --- 3. FINEWEB VERARBEITEN (Wissenserhalt) ---
81
+ # Hier setzen wir die Maske auf 1 für den GANZEN Text.
82
+ # Warum? Das Modell soll das Weltwissen aktiv auffrischen, nicht nur als Prompt sehen.
83
+ print("⚙️ Verarbeite FineWeb (Anti-Forgetting)...")
84
+ fw_iter = iter(fineweb)
85
+ fw_count = 0
86
+ fw_tokens_count = 0
87
+
88
+ for _ in tqdm(range(FINEWEB_SAMPLES), desc="FineWeb"):
89
+ try:
90
+ ex = next(fw_iter)
91
+ text = ex['text'] + EOS_TOKEN
92
+ ids = enc.encode(text, allowed_special={EOS_TOKEN})
93
+
94
+ all_tokens.extend(ids)
95
+ # Alles lernen!
96
+ all_masks.extend([1] * len(ids))
97
+
98
+ fw_tokens_count += len(ids)
99
+ fw_count += 1
100
+ except StopIteration:
101
+ break
102
+
103
+ print(f" -> FineWeb Tokens: {fw_tokens_count:,} (aus {fw_count} Dokumenten)")
104
+
105
+ # --- 4. SPEICHERN ---
106
+ total_tokens = len(all_tokens)
107
+ print(f"\n💾 Speichere {total_tokens:,} Tokens in '{OUTPUT_DIR}'...")
108
+
109
+ # Tokens als uint16 (spart Platz, reicht für GPT-2 Vocab ~50k)
110
+ token_arr = np.array(all_tokens, dtype=np.uint16)
111
+ token_arr.tofile(os.path.join(OUTPUT_DIR, "train.bin"))
112
+
113
+ # Maske als uint8 (braucht nur 0 oder 1)
114
+ mask_arr = np.array(all_masks, dtype=np.uint8)
115
+ mask_arr.tofile(os.path.join(OUTPUT_DIR, "train_mask.bin"))
116
+
117
+ # --- 5. SANITY CHECK (Ganz wichtig!) ---
118
+ print("\n🔍 --- SANITY CHECK ---")
119
+ print("Ich dekodiere die ersten 50 Tokens des ersten Beispiels, um zu prüfen, ob alles stimmt.")
120
+ print("Grün (TRAIN) = Was das Modell lernt. Grau (IGNORE) = Was das Modell nur liest.")
121
+
122
+ check_len = 100
123
+ sample_ids = all_tokens[:check_len]
124
+ sample_mask = all_masks[:check_len]
125
+
126
+ # Wir rekonstruieren den Text und zeigen, was maskiert ist
127
+ decoded_parts = []
128
+ for t_id, m_val in zip(sample_ids, sample_mask):
129
+ token_str = enc.decode([t_id])
130
+ if m_val == 1:
131
+ decoded_parts.append(f"\033[92m{token_str}\033[0m") # Grün für Training
132
+ else:
133
+ decoded_parts.append(f"\033[90m{token_str}\033[0m") # Grau für Prompt
134
+
135
+ print("".join(decoded_parts))
136
+ print("\n(Legende: \033[90mGrau=Prompt/Ignoriert\033[0m, \033[92mGrün=Response/Gelernt\033[0m)")
137
+
138
+ if len(token_arr) != len(mask_arr):
139
+ print("\n❌ ACHTUNG: Token und Mask Array sind unterschiedlich lang! Irgendwas stimmt nicht!")
140
+ else:
141
+ print("\n✅ Alles perfekt. Arrays sind synchron. Du kannst trainieren.")
142
+
143
+ if __name__ == "__main__":
144
+ main()