mirror of
https://github.com/karpathy/llm.c.git
synced 2026-07-28 20:35:09 -04:00
small touchups to grad clip
This commit is contained in:
parent
da9f4e21b0
commit
c2d12f725e
3 changed files with 24 additions and 26 deletions
|
|
@ -59,10 +59,15 @@ __global__ void norm_kernel1(float* out, const T* data, size_t count) {
|
|||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
template<class T>
|
||||
__global__ void norm_kernel2(float* out, const T* data, size_t count) {
|
||||
// concrete example for an A100 GPU (108 SMs, 2048 max threads each)
|
||||
// so there are 2048 * 108 = 221,184 threads total
|
||||
// say the block_size is 512, then we would launch 432 blocks in total
|
||||
// say num_params is ~100M, each thread will process ~500 elements
|
||||
// warps reduce with warp-level reduce, we have 221,184/32 = 6,912 warps
|
||||
// and then each warp atomicAdd's to global memory, total of 6,912 atomics
|
||||
|
||||
// no shared memory; but one atomic per warp instead of per block
|
||||
namespace cg = cooperative_groups;
|
||||
cg::thread_block block = cg::this_thread_block();
|
||||
|
|
@ -84,8 +89,6 @@ __global__ void norm_kernel2(float* out, const T* data, size_t count) {
|
|||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
template<typename T>
|
||||
void global_norm1(float* out, const T* values, size_t count, int block_size) {
|
||||
// launch just enough blocks to fill the grid. deliberately no DIV_CEIL.
|
||||
|
|
|
|||
|
|
@ -1151,31 +1151,26 @@ __device__ float lerp(float start, float end, float weight) {
|
|||
template <typename Tp, typename Tg>
|
||||
__global__ void adamw_kernel3(Tp* params_memory, float* master_params_memory, Tg* grads_memory, float* m_memory, float* v_memory, size_t num_parameters,
|
||||
float learning_rate, float beta1, float beta2, float beta1_correction, float beta2_correction, float eps, float weight_decay,
|
||||
float* grad_norm, float max_grad_norm,
|
||||
float* grad_norm, float grad_clip,
|
||||
unsigned int seed) {
|
||||
int idx = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
if (idx >= num_parameters) { return; } // guard
|
||||
|
||||
float scale = 1.f;
|
||||
if(!isfinite(*grad_norm)) {
|
||||
// if we had a numerical problem (e.g, overflow)
|
||||
// in our gradient calculation, don't mess up the
|
||||
// existing weights.
|
||||
// in our gradient norm calculation, don't mess up the existing weights.
|
||||
// TODO increase a global counter somewhere so we actually know if/how often this happens
|
||||
if(threadIdx.x == 0 && blockIdx.x == 0) {
|
||||
if(threadIdx.x == 0 && blockIdx.x == 0) {
|
||||
printf("[WARNING] weight update skipped due to non-finite gradients!\n");
|
||||
}
|
||||
return;
|
||||
}
|
||||
if(*grad_norm > max_grad_norm) {
|
||||
scale = max_grad_norm / sqrtf(*grad_norm);
|
||||
// TODO just for debugging, remove this
|
||||
if(threadIdx.x == 0 && blockIdx.x == 0) {
|
||||
printf("[norm %f]\n", sqrtf(*grad_norm));
|
||||
}
|
||||
}
|
||||
// get the gradient, m, and v for this parameter
|
||||
float grad = scale * (float)grads_memory[idx];
|
||||
float grad = (float)grads_memory[idx];
|
||||
// clip the gradients if their norm surpasses grad_clip
|
||||
if(*grad_norm > grad_clip) {
|
||||
grad *= grad_clip / sqrtf(*grad_norm);
|
||||
}
|
||||
float m = m_memory[idx];
|
||||
float v = v_memory[idx];
|
||||
// update the first moment (momentum)
|
||||
|
|
@ -2388,7 +2383,7 @@ void gpt2_multi_gpu_accumulate(GPT2* model, MultiGpuConfig* multi_gpu_config) {
|
|||
#endif
|
||||
}
|
||||
|
||||
void gpt2_update(GPT2 *model, float learning_rate, float beta1, float beta2, float eps, float weight_decay, float grad_clipping, int t, MultiGpuConfig* multi_gpu_config) {
|
||||
void gpt2_update(GPT2 *model, float learning_rate, float beta1, float beta2, float eps, float weight_decay, float grad_clip, int t, MultiGpuConfig* multi_gpu_config) {
|
||||
NVTX_RANGE_FN();
|
||||
size_t num_parameters = multi_gpu_config->shard_num_parameters;
|
||||
floatX* params_memory = (floatX*)model->params_memory + multi_gpu_config->shard_offset;
|
||||
|
|
@ -2425,7 +2420,7 @@ void gpt2_update(GPT2 *model, float learning_rate, float beta1, float beta2, flo
|
|||
adamw_kernel3<<<num_blocks, block_size>>>(params_memory, model->master_weights, grads_memory,
|
||||
model->m_memory, model->v_memory, num_parameters,
|
||||
learning_rate, beta1, beta2, beta1_correction, beta2_correction, eps, weight_decay,
|
||||
grad_norm, grad_clipping,
|
||||
grad_norm, grad_clip,
|
||||
seed);
|
||||
cudaCheck(cudaGetLastError());
|
||||
}
|
||||
|
|
@ -2672,7 +2667,7 @@ int main(int argc, char *argv[]) {
|
|||
int use_master_weights = 1;
|
||||
int recompute = 1; // recompute during backward setting, 0 = none, 1 = recompute gelu
|
||||
int zero_stage = 0; // Zero Optimization Stage for Multi-GPU training
|
||||
float grad_clipping = 1.f;
|
||||
float grad_clip = 1.0f;
|
||||
for (int i = 1; i < argc; i+=2) {
|
||||
if (i + 1 >= argc) { error_usage(); } // must have arg after flag
|
||||
if (argv[i][0] != '-') { error_usage(); } // must start with dash
|
||||
|
|
@ -2693,7 +2688,7 @@ int main(int argc, char *argv[]) {
|
|||
else if (argv[i][1] == 'a') { overfit_single_batch = atoi(argv[i+1]); }
|
||||
else if (argv[i][1] == 'f') { override_enable_tf32 = atoi(argv[i+1]); }
|
||||
else if (argv[i][1] == 'w') { use_master_weights = atoi(argv[i+1]); }
|
||||
else if (argv[i][1] == 'c') { grad_clipping = atof(argv[i+1]); }
|
||||
else if (argv[i][1] == 'c') { grad_clip = atof(argv[i+1]); }
|
||||
else if (argv[i][1] == 'z') { zero_stage = atoi(argv[i+1]); }
|
||||
else if (argv[i][1] == 'r') { recompute = atoi(argv[i+1]); }
|
||||
else { error_usage(); }
|
||||
|
|
@ -2709,7 +2704,7 @@ int main(int argc, char *argv[]) {
|
|||
printf0("| sequence length T | %-50d |\n", T);
|
||||
printf0("| total batch size | %-50d |\n", total_batch_size);
|
||||
printf0("| learning rate | %-50e |\n", learning_rate);
|
||||
printf0("| grad_clipping | %-50e |\n", grad_clipping);
|
||||
printf0("| grad_clip | %-50e |\n", grad_clip);
|
||||
printf0("| max_steps | %-50d |\n", max_steps);
|
||||
printf0("| val_loss_every | %-50d |\n", val_loss_every);
|
||||
printf0("| val_max_batches | %-50d |\n", val_max_batches);
|
||||
|
|
@ -2894,7 +2889,7 @@ int main(int argc, char *argv[]) {
|
|||
model.mean_loss = lossf;
|
||||
// update the parameters
|
||||
gpt2_multi_gpu_accumulate(&model, &multi_gpu_config);
|
||||
gpt2_update(&model, learning_rate, 0.9f, 0.999f, 1e-8f, 0.0f, grad_clipping, step+1, &multi_gpu_config);
|
||||
gpt2_update(&model, learning_rate, 0.9f, 0.999f, 1e-8f, 0.0f, grad_clip, step+1, &multi_gpu_config);
|
||||
gpt2_multi_gpu_gather(&model, &multi_gpu_config);
|
||||
// zero out the gradients for the next iteration
|
||||
gpt2_zero_grad(&model);
|
||||
|
|
|
|||
|
|
@ -407,7 +407,7 @@ if __name__ == "__main__":
|
|||
parser.add_argument("--batch_size", type=int, default=4, help="batch size, in units of #batch dimensions")
|
||||
parser.add_argument("--sequence_length", type=int, default=64, help="sequence length")
|
||||
parser.add_argument("--total_batch_size", type=int, default=256, help="total desired batch size, in units of #tokens")
|
||||
parser.add_argument("--grad_clipping", type=float, default=1, help="maximum gradient magnitude")
|
||||
parser.add_argument("--grad_clip", type=float, default=1.0, help="maximum gradient magnitude")
|
||||
args = parser.parse_args()
|
||||
B, T = args.batch_size, args.sequence_length
|
||||
assert 1 <= T <= 1024
|
||||
|
|
@ -553,7 +553,7 @@ if __name__ == "__main__":
|
|||
if device == "cuda":
|
||||
torch.cuda.reset_peak_memory_stats()
|
||||
timings = []
|
||||
norm = -1 # dummy value to print in inference-only mode
|
||||
norm = -1.0 # dummy value to print in inference-only mode
|
||||
for step in range(args.num_iterations):
|
||||
t0 = time.time()
|
||||
|
||||
|
|
@ -577,7 +577,7 @@ if __name__ == "__main__":
|
|||
# backward pass
|
||||
if not args.inference_only:
|
||||
loss.backward()
|
||||
norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clipping)
|
||||
norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
|
||||
optimizer.step()
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue