Skip to content

Commit b69d284

Browse files
committed
Add lr_mult_visitor for visit_layers_range
1 parent da591d3 commit b69d284

1 file changed

Lines changed: 95 additions & 15 deletions

File tree

examples/slm_chatbot_ex.cpp

Lines changed: 95 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -195,6 +195,34 @@ void display_random_qa_samples(size_t num_samples = 3)
195195
}
196196
}
197197

198+
// Visitor for setting learning rate multiplier on computational layers
199+
struct lr_mult_visitor
200+
{
201+
double mult;
202+
203+
lr_mult_visitor(double m) : mult(m) {}
204+
205+
template <typename layer_type>
206+
void operator()(size_t, layer_type& l) const
207+
{
208+
set_learning_rate_multiplier_impl(l, mult);
209+
}
210+
211+
private:
212+
template <typename T>
213+
static auto set_learning_rate_multiplier_impl(T& layer, double m)
214+
-> decltype(layer.layer_details().set_learning_rate_multiplier(m), void())
215+
{
216+
layer.layer_details().set_learning_rate_multiplier(m);
217+
}
218+
219+
template <typename T>
220+
static void set_learning_rate_multiplier_impl(T&, ...)
221+
{
222+
// No-op for layers without this method
223+
}
224+
};
225+
198226
int main(int argc, char** argv)
199227
{
200228
try
@@ -241,7 +269,7 @@ int main(int argc, char** argv)
241269
const std::string tokenizer_file = get_option(parser, "tokenizer-file", std::string("dlib_lm_tokenizer.vocab"));
242270

243271
// Configuration parameters
244-
const long vocab_size = 3500;
272+
const long vocab_size = 2000;
245273
const long max_seq_len = 128;
246274
using config = chatbot_config<vocab_size>;
247275
using train_net = config::network_type<true>;
@@ -347,26 +375,57 @@ int main(int argc, char** argv)
347375
// Release memory
348376
qa_tokens.clear();
349377

350-
cout << "Applying freezing strategy...\n";
351-
set_all_learning_rate_multipliers(net, 0.1);
352-
layer<1>(net).layer_details().set_learning_rate_multiplier(1.0); // linear
353-
layer<2>(net).layer_details().set_learning_rate_multiplier(0.5); // rms_norm
354-
layer<111>(net).layer_details().set_learning_rate_multiplier(0.5); // embeddings
355-
cout << "Fine-tuning learning rate strategy applied.\n\n";
378+
// Strategy: Freeze embeddings and lower transformer layers, fine-tune upper layers
379+
// - Embeddings: frozen (preserve learned token representations)
380+
// - Lower transformer blocks: frozen or very slow (preserve general language understanding)
381+
// - Upper transformer blocks: slow learning (adapt to domain)
382+
// - Classification head: normal learning (specialize for task)
383+
cout << "Applying freezing strategy for fine-tuning\n";
384+
// Step 1: freeze everything first (multiplier = 0)
385+
set_all_learning_rate_multipliers(net, 0.0);
386+
// Step 2: unfreeze classification head (layers 1-2: linear + rms_norm)
387+
layer<1>(net).layer_details().set_learning_rate_multiplier(1.0); // linear (classification)
388+
layer<2>(net).layer_details().set_learning_rate_multiplier(1.0); // rms_norm
389+
// Step 3: partially unfreeze upper transformer layers with gradual unfreezing
390+
// For a 3-layer transformer, unfreeze the last 1-2 blocks with reduced LR
391+
// Layer indices depend on architecture - adjust based on `net` output
392+
// Top transformer block: moderate learning
393+
visit_layers_range<3, 40>(net, lr_mult_visitor(0.3));
394+
// Middle transformer block: slower learning
395+
visit_layers_range<40, 75>(net, lr_mult_visitor(0.1));
356396
cout << net << endl;
357397

358398
size_t epoch = 0, steps = 0;
359399
size_t batches_count = 0, batches_seen = 0, samples_seen = 0;
360400
double total_loss = 0.0;
361401
auto epoch_start = std::chrono::high_resolution_clock::now();
362402

403+
// Setup learning rate scheduler with warmup
404+
const size_t steps_per_epoch = (samples.size() + batch_size - 1) / batch_size;
405+
const size_t total_steps = steps_per_epoch * max_epochs;
406+
const size_t warmup_steps = std::min(size_t(500), total_steps / 10); // 10% or 500 steps max
407+
408+
lr_scheduler scheduler(
409+
learning_rate, // peak_lr
410+
warmup_steps, // warmup_steps
411+
total_steps, // total_steps
412+
1e-7, // min_lr
413+
lr_decay_type::COSINE // decay_type
414+
);
415+
cout << "Learning rate schedule:\n"
416+
<< " - peak learning rate: " << learning_rate << "\n"
417+
<< " - warmup steps: " << warmup_steps << "\n"
418+
<< " - total steps: " << total_steps << "\n"
419+
<< " - decay: cosine\n\n";
420+
cout << "Starting fine-tuning with warmup...\n";
421+
363422
// Training loop
364-
cout << "Starting standard training...\n";
365-
while (trainer.get_learning_rate() >= trainer.get_min_learning_rate()
423+
while (!scheduler.is_training_complete()
366424
&& epoch < max_epochs && !g_terminate_flag.load())
367425
{
368426
total_loss = 0.0;
369-
batches_seen = 0, samples_seen = 0;
427+
batches_seen = 0;
428+
samples_seen = 0;
370429
epoch_start = std::chrono::high_resolution_clock::now();
371430

372431
// Shuffle the dataset
@@ -380,16 +439,24 @@ int main(int argc, char** argv)
380439
std::vector<unsigned long> batch_labels(
381440
labels.begin() + i, labels.begin() + batch_end);
382441

442+
// Update learning rate from scheduler
443+
double current_lr = scheduler.get_learning_rate();
444+
trainer.set_learning_rate(current_lr);
445+
383446
std::vector<long> pad_lengths(batch_samples.size());
384447
for (size_t j = 0; j < batch_samples.size(); ++j)
385448
pad_lengths[j] = count_leading_padding(batch_samples[j], pad_token);
386449
tril_padding_context::set_from_lengths(pad_lengths);
387450

451+
// Train
388452
trainer.train_one_step(batch_samples, batch_labels);
453+
454+
// Advance scheduler
455+
scheduler.step();
456+
389457
total_loss += trainer.get_average_loss();
390458
batches_seen++;
391459
samples_seen += batch_samples.size();
392-
steps += batch_samples.size();
393460

394461
// Progress reporting
395462
if (batches_count++ % 100 == 0) {
@@ -398,21 +465,34 @@ int main(int argc, char** argv)
398465
std::chrono::high_resolution_clock::now() - epoch_start).count();
399466
double samples_per_sec = samples_seen / (elapsed > 0 ? elapsed : 1);
400467

468+
std::ios_base::fmtflags old_flags = cout.flags();
469+
std::streamsize old_precision = cout.precision();
470+
401471
cout << "epoch#: " << (epoch + 1) << "/" << max_epochs
402-
<< " (ksteps: " << (steps / 1000) << ")"
403-
<< " \t loss: " << avg_loss
404-
<< " \t patience: " << trainer.get_steps_without_progress()
472+
<< " \t loss: " << std::fixed << std::setprecision(3) << avg_loss
473+
<< " \t lr: " << std::scientific << std::setprecision(2) << current_lr
474+
<< " \t phase: " << scheduler.get_phase_name()
475+
<< " \t progress: " << std::fixed << std::setprecision(1)
476+
<< (scheduler.get_total_progress() * 100) << "%"
405477
<< " \t speed: " << samples_per_sec << " samples/sec\n";
406478
cout.flush();
479+
480+
cout.flags(old_flags);
481+
cout.precision(old_precision);
407482
}
483+
484+
// Check if scheduler indicates training is complete
485+
if (scheduler.is_training_complete()) break;
408486
}
409487
epoch++;
410488
}
411489
tril_padding_context::clear();
412490

413491
// Save fine-tuned model
414-
set_all_learning_rate_multipliers(net, 1.0);
492+
set_all_learning_rate_multipliers(net, 1.0); // Reset multipliers before saving
415493
cout << "\nFine-tuning complete, saving specialized model...\n";
494+
cout << "Final step: " << scheduler.get_current_step()
495+
<< ", final learning rate: " << scheduler.get_learning_rate() << "\n";
416496
net.clean();
417497

418498
serialize(finetuned_model) << net << tokenizer;

0 commit comments

Comments
 (0)