Simplify logic since they all share similar args

This commit is contained in:
Aleksa Gordic 2024-06-19 20:39:17 +02:00
parent 9d141ddab9
commit 7416ee425f
2 changed files with 42 additions and 89 deletions

View file

@ -41,7 +41,7 @@ LRSchedulerType get_lr_scheduler_type_from_name(const char* name) {
}
//
// Learning rate scheduler structs
// Learning rate scheduler structs and init
//
typedef struct {
@ -49,33 +49,39 @@ typedef struct {
int warmup_iterations;
int train_num_batches;
float final_learning_rate_frac;
} CosineLearningRateScheduler;
} LearningRateScheduler;
// Linear with warmup learning rate scheduler
typedef struct {
float learning_rate;
int warmup_iterations;
int train_num_batches;
float final_learning_rate_frac;
} LinearLearningRateScheduler;
typedef struct {
float min_lr;
float max_lr;
int step_size;
} CyclicTriangularLearningRateScheduler;
// Constant learning rate scheduler
typedef struct {
float learning_rate;
} ConstantLearningRateScheduler;
void lr_scheduler_init(LearningRateScheduler *scheduler, float learning_rate, int warmup_iterations, int train_num_batches, float final_learning_rate_frac) {
scheduler->learning_rate = learning_rate;
scheduler->warmup_iterations = warmup_iterations;
scheduler->train_num_batches = train_num_batches;
scheduler->final_learning_rate_frac = final_learning_rate_frac;
}
//
// Learning rate scheduler functions
//
// switch to the appropriate learning rate scheduler
float get_learning_rate(LRSchedulerType lr_scheduler_type, LearningRateScheduler *scheduler, int step) {
float step_learning_rate;
if (lr_scheduler_type == LR_SCHEDULER_COSINE) {
step_learning_rate = get_learning_rate_cosine(scheduler, step);
} else if (lr_scheduler_type == LR_SCHEDULER_LINEAR) {
step_learning_rate = get_learning_rate_linear(scheduler, step);
} else if (lr_scheduler_type == LR_SCHEDULER_TRIANGULAR) {
step_learning_rate = get_learning_rate_triangular(scheduler, step);
} else if (lr_scheduler_type == LR_SCHEDULER_CONSTANT) {
step_learning_rate = get_learning_rate_constant(scheduler, step);
} else {
printf("Unknown learning rate scheduler type\n");
exit(EXIT_FAILURE);
}
return step_learning_rate;
}
// cosine learning rate schedule: warmup linearly to max LR, then cosine decay to LR * final_learning_rate_frac
float get_learning_rate_cosine(CosineLearningRateScheduler *scheduler, int step) {
float get_learning_rate_cosine(LearningRateScheduler *scheduler, int step) {
float lr = scheduler->learning_rate;
if (step < scheduler->warmup_iterations) {
lr = scheduler->learning_rate * ((float)(step + 1)) / scheduler->warmup_iterations;
@ -91,7 +97,7 @@ float get_learning_rate_cosine(CosineLearningRateScheduler *scheduler, int step)
}
// linear warmup learning rate schedule: warmup linearly to max LR, then decay linearly to LR * final_learning_rate_frac
float get_learning_rate_linear(LinearLearningRateScheduler *scheduler, int step) {
float get_learning_rate_linear(LearningRateScheduler *scheduler, int step) {
float lr = scheduler->learning_rate;
if (step < scheduler->warmup_iterations) {
lr = scheduler->learning_rate * ((float)(step + 1)) / scheduler->warmup_iterations;
@ -105,44 +111,21 @@ float get_learning_rate_linear(LinearLearningRateScheduler *scheduler, int step)
}
// cyclic triangular learning rate schedule: linearly increase LR from min LR to max LR, then linearly decrease LR to min LR (repeat)
float get_learning_rate_triangular(CyclicTriangularLearningRateScheduler *scheduler, int step) {
int cycle_index = 1 + step / (2 * scheduler->step_size); // tells us which cycle we are in, starting at 1
float x = fabsf((float)step / scheduler->step_size - 2 * cycle_index + 1); // goes from 0 to 1 to 0
float lr = scheduler->min_lr + (scheduler->max_lr - scheduler->min_lr) * fmaxf(0, (1 - x));
// currently hardcoded to support only a single cycle
float get_learning_rate_triangular(LearningRateScheduler *scheduler, int step) {
int step_size = scheduler->train_num_batches / 2; // number of steps in half a cycle
float min_lr = scheduler->learning_rate * scheduler->final_learning_rate_frac;
float max_lr = scheduler->learning_rate;
int cycle_index = 1 + step / (2 * step_size); // tells us which cycle we are in, starting at 1
float x = fabsf((float)step / step_size - 2 * cycle_index + 1); // goes from 0 to 1 to 0
float lr = min_lr + (max_lr - min_lr) * fmaxf(0, (1 - x));
return lr;
}
// constant learning rate schedule
float get_learning_rate_constant(ConstantLearningRateScheduler *scheduler, int step) {
float get_learning_rate_constant(LearningRateScheduler *scheduler, int step) {
return scheduler->learning_rate;
}
//
// Init functions
//
void lr_scheduler_init_cosine(CosineLearningRateScheduler *scheduler, float learning_rate, int warmup_iterations, int train_num_batches, float final_learning_rate_frac) {
scheduler->learning_rate = learning_rate;
scheduler->warmup_iterations = warmup_iterations;
scheduler->train_num_batches = train_num_batches;
scheduler->final_learning_rate_frac = final_learning_rate_frac;
}
void lr_scheduler_init_linear(LinearLearningRateScheduler *scheduler, float learning_rate, int warmup_iterations, int train_num_batches, float final_learning_rate_frac) {
scheduler->learning_rate = learning_rate;
scheduler->warmup_iterations = warmup_iterations;
scheduler->train_num_batches = train_num_batches;
scheduler->final_learning_rate_frac = final_learning_rate_frac;
}
void lr_scheduler_init_triangular(CyclicTriangularLearningRateScheduler *scheduler, float min_lr, float max_lr, int step_size) {
scheduler->min_lr = min_lr;
scheduler->max_lr = max_lr;
scheduler->step_size = step_size;
}
void lr_scheduler_init_constant(ConstantLearningRateScheduler *scheduler, float learning_rate) {
scheduler->learning_rate = learning_rate;
}
#endif // SCHEDULERS_H

View file

@ -1385,7 +1385,7 @@ int main(int argc, char *argv[]) {
else if (argv[i][1] == 'z') { zero_stage = atoi(argv[i+1]); }
else if (argv[i][1] == 'r') { recompute = atoi(argv[i+1]); }
else if (argv[i][1] == 'h') { hellaswag_eval = atoi(argv[i+1]); }
else if (argv[i][1] == 'k') { lr_scheduler_type = get_lr_scheduler_type_from_name(argv[i + 1]); }
else if (argv[i][1] == 'k') { lr_scheduler_type = get_lr_scheduler_type_from_name(argv[i+1]); }
else { error_usage(); }
}
// should do a bit more error checking here
@ -1553,26 +1553,8 @@ int main(int argc, char *argv[]) {
tokenizer_init(&tokenizer, "gpt2_tokenizer.bin");
// set up learning rate scheduler
CosineLearningRateScheduler lr_scheduler_cosine;
LinearLearningRateScheduler lr_scheduler_linear;
CyclicTriangularLearningRateScheduler lr_scheduler_triangular;
ConstantLearningRateScheduler lr_scheduler_constant;
if (lr_scheduler_type == LR_SCHEDULER_COSINE) {
lr_scheduler_init_cosine(&lr_scheduler_cosine, learning_rate, warmup_iterations, train_num_batches, final_learning_rate_frac);
} else if (lr_scheduler_type == LR_SCHEDULER_LINEAR) {
lr_scheduler_init_linear(&lr_scheduler_linear, learning_rate, warmup_iterations, train_num_batches, final_learning_rate_frac);
} else if (lr_scheduler_type == LR_SCHEDULER_TRIANGULAR) {
// Hardcode some reasonable defaults for now
float min_lr = learning_rate / 10.0f;
float max_lr = learning_rate;
int step_size = train_num_batches / 4;
lr_scheduler_init_triangular(&lr_scheduler_triangular, min_lr, max_lr, step_size);
} else if (lr_scheduler_type == LR_SCHEDULER_CONSTANT) {
lr_scheduler_init_constant(&lr_scheduler_constant, learning_rate);
} else {
printf("Unknown learning rate scheduler type\n");
exit(EXIT_FAILURE);
}
LearningRateScheduler lr_scheduler;
lr_scheduler_init(&lr_scheduler, learning_rate, warmup_iterations, train_num_batches, final_learning_rate_frac);
// some memory for generating samples from the model
int* gen_tokens = (int*)mallocCheck(B * T * sizeof(int));
@ -1733,19 +1715,7 @@ int main(int argc, char *argv[]) {
// average the loss and the gradients between all processes
gpt2_multi_gpu_loss_reduce(&model, &multi_gpu_config);
// fetch the next learning rate
float step_learning_rate;
if (lr_scheduler_type == LR_SCHEDULER_COSINE) {
step_learning_rate = get_learning_rate_cosine(&lr_scheduler_cosine, step);
} else if (lr_scheduler_type == LR_SCHEDULER_LINEAR) {
step_learning_rate = get_learning_rate_linear(&lr_scheduler_linear, step);
} else if (lr_scheduler_type == LR_SCHEDULER_TRIANGULAR) {
step_learning_rate = get_learning_rate_triangular(&lr_scheduler_triangular, step);
} else if (lr_scheduler_type == LR_SCHEDULER_CONSTANT) {
step_learning_rate = get_learning_rate_constant(&lr_scheduler_constant, step);
} else {
printf("Unknown learning rate scheduler type\n");
exit(EXIT_FAILURE);
}
float step_learning_rate = get_learning_rate(lr_scheduler_type, &lr_scheduler, step);
// update the model parameters
float grad_norm = gpt2_update(&model, step_learning_rate, 0.9f, 0.95f, 1e-8f, weight_decay, 1.0f, step+1, &multi_gpu_config);
// zero out the gradients for the next iteration