From 4f7d8d9f44c4186a102f754161abc2b5cf9ed8a9 Mon Sep 17 00:00:00 2001 From: Andrej Karpathy Date: Sun, 28 Apr 2024 16:08:27 +0000 Subject: [PATCH] allow user to make different precisions, add prints and error handling around precisions --- Makefile | 18 ++++++++++++++++-- train_gpt2.cu | 32 ++++++++++++++++++++++++++------ 2 files changed, 42 insertions(+), 8 deletions(-) diff --git a/Makefile b/Makefile index 2908e27..33f3c1f 100644 --- a/Makefile +++ b/Makefile @@ -85,6 +85,20 @@ else endif endif +# Precision settings, default to bf16 but ability to override +PRECISION ?= BF16 +VALID_PRECISIONS := FP32 FP16 BF16 +ifeq ($(filter $(PRECISION),$(VALID_PRECISIONS)),) + $(error Invalid precision $(PRECISION), valid precisions are $(VALID_PRECISIONS)) +endif +ifeq ($(PRECISION), FP32) + PFLAGS = -DENABLE_FP32 +else ifeq ($(PRECISION), FP16) + PFLAGS = -DENABLE_FP16 +else + PFLAGS = -DENABLE_BF16 +endif + # PHONY means these targets will always be executed .PHONY: all train_gpt2 test_gpt2 train_gpt2cu test_gpt2cu train_gpt2fp32cu test_gpt2fp32cu @@ -108,13 +122,13 @@ test_gpt2: test_gpt2.c $(CC) $(CFLAGS) $(INCLUDES) $(LDFLAGS) $< $(LDLIBS) -o $@ train_gpt2cu: train_gpt2.cu - $(NVCC) $(NVCC_FLAGS) $< $(NVCC_LDFLAGS) $(NVCC_INCLUDES) $(NVCC_LDLIBS) $(NVCC_LDFLAGS) -o $@ + $(NVCC) $(NVCC_FLAGS) $(PFLAGS) $< $(NVCC_LDFLAGS) $(NVCC_INCLUDES) $(NVCC_LDLIBS) $(NVCC_LDFLAGS) -o $@ train_gpt2fp32cu: train_gpt2_fp32.cu $(NVCC) $(NVCC_FLAGS) $< $(NVCC_LDFLAGS) $(NVCC_INCLUDES) $(NVCC_LDLIBS) $(NVCC_LDFLAGS) -o $@ test_gpt2cu: test_gpt2.cu - $(NVCC) $(NVCC_FLAGS) $< $(NVCC_LDFLAGS) $(NVCC_INCLUDES) $(NVCC_LDLIBS) $(NVCC_LDFLAGS) -o $@ + $(NVCC) $(NVCC_FLAGS) $(PFLAGS) $< $(NVCC_LDFLAGS) $(NVCC_INCLUDES) $(NVCC_LDLIBS) $(NVCC_LDFLAGS) -o $@ test_gpt2fp32cu: test_gpt2_fp32.cu $(NVCC) $(NVCC_FLAGS) $< $(NVCC_LDFLAGS) $(NVCC_INCLUDES) $(NVCC_LDLIBS) $(NVCC_LDFLAGS) -o $@ diff --git a/train_gpt2.cu b/train_gpt2.cu index 835efda..5a3d682 100644 --- a/train_gpt2.cu +++ b/train_gpt2.cu @@ -37,12 +37,13 @@ mpirun -np 4 ./train_gpt2cu -b 8 -v 200 -s 200 -i data/TinyStories #include #include #include +// GPU / CUDA related #include #include #include #include #include - +// Multi-GPU related #ifdef MULTI_GPU #include #include @@ -51,9 +52,11 @@ mpirun -np 4 ./train_gpt2cu -b 8 -v 200 -s 200 -i data/TinyStories // ---------------------------------------------------------------------------- // CUDA precision settings -// turn on bf16 as default, done up here for now -#define ENABLE_BF16 -// #define ENABLE_FP32 +enum PrecisionMode { + PRECISION_FP32, + PRECISION_FP16, + PRECISION_BF16 +}; // use bf16 (bfloat 16) #if defined(ENABLE_BF16) @@ -62,6 +65,8 @@ typedef float floatN; #define CUBLAS_LOWP CUDA_R_16BF #define CUBLAS_LOWP_COMPUTE CUBLAS_COMPUTE_32F const char* load_filename = "gpt2_124M_bf16.bin"; // bf16 weights +PrecisionMode PRECISION_MODE = PRECISION_BF16; +const char* precision_mode_str = "bf16"; #ifdef MULTI_GPU const ncclDataType_t ncclFloatX = ncclBfloat16; @@ -75,6 +80,8 @@ typedef float floatN; #define CUBLAS_LOWP CUDA_R_16F #define CUBLAS_LOWP_COMPUTE CUBLAS_COMPUTE_32F const char* load_filename = "gpt2_124M.bin"; // fp32 weights +PrecisionMode PRECISION_MODE = PRECISION_FP16; +const char* precision_mode_str = "fp16"; #ifdef MULTI_GPU const ncclDataType_t ncclFloatX = ncclHalf; @@ -88,6 +95,8 @@ typedef float floatN; #define CUBLAS_LOWP CUDA_R_32F #define CUBLAS_LOWP_COMPUTE cublas_compute_type // auto-select FP32 vs TF32 const char* load_filename = "gpt2_124M.bin"; // fp32 weights +PrecisionMode PRECISION_MODE = PRECISION_FP32; +const char* precision_mode_str = "fp32"; #ifdef MULTI_GPU const ncclDataType_t ncclFloatX = ncclFloat; @@ -271,6 +280,7 @@ FILE *fopen_check(const char *path, const char *mode, const char *file, int line fprintf(stderr, " Line: %d\n", line); fprintf(stderr, " Path: %s\n", path); fprintf(stderr, " Mode: %s\n", mode); + fprintf(stderr, "---> HINT: try to re-run `python train_gpt2.py`\n"); exit(EXIT_FAILURE); } return fp; @@ -1700,16 +1710,25 @@ typedef struct { void gpt2_build_from_checkpoint(GPT2 *model, const char* checkpoint_path) { + if (PRECISION_MODE == PRECISION_FP16) { + // TODO for later perhaps, would require us dynamically converting the + // model weights from fp32 to fp16 online, here in this function, or writing + // the fp16 weights directly from Python, which we only do for fp32/bf16 atm. + fprintf(stderr, "build_from_checkpoint() does not support fp16 right now.\n"); + exit(EXIT_FAILURE); + } + // read in model from a checkpoint file FILE *model_file = fopenCheck(checkpoint_path, "rb"); int model_header[256]; freadCheck(model_header, sizeof(int), 256, model_file); - if (model_header[0] != 20240326) { printf("Bad magic model file"); exit(EXIT_FAILURE); } + if (model_header[0] != 20240326) { printf("Bad magic model file\n"); exit(EXIT_FAILURE); } int version = model_header[1]; if (!(version == 1 || version == 2)) { // 1 = fp32, ordered layernorm at the end // 2 = bf16, ordered layernorm at the end - printf("Bad version in model file"); + fprintf(stderr, "Bad version in model file\n"); + fprintf(stderr, "---> HINT: try to re-run `python train_gpt2.py`\n"); exit(EXIT_FAILURE); } @@ -2410,6 +2429,7 @@ int main(int argc, char *argv[]) { cudaCheck(cudaMalloc(&cublaslt_workspace, cublaslt_workspace_size)); printf0("| device | %-50s |\n", deviceProp.name); printf0("| TF32 | %-50s |\n", enable_tf32 ? "enabled" : "disabled"); + printf0("| precision | %-50s |\n", precision_mode_str); printf0("+-----------------------+----------------------------------------------------+\n"); // build the GPT-2 model from a checkpoint