Views
No views yet
[!TIP]🐧 If you're impppatient, get the trained checkpoint file that runs on 1 cpu core:wget https://huggingface.co/nisten/Biggie-SmoLlm-0.15B-Base/resolve/main/biggie_groked_int8_q8_0.ggufmake sure to install latest llama.cpp first, it's easy on linux & mac:git clone https://github.com/ggerganov/llama.cpp && cd llama.cpp && make -j
1./llama-cli -fa -b 512 -ctv q8_0 -ctk q8_0 --min-p 0.3 --top-p 0.85 --keep -1 \
2 -p "You are a NASA JPL Scientists. Human: I want to bring my cat to mars." \
3 --in-prefix "<|im_start|>Human:" --reverse-prompt "Human:" \
4 -m biggie_groked_int8_q8_0.gguf -co -cnv \
5 -c 1024 -n 700 --temp 1.5 -ngl 0 -t 1wget https://huggingface.co/nisten/Biggie-SmoLlm-0.15B-Base/resolve/main/biggie_groked_int8_q8_0.gguf 
./llama-cli -n -1 -fa -b 512 -ctv q8_0 -ctk q8_0 -fa --min-p 0.3 --top-p 0.85 --keep -1 -p "You are a NASA JPL Scientists. Human: I want to bring my cat to mars." -m biggie_groked_int8_q8_0.gguf -co -cnv --in-prefix "<|im_start|>Human:" --reverse-prompt "Human:" -c 1024 -n 512 --temp 1.5 -ngl 01git clone https://github.com/nisten/grokadamw
2cd grokadamwpython smoltrainer.pypip install torch transformers datasets tqdmpython meow.py1import torch
2import torch.nn as nn
3import logging
4from datasets import load_dataset, Dataset
5from transformers import AutoConfig, AutoTokenizer, AutoModelForCausalLM, TrainingArguments, Trainer, DataCollatorForLanguageModeling
6from torch.cuda.amp import autocast
7import warnings
8from tqdm import tqdm
9
10warnings.filterwarnings("ignore", category=FutureWarning)
11warnings.filterwarnings("ignore", category=UserWarning)
12
13logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
14logger = logging.getLogger(__name__)
15
16MODEL_NAME = "nisten/Biggie-SmoLlm-0.15B-Base"
17MAX_LENGTH = 2048
18BATCH_SIZE = 8
19LEARNING_RATE = 2e-4
20MAX_STEPS = 3000
21GRADIENT_ACCUMULATION_STEPS = 2
22NUM_WARMUP_STEPS = 30
23OUTPUT_DIR = "./capybara_finetuned_results"
24
25torch.backends.cuda.matmul.allow_tf32 = True
26torch.backends.cudnn.allow_tf32 = True
27
28class GrokAdamW(torch.optim.Optimizer):
29 def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=1e-2,
30 alpha_init=0.98, lamb=2.0, gamma=0.1, grokking_signal_fns=None,
31 grokking_signal_decay_rate=0.1, gradient_clipping=1.0):
32 defaults = dict(lr=lr, betas=betas, eps=eps, weight_decay=weight_decay,
33 alpha_init=alpha_init, lamb=lamb, gamma=gamma,
34 grokking_signal_fns=grokking_signal_fns,
35 grokking_signal_decay_rate=grokking_signal_decay_rate,
36 gradient_clipping=gradient_clipping)
37 super(GrokAdamW, self).__init__(params, defaults)
38
39 @torch.no_grad()
40 def step(self, closure=None):
41 loss = None
42 if closure is not None:
43 with torch.enable_grad():
44 loss = closure()
45
46 for group in self.param_groups:
47 grokking_signal = self._compute_grokking_signal(group)
48 for i, p in enumerate(group['params']):
49 if p.grad is None:
50 continue
51 grad = p.grad
52
53 if group['gradient_clipping'] > 0:
54 grad = torch.clamp(grad, -group['gradient_clipping'], group['gradient_clipping'])
55
56 state = self.state[p]
57
58 if len(state) == 0:
59 state['step'] = 0
60 state['exp_avg'] = torch.zeros_like(p, memory_format=torch.preserve_format)
61 state['exp_avg_sq'] = torch.zeros_like(p, memory_format=torch.preserve_format)
62 state['grok_ema'] = torch.zeros_like(p, memory_format=torch.preserve_format)
63
64 exp_avg, exp_avg_sq, grok_ema = state['exp_avg'], state['exp_avg_sq'], state['grok_ema']
65 beta1, beta2 = group['betas']
66
67 state['step'] += 1
68
69 layer_beta1 = beta1 * (1 - group['gamma'])**i
70
71 alpha = group['alpha_init'] * torch.exp(torch.tensor(-group['grokking_signal_decay_rate'] * grokking_signal))
72 grok_ema.mul_(alpha).add_(grad, alpha=1 - alpha)
73 grok_grad = grad + group['lamb'] * grok_ema
74
75 exp_avg.mul_(layer_beta1).add_(grok_grad, alpha=1 - layer_beta1)
76 exp_avg_sq.mul_(beta2).addcmul_(grok_grad, grok_grad, value=1 - beta2)
77
78 denom = exp_avg_sq.sqrt().add_(group['eps'])
79 step_size = group['lr']
80
81 if group['weight_decay'] != 0:
82 p.data.mul_(1 - group['lr'] * group['weight_decay'])
83
84 p.addcdiv_(exp_avg, denom, value=-step_size)
85
86 return loss
87
88 def _compute_grokking_signal(self, group):
89 if group['grokking_signal_fns'] is None:
90 return 0.0
91
92 signals = []
93 for fn in group['grokking_signal_fns']:
94 try:
95 signal = fn()
96 if signal is not None:
97 signals.append(signal)
98 except Exception as e:
99 logger.warning(f"Error in grokking_signal_fn: {e}. Ignoring this function.")
100
101 if not signals:
102 return 0.0
103
104 return sum(signals) / len(signals)
105
106def format_capybara_prompts(examples):
107 texts = []
108 for conversation in examples['conversation']:
109 formatted_text = ""
110 for turn in conversation:
111 if 'input' in turn:
112 formatted_text += f"Human: {turn['input']}\n\n"
113 if 'output' in turn:
114 formatted_text += f"Assistant: {turn['output']}\n\n"
115 texts.append(formatted_text.strip())
116 return {"text": texts}
117
118class CustomTrainer(Trainer):
119 def __init__(self, *args, **kwargs):
120 super().__init__(*args, **kwargs)
121 self.grokking_signal = 0.0
122
123 def compute_loss(self, model, inputs, return_outputs=False):
124 labels = inputs.pop("labels")
125 outputs = model(**inputs)
126 logits = outputs.logits
127 shift_logits = logits[..., :-1, :].contiguous()
128 shift_labels = labels[..., 1:].contiguous()
129 loss_fct = nn.CrossEntropyLoss()
130 loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1))
131 return (loss, outputs) if return_outputs else loss
132
133 def training_step(self, model, inputs):
134 model.train()
135 inputs = self._prepare_inputs(inputs)
136
137 with autocast(dtype=torch.bfloat16):
138 loss = self.compute_loss(model, inputs)
139
140 if self.args.gradient_accumulation_steps > 1:
141 loss = loss / self.args.gradient_accumulation_steps
142
143 loss.backward()
144
145 self.grokking_signal = loss.item()
146
147 return loss.detach()
148
149def grokking_signal_fn():
150 return trainer.grokking_signal
151
152def main():
153 logger.info(f"🚀 Initializing {MODEL_NAME} finetuning with GrokAdamW")
154
155 try:
156 config = AutoConfig.from_pretrained(MODEL_NAME)
157 tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
158 model = AutoModelForCausalLM.from_pretrained(MODEL_NAME, torch_dtype=torch.bfloat16)
159 except Exception as e:
160 logger.error(f"❌ Failed to load model or tokenizer: {str(e)}")
161 return
162
163 if tokenizer.pad_token is None:
164 tokenizer.pad_token = tokenizer.eos_token
165 model.config.pad_token_id = model.config.eos_token_id
166
167 logger.info("📚 Loading Capybara dataset")
168 try:
169 capybara_dataset = load_dataset("LDJnr/Capybara", split="train")
170 capybara_dataset = capybara_dataset.map(format_capybara_prompts, batched=True, remove_columns=capybara_dataset.column_names)
171 except Exception as e:
172 logger.error(f"❌ Failed to load Capybara dataset: {str(e)}")
173 return
174
175 logger.info(f"📊 Capybara dataset size: {len(capybara_dataset)}")
176
177 def tokenize_function(examples):
178 return tokenizer(examples["text"], truncation=True, padding="max_length", max_length=MAX_LENGTH)
179
180 logger.info("🔢 Tokenizing dataset")
181 tokenized_dataset = capybara_dataset.map(tokenize_function, batched=True, remove_columns=capybara_dataset.column_names)
182
183 logger.info("🏋️ Setting up the training arguments")
184 training_args = TrainingArguments(
185 output_dir=OUTPUT_DIR,
186 num_train_epochs=3,
187 per_device_train_batch_size=BATCH_SIZE,
188 gradient_accumulation_steps=GRADIENT_ACCUMULATION_STEPS,
189 learning_rate=LEARNING_RATE,
190 weight_decay=0.01,
191 bf16=True,
192 logging_steps=10,
193 save_steps=300,
194 save_total_limit=10,
195 dataloader_num_workers=4,
196 warmup_steps=NUM_WARMUP_STEPS,
197 gradient_checkpointing=True,
198 evaluation_strategy="steps",
199 eval_steps=300,
200 max_steps=MAX_STEPS,
201 fp16=False,
202 optim="adamw_hf",
203 lr_scheduler_type="cosine",
204 load_best_model_at_end=True,
205 metric_for_best_model="loss",
206 greater_is_better=False,
207 )
208
209 data_collator = DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm=False)
210
211 optimizer = GrokAdamW(
212 model.parameters(),
213 lr=LEARNING_RATE,
214 betas=(0.9, 0.999),
215 eps=1e-8,
216 weight_decay=0.01,
217 alpha_init=0.98,
218 lamb=2.0,
219 gamma=0.1,
220 grokking_signal_fns=[grokking_signal_fn],
221 grokking_signal_decay_rate=0.1,
222 gradient_clipping=1.0
223 )
224
225 logger.info("🏃♂️ Initializing Trainer with GrokAdamW")
226 global trainer
227 trainer = CustomTrainer(
228 model=model,
229 args=training_args,
230 train_dataset=tokenized_dataset,
231 eval_dataset=tokenized_dataset.select(range(min(1000, len(tokenized_dataset)))),
232 data_collator=data_collator,
233 optimizers=(optimizer, None),
234 )
235
236 logger.info("🔥 Starting the training with GrokAdamW")
237 try:
238 trainer.train()
239 except Exception as e:
240 logger.error(f"❌ Training failed: {str(e)}")
241 return
242
243 logger.info("💾 Saving the model")
244 try:
245 trainer.save_model(OUTPUT_DIR)
246 except Exception as e:
247 logger.error(f"❌ Failed to save model: {str(e)}")
248
249 logger.info("🎉 Finetuning with GrokAdamW completed!")
250
251if __name__ == "__main__":
252 main()Note: You'll need about 14GB of VRAM. If you have 8GB, change to batch size 4.
./capybara_finetuned_results