4141#include < dlib/data_io.h>
4242#include < dlib/cmd_line_parser.h>
4343#include < dlib/tokenizer/bpe_tokenizer.h>
44+ #include < dlib/misc_api.h>
4445
4546// Include internal dataset
4647#include " slm_data.h"
@@ -115,39 +116,6 @@ namespace dlib
115116 };
116117}
117118
118- // Signal handling for clean termination
119- namespace {
120- std::atomic<bool > g_terminate_flag (false );
121-
122- #ifdef _WIN32
123- // Windows-specific handler
124- BOOL WINAPI console_ctrl_handler (DWORD ctrl_type) {
125- if (ctrl_type == CTRL_C_EVENT ) {
126- g_terminate_flag.store (true );
127- cout << " \n Ctrl+C detected, cleaning up and closing the program..." << endl;
128- return TRUE ;
129- }
130- return FALSE ;
131- }
132-
133- void setup_interrupt_handler () {
134- SetConsoleCtrlHandler (console_ctrl_handler, TRUE );
135- }
136- #else
137- // Unix/Linux-specific handler
138- void signal_handler (int signal) {
139- if (signal == SIGINT ) {
140- g_terminate_flag.store (true );
141- cout << " \n Ctrl+C detected, cleaning up and closing the program..." << endl;
142- }
143- }
144-
145- void setup_interrupt_handler () {
146- std::signal (SIGINT , signal_handler);
147- }
148- #endif
149- }
150-
151119// Utility functions
152120std::string generate_tokens_filename (size_t max_bytes)
153121{
@@ -252,7 +220,7 @@ int main(int argc, char** argv)
252220 try
253221 {
254222 // Setup interrupt handling for clean termination
255- setup_interrupt_handler ();
223+ signal_handler::setup ();
256224
257225 command_line_parser parser;
258226 parser.add_option (" train" , " Train a transformer model on internal dataset" );
@@ -450,7 +418,6 @@ int main(int argc, char** argv)
450418 using net_type = my_transformer::network_type<true >;
451419 net_type net;
452420 const int pad_token = tokenizer.get_special_token_id (" <pad>" );
453- layer<0 >(net).loss_details ().set_ignore_index (pad_token);
454421 cout << my_transformer::model_info::describe () << endl;
455422
456423 // Tokenizer stored with model for simplified inference
@@ -475,7 +442,8 @@ int main(int argc, char** argv)
475442 auto epoch_start = std::chrono::high_resolution_clock::now ();
476443
477444 // Training loop
478- while (trainer.get_learning_rate () >= 1e-6 && epoch < max_epochs && !g_terminate_flag.load ())
445+ while (trainer.get_learning_rate () >= 1e-6 && epoch < max_epochs
446+ && !signal_handler::is_triggered ())
479447 {
480448 total_loss = 0.0 ;
481449 batches_seen = samples_seen = 0 ;
@@ -484,7 +452,7 @@ int main(int argc, char** argv)
484452 // Shuffle the dataset
485453 shuffle_training_dataset (samples, labels);
486454
487- for (size_t i = 0 ; i < samples.size () && !g_terminate_flag. load (); i += batch_size)
455+ for (size_t i = 0 ; i < samples.size () && !signal_handler::is_triggered (); i += batch_size)
488456 {
489457 size_t batch_end = std::min (i + batch_size, samples.size ());
490458 std::vector<matrix<int , 0 , 1 >> batch_samples (
@@ -529,7 +497,7 @@ int main(int argc, char** argv)
529497
530498 // Evaluate on training set
531499 {
532- if (!g_terminate_flag. load ()) {
500+ if (!signal_handler::is_triggered ()) {
533501 cout << " Evaluating model accuracy...\n " ;
534502 my_transformer::network_type<false > g_infer;
535503 deserialize (model_file) >> g_infer >> tokenizer;
@@ -665,7 +633,8 @@ int main(int argc, char** argv)
665633
666634 // Generate until target size is reached
667635 int end_of_text = tokenizer.get_special_token_id (" </text>" ), next_token = 0 ;
668- while (total_bytes < target_size && next_token != end_of_text && !g_terminate_flag.load ()) {
636+ while (total_bytes < target_size && next_token != end_of_text
637+ && !signal_handler::is_triggered ()) {
669638 // Predict next token
670639 long pad_len = count_leading_padding (input_seq, pad_token);
671640 tril_padding_context::set_uniform (pad_len, 1 );
0 commit comments