CompressedGemma commited on
Commit
04f85ce
Β·
verified Β·
1 Parent(s): ddf2414

Upload 2 files

Browse files
Files changed (2) hide show
  1. gemma_inject.py +194 -0
  2. train_gemma.py +234 -0
gemma_inject.py ADDED
@@ -0,0 +1,194 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ gemma_inject.py β€” FABLE 5 β†’ HPC gate injection for Gemma 4 12B
4
+
5
+ Adapted from hpc_fable_inject.py for Gemma architecture:
6
+ - hidden_size=3840, intermediate_size=15360 (4Γ— hidden_size)
7
+ - Pauli state: [E/Ο„ | ΞΌ/Ο„ | Ξ΅/Ο„ | ΞΌ/Ο„] β€” 4 blocks filling 15360 dims
8
+ - Single model.safetensors file (~23 GB)
9
+ - tie_word_embeddings=True (lm_head == embed_tokens)
10
+ """
11
+ import argparse, gc, json, math, os, sys, time
12
+ from collections import Counter, defaultdict
13
+ from pathlib import Path
14
+
15
+ import numpy as np
16
+ import torch
17
+ from safetensors import safe_open
18
+ from safetensors.torch import save_file
19
+ from transformers import AutoTokenizer
20
+
21
+
22
+ def tokenize_fable(data_path, tokenizer):
23
+ all_ids = []
24
+ with open(data_path) as f:
25
+ for line in f:
26
+ d = json.loads(line)
27
+ text = d.get("text", "")
28
+ for part in text.split("<|im_start|>"):
29
+ if part.startswith("assistant\n"):
30
+ content = part[len("assistant\n"):].replace("<|im_end|>", "").strip()
31
+ if content:
32
+ all_ids.extend(tokenizer.encode(content, add_special_tokens=False))
33
+ return all_ids
34
+
35
+
36
+ def compute_bigram_stats(all_ids, vocab_size, topk):
37
+ cnt = Counter()
38
+ bigram = defaultdict(lambda: Counter())
39
+ for i, tid in enumerate(all_ids):
40
+ cnt[tid] += 1
41
+ if i > 0:
42
+ bigram[all_ids[i - 1]][tid] += 1
43
+ top_tokens = [t for t, _ in cnt.most_common(topk)]
44
+ token_to_idx = {t: i for i, t in enumerate(top_tokens)}
45
+ return cnt, bigram, top_tokens, token_to_idx
46
+
47
+
48
+ def recover_gate_weight(embed, cnt, bigram, top_tokens, token_to_idx, tau):
49
+ D = embed.shape[1] # 3840
50
+ V = len(top_tokens)
51
+
52
+ E_raw = np.array([embed[t] for t in top_tokens]).astype(np.float64)
53
+ mu = np.zeros_like(E_raw)
54
+ eps = np.zeros_like(E_raw)
55
+
56
+ for i, ti in enumerate(top_tokens):
57
+ total = sum(bigram[ti].values())
58
+ fi = cnt[ti] or 1
59
+ if total == 0:
60
+ continue
61
+ for tj, c in bigram[ti].items():
62
+ if tj not in token_to_idx:
63
+ continue
64
+ fj = cnt[tj] or 1
65
+ p = c / total
66
+ w = c / math.sqrt(fi * fj)
67
+ c_Z = (1.0 - w) / 2.0
68
+ mu[i] += p * embed[tj]
69
+ eps[i] += c_Z * embed[tj]
70
+
71
+ # 4-block Pauli state: [E/Ο„ | ΞΌ/Ο„ | Ξ΅/Ο„ | ΞΌ/Ο„]
72
+ z = np.hstack([E_raw / tau, mu / tau, eps / tau, mu / tau])
73
+ z = np.clip(z, -30.0, 30.0)
74
+
75
+ E_n = (E_raw - E_raw.mean(0, keepdims=True)) / (E_raw.std(0, keepdims=True) + 1e-10)
76
+ reg = 1e-3 * np.eye(D)
77
+ W = np.linalg.solve(E_n.T @ E_n + reg, E_n.T @ z)
78
+ return W.T.astype(np.float32) # (15360, 3840)
79
+
80
+
81
+ def inject_safetensors(src_path, dst_path, W_hpc_torch, alpha, layer_count):
82
+ """
83
+ Read source (mmap'd), modify gate_proj in RAM, save to new file.
84
+ Unmodified tensors stay mmap'd (zero RAM cost).
85
+ Only 48 gate_proj tensors (~5.6 GB) materialize in RAM.
86
+ """
87
+ with safe_open(str(src_path), framework='pt') as sf:
88
+ keys = list(sf.keys())
89
+ out = {}
90
+ for idx, k in enumerate(keys):
91
+ t = sf.get_tensor(k)
92
+ if 'gate_proj' in k and 'mlp' in k:
93
+ is_gate, layer_idx = _parse_gate_key(k)
94
+ if is_gate and layer_idx < layer_count:
95
+ t = t.to(torch.float32) + alpha * W_hpc_torch.to(torch.float32)
96
+ t = t.to(torch.bfloat16).contiguous()
97
+ out[k] = t
98
+
99
+ if (idx + 1) % 50 == 0:
100
+ print(f" [{idx+1}/{len(keys)}] processed", flush=True)
101
+ gc.collect()
102
+
103
+ os.makedirs(os.path.dirname(dst_path) or '.', exist_ok=True)
104
+ save_file(out, dst_path, metadata={'format': 'pt'})
105
+ print(f"Wrote {len(keys)} tensors to {dst_path}")
106
+
107
+
108
+ def _parse_gate_key(key):
109
+ parts = key.split('.')
110
+ for i, p in enumerate(parts):
111
+ if p == 'layers' and i + 1 < len(parts):
112
+ try:
113
+ return True, int(parts[i + 1])
114
+ except ValueError:
115
+ pass
116
+ return False, -1
117
+
118
+
119
+ def main():
120
+ parser = argparse.ArgumentParser(description="Gemma FABLE β†’ HPC gate injector")
121
+ parser.add_argument("--model", default="./gemma-4-12B-it", help="Gemma model directory")
122
+ parser.add_argument("--output", default="./gemma-4-12B-it-hpc", help="Output directory")
123
+ parser.add_argument("--data", default="/tmp/fable5_sft.jsonl", help="FABLE 5 JSONL path")
124
+ parser.add_argument("--alpha", type=float, default=0.3, help="Injection strength")
125
+ parser.add_argument("--tau", type=float, default=0.003, help="Temperature")
126
+ parser.add_argument("--topk", type=int, default=30000, help="Top-k tokens")
127
+ args = parser.parse_args()
128
+
129
+ t0 = time.time()
130
+ model_dir = Path(args.model)
131
+ src_safetensors = model_dir / "model.safetensors"
132
+ dst_safetensors = Path(args.output) / "model.safetensors"
133
+
134
+ # ── 1. Tokenizer & embedding ──
135
+ print("[1/6] Loading tokenizer & embedding...")
136
+ tokenizer = AutoTokenizer.from_pretrained(str(model_dir), trust_remote_code=True)
137
+ embed = None
138
+ with safe_open(str(src_safetensors), framework='pt') as sf:
139
+ for k in sf.keys():
140
+ if 'embed_tokens' in k:
141
+ embed = sf.get_tensor(k).float().numpy()
142
+ break
143
+ assert embed is not None, "embed_tokens not found"
144
+ print(f" Embedding: {embed.shape}")
145
+
146
+ # ── 2. Tokenize FABLE ──
147
+ print(f"[2/6] Tokenizing {args.data}...")
148
+ all_ids = tokenize_fable(args.data, tokenizer)
149
+ print(f" {len(all_ids)} tokens")
150
+
151
+ # ── 3. Bigram stats ──
152
+ print(f"[3/6] Bigrams (top-{args.topk})...")
153
+ cnt, bigram, top_tokens, token_to_idx = compute_bigram_stats(all_ids, tokenizer.vocab_size, args.topk)
154
+ print(f" {len(top_tokens)} tokens")
155
+
156
+ # ── 4. Recover gate weight ──
157
+ print(f"[4/6] Recovering gate weight (Ο„={args.tau})...")
158
+ W_hpc = recover_gate_weight(embed, cnt, bigram, top_tokens, token_to_idx, args.tau)
159
+ target_std = 0.012
160
+ scale = target_std / W_hpc.std()
161
+ W_hpc *= scale
162
+ print(f" W_hpc: {W_hpc.shape}, std={W_hpc.std():.6f}")
163
+
164
+ # ── 5. Get layer count ──
165
+ print(f"[5/6] Scanning source file...")
166
+ layer_count = 0
167
+ with safe_open(str(src_safetensors), framework='pt') as sf:
168
+ for k in sf.keys():
169
+ is_gate, layer = _parse_gate_key(k)
170
+ if is_gate:
171
+ layer_count = max(layer_count, layer + 1)
172
+ print(f" {layer_count} layers detected")
173
+
174
+ W_hpc_t = torch.from_numpy(W_hpc).to(torch.bfloat16)
175
+
176
+ # ── 6. Inject ──
177
+ print(f"[6/6] Injecting (Ξ±={args.alpha}) β†’ {args.output}...")
178
+ os.makedirs(args.output, exist_ok=True)
179
+ inject_safetensors(str(src_safetensors), str(dst_safetensors),
180
+ W_hpc_t, args.alpha, layer_count)
181
+
182
+ import shutil
183
+ for fn in ["config.json", "tokenizer.json", "tokenizer_config.json",
184
+ "generation_config.json", "chat_template.jinja", "processor_config.json"]:
185
+ src = model_dir / fn
186
+ if src.exists():
187
+ shutil.copy2(src, Path(args.output) / fn)
188
+
189
+ elapsed = time.time() - t0
190
+ print(f"Done in {elapsed:.0f}s. Output: {args.output}/")
191
+
192
+
193
+ if __name__ == "__main__":
194
+ main()
train_gemma.py ADDED
@@ -0,0 +1,234 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ train_gemma.py β€” QLoRA fine-tune HPC-injected Gemma 4 12B on FABLE 5
4
+
5
+ Loads Gemma-4-12B-it-hpc in 4-bit, adds LoRA to MLP projections,
6
+ fine-tunes on FABLE 5 assistant conversations.
7
+
8
+ Usage:
9
+ python3 train_gemma.py --model ./gemma-4-12B-it-hpc \\
10
+ --data /tmp/fable5_sft.jsonl \\
11
+ --output ./gemma-4-12B-it-fable \\
12
+ --epochs 1 \\
13
+ --lr 2e-4
14
+ """
15
+ import argparse, gc, json, math, os, sys, time
16
+ from functools import partial
17
+
18
+ import torch
19
+ import torch.nn as nn
20
+ from torch.utils.data import Dataset, DataLoader
21
+ from transformers import (
22
+ AutoTokenizer,
23
+ AutoModelForCausalLM,
24
+ BitsAndBytesConfig,
25
+ get_linear_schedule_with_warmup,
26
+ )
27
+ from peft import LoraConfig, get_peft_model
28
+
29
+
30
+ # ── Dataset ──────────────────────────────────────────────────────────────────
31
+
32
+ class FableDataset(Dataset):
33
+ """Tokenized FABLE 5 conversations, formatted for Gemma chat template."""
34
+
35
+ def __init__(self, data_path, tokenizer, max_length=2048):
36
+ self.tokenizer = tokenizer
37
+ self.max_length = max_length
38
+ self.samples = []
39
+
40
+ with open(data_path) as f:
41
+ for line in f:
42
+ d = json.loads(line)
43
+ text = d.get("text", "")
44
+ if not text:
45
+ continue
46
+
47
+ # Parse FABLE <|im_start|> format β†’ messages list
48
+ turns = text.split("<|im_start|>")
49
+ messages = []
50
+ for turn in turns:
51
+ turn = turn.strip()
52
+ if not turn:
53
+ continue
54
+ if turn.startswith("user\n"):
55
+ messages.append({
56
+ "role": "user",
57
+ "content": turn[len("user\n"):].replace("<|im_end|>", "").strip()
58
+ })
59
+ elif turn.startswith("assistant\n"):
60
+ messages.append({
61
+ "role": "assistant",
62
+ "content": turn[len("assistant\n"):].replace("<|im_end|>", "").strip()
63
+ })
64
+
65
+ if len(messages) <= 1:
66
+ continue
67
+
68
+ # Format with Gemma's chat template
69
+ formatted = tokenizer.apply_chat_template(
70
+ messages, tokenize=False, add_generation_prompt=False
71
+ )
72
+ tokens = tokenizer.encode(
73
+ formatted, add_special_tokens=False,
74
+ truncation=True, max_length=max_length
75
+ )
76
+ self.samples.append(tokens)
77
+
78
+ def __len__(self):
79
+ return len(self.samples)
80
+
81
+ def __getitem__(self, idx):
82
+ return torch.tensor(self.samples[idx], dtype=torch.long)
83
+
84
+
85
+ def collate_fn(batch, pad_token_id):
86
+ max_len = max(len(x) for x in batch)
87
+ padded = torch.full((len(batch), max_len), pad_token_id, dtype=torch.long)
88
+ for i, seq in enumerate(batch):
89
+ padded[i, :len(seq)] = seq
90
+ return padded
91
+
92
+
93
+ # ── Training ─────────────────────────────────────────────────────────────────
94
+
95
+ def train():
96
+ parser = argparse.ArgumentParser(description="Fine-tune HPC-injected Gemma on FABLE 5")
97
+ parser.add_argument("--model", default="./gemma-4-12B-it-hpc", help="Injected model path")
98
+ parser.add_argument("--data", default="/tmp/fable5_sft.jsonl", help="FABLE 5 JSONL path")
99
+ parser.add_argument("--output", default="./gemma-4-12B-it-fable", help="Output path")
100
+ parser.add_argument("--epochs", type=int, default=1, help="Training epochs")
101
+ parser.add_argument("--lr", type=float, default=2e-4, help="Peak learning rate")
102
+ parser.add_argument("--batch_size", type=int, default=1, help="Per-device batch size")
103
+ parser.add_argument("--grad_accum", type=int, default=8, help="Gradient accumulation steps")
104
+ parser.add_argument("--max_length", type=int, default=2048, help="Max sequence length")
105
+ parser.add_argument("--lora_r", type=int, default=16, help="LoRA rank")
106
+ parser.add_argument("--lora_alpha", type=int, default=32, help="LoRA alpha")
107
+ parser.add_argument("--lora_dropout", type=float, default=0.05, help="LoRA dropout")
108
+ parser.add_argument("--save_steps", type=int, default=200, help="Checkpoint interval (steps)")
109
+ args = parser.parse_args()
110
+
111
+ t0 = time.time()
112
+
113
+ # ── 1. Tokenizer & 4-bit model ──
114
+ print("[1/6] Loading tokenizer & 4-bit model...")
115
+ tokenizer = AutoTokenizer.from_pretrained(args.model, trust_remote_code=True, use_fast=False)
116
+ tokenizer.pad_token = tokenizer.eos_token
117
+ tokenizer.padding_side = "right"
118
+
119
+ bnb = BitsAndBytesConfig(
120
+ load_in_4bit=True,
121
+ bnb_4bit_quant_type="nf4",
122
+ bnb_4bit_use_double_quant=True,
123
+ bnb_4bit_compute_dtype=torch.bfloat16,
124
+ )
125
+
126
+ gpu_mem = torch.cuda.get_device_properties(0).total_memory // (1024**3)
127
+ max_memory = {0: f"{gpu_mem - 3}GiB", "cpu": "64GiB"}
128
+
129
+ model = AutoModelForCausalLM.from_pretrained(
130
+ args.model,
131
+ trust_remote_code=True,
132
+ quantization_config=bnb,
133
+ device_map="auto",
134
+ max_memory=max_memory,
135
+ torch_dtype=torch.bfloat16,
136
+ low_cpu_mem_usage=True,
137
+ )
138
+ model.config.use_cache = False
139
+ model.gradient_checkpointing_enable()
140
+
141
+ # ── 2. LoRA ──
142
+ print(f"[2/6] Adding LoRA (r={args.lora_r}, alpha={args.lora_alpha})...")
143
+ lora_config = LoraConfig(
144
+ r=args.lora_r,
145
+ lora_alpha=args.lora_alpha,
146
+ lora_dropout=args.lora_dropout,
147
+ bias="none",
148
+ task_type="CAUSAL_LM",
149
+ target_modules=["gate_proj", "up_proj", "down_proj"],
150
+ )
151
+ model = get_peft_model(model, lora_config)
152
+ model.print_trainable_parameters()
153
+
154
+ # ── 3. Data ──
155
+ print("[3/6] Loading FABLE 5 dataset...")
156
+ dataset = FableDataset(args.data, tokenizer, max_length=args.max_length)
157
+ loader = DataLoader(
158
+ dataset,
159
+ batch_size=args.batch_size,
160
+ shuffle=True,
161
+ collate_fn=partial(collate_fn, pad_token_id=tokenizer.pad_token_id),
162
+ num_workers=2,
163
+ pin_memory=True,
164
+ )
165
+ print(f" {len(dataset)} samples, {len(loader)} batches/epoch")
166
+
167
+ # ── 4. Optimizer & scheduler ──
168
+ print("[4/6] Setting up optimizer...")
169
+ opt = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad], lr=args.lr)
170
+ total_steps = len(loader) * args.epochs // args.grad_accum
171
+ scheduler = get_linear_schedule_with_warmup(
172
+ opt, num_warmup_steps=int(0.05 * total_steps), num_training_steps=total_steps
173
+ )
174
+
175
+ # ── 5. Training loop ──
176
+ print(f"[5/6] Training ({args.epochs} epoch(s))...")
177
+ os.makedirs(args.output, exist_ok=True)
178
+ global_step = 0
179
+ best_loss = float("inf")
180
+
181
+ for epoch in range(args.epochs):
182
+ model.train()
183
+ total_loss = 0.0
184
+ n_batches = 0
185
+ epoch_t0 = time.time()
186
+
187
+ for batch_idx, batch in enumerate(loader):
188
+ batch = batch.to(model.device)
189
+ labels = batch.clone()
190
+
191
+ loss = model(input_ids=batch, labels=labels).loss
192
+ loss = loss / args.grad_accum
193
+ loss.backward()
194
+
195
+ total_loss += loss.item() * args.grad_accum
196
+ n_batches += 1
197
+
198
+ if (batch_idx + 1) % args.grad_accum == 0 or (batch_idx + 1) == len(loader):
199
+ torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
200
+ opt.step()
201
+ scheduler.step()
202
+ opt.zero_grad()
203
+ global_step += 1
204
+
205
+ if global_step % args.save_steps == 0:
206
+ avg_loss = total_loss / n_batches
207
+ ppl = math.exp(avg_loss)
208
+ save_path = os.path.join(args.output, f"checkpoint-{global_step}")
209
+ model.save_pretrained(save_path)
210
+ tokenizer.save_pretrained(save_path)
211
+ print(f" Step {global_step}: loss={avg_loss:.4f}, ppl={ppl:.2f}")
212
+
213
+ if (batch_idx + 1) % 20 == 0:
214
+ current_loss = total_loss / n_batches
215
+ print(f" Epoch {epoch+1}, batch {batch_idx+1}/{len(loader)}: loss={current_loss:.4f}")
216
+
217
+ avg_loss = total_loss / n_batches
218
+ ppl = math.exp(avg_loss)
219
+ epoch_time = time.time() - epoch_t0
220
+ print(f" Epoch {epoch+1} done: loss={avg_loss:.4f}, ppl={ppl:.2f}, time={epoch_time:.0f}s")
221
+
222
+ if avg_loss < best_loss:
223
+ best_loss = avg_loss
224
+ model.save_pretrained(os.path.join(args.output, "best"))
225
+ tokenizer.save_pretrained(os.path.join(args.output, "best"))
226
+
227
+ model.save_pretrained(os.path.join(args.output, "final"))
228
+ tokenizer.save_pretrained(os.path.join(args.output, "final"))
229
+ print(f"\nDone in {time.time()-t0:.0f}s. Final model: {args.output}/final")
230
+ print(f"Best loss: {best_loss:.4f} (PPL={math.exp(best_loss):.2f})")
231
+
232
+
233
+ if __name__ == "__main__":
234
+ train()