@@ -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+
198226int 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 << " \n Fine-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