mirror of
https://github.com/karpathy/llm.c.git
synced 2026-07-28 20:35:09 -04:00
Merge pull request #636 from karpathy/feature/rolling_checkpoints
rolling checkpoints
This commit is contained in:
commit
16b5bd53fb
1 changed files with 60 additions and 22 deletions
|
|
@ -65,6 +65,9 @@ GPT-2 Transformer Neural Net training loop. See README.md for usage.
|
|||
// ----------- Multi-GPU support -----------
|
||||
#include "llmc/zero.cuh"
|
||||
|
||||
// ----------------------------------------------------------------------------
|
||||
// global vars for I/O
|
||||
char filename_buffer[512];
|
||||
|
||||
// ----------------------------------------------------------------------------
|
||||
// global vars containing information about the GPU this process is running on
|
||||
|
|
@ -1268,6 +1271,42 @@ void load_state(int* step, GPT2* model, DataLoader* loader, const char* filename
|
|||
fcloseCheck(state_file);
|
||||
}
|
||||
|
||||
void write_checkpoint(const char* output_log_dir, int step, GPT2* model, DataLoader* train_loader, MultiGpuConfig* multi_gpu_config) {
|
||||
// a checkpoint contains: model weights, optimizer/dataloader state, and a DONE file
|
||||
printf0("Writing checkpoint at step %d\n", step);
|
||||
int rank = multi_gpu_config->process_rank;
|
||||
// only rank 0 writes the model file because it is the same across all ranks
|
||||
if (rank == 0) {
|
||||
snprintf(filename_buffer, sizeof(filename_buffer), "%s/model_%08d.bin", output_log_dir, step);
|
||||
gpt2_write_to_checkpoint(model, filename_buffer);
|
||||
}
|
||||
// all ranks write their state file
|
||||
snprintf(filename_buffer, sizeof(filename_buffer), "%s/state_%08d_%05d.bin", output_log_dir, step, rank);
|
||||
save_state(filename_buffer, step, model, train_loader);
|
||||
// DONE file is a signal that this checkpoint as a whole is complete
|
||||
multi_gpu_barrier(multi_gpu_config);
|
||||
if (rank == 0) {
|
||||
snprintf(filename_buffer, sizeof(filename_buffer), "%s/DONE_%08d", output_log_dir, step);
|
||||
FILE* done_file = fopenCheck(filename_buffer, "w");
|
||||
fcloseCheck(done_file);
|
||||
}
|
||||
}
|
||||
|
||||
void delete_checkpoint(const char* output_log_dir, int step, MultiGpuConfig* multi_gpu_config) {
|
||||
// mirrors write_checkpoint function, cleans up checkpoint from disk
|
||||
printf0("Deleting checkpoint at step %d\n", step);
|
||||
int rank = multi_gpu_config->process_rank;
|
||||
if (rank == 0) {
|
||||
snprintf(filename_buffer, sizeof(filename_buffer), "%s/model_%08d.bin", output_log_dir, step);
|
||||
remove(filename_buffer);
|
||||
}
|
||||
snprintf(filename_buffer, sizeof(filename_buffer), "%s/state_%08d_%05d.bin", output_log_dir, step, rank);
|
||||
remove(filename_buffer);
|
||||
if (rank == 0) {
|
||||
snprintf(filename_buffer, sizeof(filename_buffer), "%s/DONE_%08d", output_log_dir, step);
|
||||
remove(filename_buffer);
|
||||
}
|
||||
}
|
||||
|
||||
#ifndef TESTING
|
||||
// if we are TESTING (see test_gpt2.cu), we'll skip everything below this point
|
||||
|
|
@ -1279,7 +1318,7 @@ void load_state(int* step, GPT2* model, DataLoader* loader, const char* filename
|
|||
|
||||
// ----------------------------------------------------------------------------
|
||||
// CLI, poor man's argparse
|
||||
// unclaimed flags lol: p
|
||||
// (all single letters have been claimed now)
|
||||
|
||||
void error_usage() {
|
||||
fprintf(stderr, "Usage: ./train_gpt2cu [options]\n");
|
||||
|
|
@ -1290,6 +1329,8 @@ void error_usage() {
|
|||
fprintf(stderr, " -e <string> input from model at this filename (default = gpt2_124M_bf16.bin)\n");
|
||||
fprintf(stderr, " -o <string> output log dir (default = NULL, no logging)\n");
|
||||
fprintf(stderr, " -n <int> write optimization checkpoints every how many steps? (default 0, don't)\n");
|
||||
fprintf(stderr, " -nk <int> max number of checkpoints to keep in the directory, removing old ones (0 = disable, default)\n");
|
||||
fprintf(stderr, " -nm <int> every how many step checkpoints are considered major? major checkpoints never get deleted.\n");
|
||||
fprintf(stderr, " -y <int> resume optimization found inside output log dir? (0=restart/overwrite, 1=resume/append)\n");
|
||||
// token layout for each step of the optimization
|
||||
fprintf(stderr, " -b <int> (per-GPU, micro) batch size B (default = 4)\n");
|
||||
|
|
@ -1336,7 +1377,9 @@ int main(int argc, char *argv[]) {
|
|||
const char* load_filename = "gpt2_124M_bf16.bin"; // bf16 weights of the model
|
||||
const char* lr_scheduler_type = "cosine";
|
||||
const char* output_log_dir = NULL;
|
||||
int checkpoint_every = 0; // write optimization checkpoints every how many steps?
|
||||
int checkpoint_every = 0; // write checkpoints every how many steps?
|
||||
int checkpoints_keep = 0; // how long checkpoint history do we keep? (in units of checkpoints)
|
||||
int major_checkpoint_every = 0; // major checkpoints never get deleted when maintaining history
|
||||
int resume = 0; // resume the optimization, if one is found inside output_log_dir?
|
||||
int B = 4; // batch size
|
||||
int T = 1024; // sequence length max
|
||||
|
|
@ -1372,7 +1415,7 @@ int main(int argc, char *argv[]) {
|
|||
else if (argv[i][1] == 'j') { val_data_pattern = argv[i+1]; }
|
||||
else if (argv[i][1] == 'e') { load_filename = argv[i+1]; }
|
||||
else if (argv[i][1] == 'o') { output_log_dir = argv[i+1]; }
|
||||
else if (argv[i][1] == 'n') { checkpoint_every = atoi(argv[i+1]); }
|
||||
else if (argv[i][1] == 'n' && argv[i][2] == '\0') { checkpoint_every = atoi(argv[i+1]); }
|
||||
else if (argv[i][1] == 'y') { resume = atoi(argv[i+1]); }
|
||||
else if (argv[i][1] == 'b') { B = atoi(argv[i+1]); } // Per-GPU (micro) batch size
|
||||
else if (argv[i][1] == 't') { T = atoi(argv[i+1]); }
|
||||
|
|
@ -1399,6 +1442,8 @@ int main(int argc, char *argv[]) {
|
|||
else if (argv[i][1] == 'p' && argv[i][2] == 'n') { num_processes = atoi(argv[i+1]); }
|
||||
else if (argv[i][1] == 'p' && argv[i][2] == 'r') { process_rank = atoi(argv[i+1]); }
|
||||
else if (argv[i][1] == 'p' && argv[i][2] == 'g') { gpus_per_node = atoi(argv[i+1]); }
|
||||
else if (argv[i][1] == 'n' && argv[i][2] == 'k') { checkpoints_keep = atoi(argv[i+1]); }
|
||||
else if (argv[i][1] == 'n' && argv[i][2] == 'm') { major_checkpoint_every = atoi(argv[i+1]); }
|
||||
else { error_usage(); }
|
||||
}
|
||||
multi_gpu_config = multi_gpu_config_init(num_processes, process_rank, gpus_per_node, server_ip, fs_path, nccl_init_method);
|
||||
|
|
@ -1453,7 +1498,6 @@ int main(int argc, char *argv[]) {
|
|||
printf0("+-----------------------+----------------------------------------------------+\n");
|
||||
|
||||
// figure out if we are going to be resuming the optimization
|
||||
char filename_buffer[512];
|
||||
int resuming = 0;
|
||||
int resume_max_step = find_max_step(output_log_dir);
|
||||
if (resume == 1) {
|
||||
|
|
@ -1462,7 +1506,7 @@ int main(int argc, char *argv[]) {
|
|||
if (resume_max_step == -1) {
|
||||
} else {
|
||||
resuming = 1;
|
||||
snprintf(filename_buffer, 512, "%s/model_%08d.bin", output_log_dir, resume_max_step);
|
||||
snprintf(filename_buffer, sizeof(filename_buffer), "%s/model_%08d.bin", output_log_dir, resume_max_step);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -1574,7 +1618,7 @@ int main(int argc, char *argv[]) {
|
|||
// if we found a checkpoint to resume from, load the optimization state
|
||||
int step = 0;
|
||||
if (resuming == 1) {
|
||||
snprintf(filename_buffer, 512, "%s/state_%08d_%05d.bin", output_log_dir, resume_max_step, multi_gpu_config.process_rank);
|
||||
snprintf(filename_buffer, sizeof(filename_buffer), "%s/state_%08d_%05d.bin", output_log_dir, resume_max_step, multi_gpu_config.process_rank);
|
||||
load_state(&step, &model, &train_loader, filename_buffer);
|
||||
}
|
||||
|
||||
|
|
@ -1675,23 +1719,17 @@ int main(int argc, char *argv[]) {
|
|||
// once in a while checkpoint the optimization state (all ranks)
|
||||
if ((checkpoint_every > 0 && output_log_dir != NULL && resuming == 0) &&
|
||||
((step > 0 && step % checkpoint_every == 0) || last_step)) {
|
||||
assert(strlen(output_log_dir) < 400); // being a bit lazy here
|
||||
// only rank 0 writes the model file because it is the same across all ranks
|
||||
if (multi_gpu_config.process_rank == 0) {
|
||||
snprintf(filename_buffer, 512, "%s/model_%08d.bin", output_log_dir, step);
|
||||
gpt2_write_to_checkpoint(&model, filename_buffer);
|
||||
// writes model .bin file, state .bin files, and DONE file for step
|
||||
write_checkpoint(output_log_dir, step, &model, &train_loader, &multi_gpu_config);
|
||||
// we only keep checkpoints_keep checkpoints on disk to save space
|
||||
// so now that we wrote a new checkpoint, delete one old one (unless it is a "major" checkpoint)
|
||||
// we only do this is checkpoint keeping is turned on (checkpoints_keep > 0)
|
||||
int step_delete = step - checkpoints_keep * checkpoint_every;
|
||||
if (checkpoints_keep > 0 && step_delete > 0 &&
|
||||
(major_checkpoint_every == 0 || step_delete % major_checkpoint_every != 0)
|
||||
) {
|
||||
delete_checkpoint(output_log_dir, step_delete, &multi_gpu_config);
|
||||
}
|
||||
// all ranks write their state file
|
||||
snprintf(filename_buffer, 512, "%s/state_%08d_%05d.bin", output_log_dir, step, multi_gpu_config.process_rank);
|
||||
save_state(filename_buffer, step, &model, &train_loader);
|
||||
// DONE file is a signal that this checkpoint as a whole is complete
|
||||
multi_gpu_barrier(&multi_gpu_config);
|
||||
if (multi_gpu_config.process_rank == 0) {
|
||||
snprintf(filename_buffer, 512, "%s/DONE_%08d", output_log_dir, step);
|
||||
FILE* done_file = fopenCheck(filename_buffer, "w");
|
||||
fclose(done_file);
|
||||
}
|
||||
multi_gpu_barrier(&multi_gpu_config);
|
||||
}
|
||||
resuming = 0;
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue