mirror of
https://github.com/karpathy/llm.c.git
synced 2026-07-28 20:35:09 -04:00
update activation backward kernels that are not accumulated into the residual stream to overwrite instead of update.
This commit is contained in:
parent
de1e87158c
commit
60090dcf2e
1 changed files with 52 additions and 43 deletions
|
|
@ -288,9 +288,9 @@ __global__ void permute_kernel_backward(float* dinp,
|
|||
int d_ = rest % d;
|
||||
|
||||
int inp_idx = (b * N * 3 * NH * d) + (n * 3 * NH * d) + (0 * NH * d) + (nh_ * d) + d_;
|
||||
dinp[inp_idx] += dq[idx];
|
||||
dinp[inp_idx + NH * d] += dk[idx];
|
||||
dinp[inp_idx + 2 * (NH * d)] += dv[idx];
|
||||
dinp[inp_idx] = dq[idx];
|
||||
dinp[inp_idx + NH * d] = dk[idx];
|
||||
dinp[inp_idx + 2 * (NH * d)] = dv[idx];
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -323,7 +323,7 @@ __global__ void unpermute_kernel_backward(float* dinp, const float *dout, int B,
|
|||
int d_ = rest % d;
|
||||
|
||||
int other_idx = (b * NH * N * d) + (n * NH * d) + (nh_ * d) + d_;
|
||||
dinp[idx] += dout[other_idx];
|
||||
dinp[idx] = dout[other_idx];
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -399,7 +399,7 @@ __global__ void residual_forward_kernel(float* out, float* inp1, float* inp2, in
|
|||
}
|
||||
}
|
||||
|
||||
__global__ void residual_backward_kernel(float* dinp1, float* dinp2, float* dout, int N) {
|
||||
__global__ void residual_backward_kernel(float* dinp1, float* dinp2, const float* dout, int N) {
|
||||
int idx = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
if (idx < N) {
|
||||
dinp1[idx] += __ldcs(&dout[idx]);
|
||||
|
|
@ -427,7 +427,7 @@ __global__ void gelu_backward_kernel(float* dinp, const float* inp, const float*
|
|||
float coshf_out = coshf(tanh_arg);
|
||||
float sech_out = 1.0f / (coshf_out * coshf_out);
|
||||
float local_grad = 0.5f * (1.0f + tanh_out) + x * 0.5f * sech_out * GELU_SCALING_FACTOR * (1.0f + 3.0f * 0.044715f * x * x);
|
||||
dinp[i] += local_grad * dout[i];
|
||||
dinp[i] = local_grad * dout[i];
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -575,7 +575,7 @@ __global__ void crossentropy_softmax_backward_kernel1(float* dlogits,
|
|||
int ix = targets[b * T + t];
|
||||
float p = probs_bt[v];
|
||||
float indicator = v == ix ? 1.0f : 0.0f;
|
||||
dlogits_bt[v] += (p - indicator) * dloss;
|
||||
dlogits_bt[v] = (p - indicator) * dloss;
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -601,12 +601,12 @@ __global__ void matmul_backward_bias_kernel_faster(float* dbias, const float* do
|
|||
}
|
||||
// write the final result (at thread 0) to global memory
|
||||
if (tid == 0) {
|
||||
dbias[o] = shared[0];
|
||||
dbias[o] += shared[0];
|
||||
}
|
||||
}
|
||||
|
||||
__global__ void layernorm_backward_kernel(float* dinp, float* dweight, float* dbias,
|
||||
float* dout, float* inp, const float* weight, const float* mean, const float* rstd,
|
||||
const float* dout, const float* inp, const float* weight, const float* mean, const float* rstd,
|
||||
int B, int T, int C) {
|
||||
namespace cg = cooperative_groups;
|
||||
cg::thread_block block = cg::this_thread_block();
|
||||
|
|
@ -620,8 +620,8 @@ __global__ void layernorm_backward_kernel(float* dinp, float* dweight, float* db
|
|||
int b = idx / T;
|
||||
int t = idx % T;
|
||||
|
||||
float* dout_bt = dout + b * T * C + t * C;
|
||||
float* inp_bt = inp + b * T * C + t * C;
|
||||
const float* dout_bt = dout + b * T * C + t * C;
|
||||
const float* inp_bt = inp + b * T * C + t * C;
|
||||
float* dinp_bt = dinp + b * T * C + t * C;
|
||||
float mean_bt = mean[b * T + t];
|
||||
float rstd_bt = rstd[b * T + t];
|
||||
|
|
@ -695,7 +695,7 @@ __global__ void softmax_autoregressive_backward_kernel(float* dpreatt, const flo
|
|||
|
||||
result = cg::reduce(warp, result, cg::plus<float>());
|
||||
if(warp.thread_rank() == 0) {
|
||||
dpreatt_bth[t3] += scale * result;
|
||||
dpreatt_bth[t3] = scale * result;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -965,12 +965,12 @@ void crossentropy_softmax_backward(float* dlogits,
|
|||
void matmul_backward(float* dinp, float* dweight, float* dbias,
|
||||
float* dout, float* inp, float* weight,
|
||||
int B, int T, int C, int OC) {
|
||||
float alpha = 1.0f;
|
||||
float beta = 1.0f; // note we must use beta = 1.0 so that we do a +=, as we should, because gradients add
|
||||
float one = 1.0f;
|
||||
float zero = 0.0f;
|
||||
// backward to input
|
||||
cublasCheck(cublasSgemm(cublas_handle, CUBLAS_OP_N, CUBLAS_OP_N, C, B*T, OC, &alpha, weight, C, dout, OC, &beta, dinp, C));
|
||||
cublasCheck(cublasSgemm(cublas_handle, CUBLAS_OP_N, CUBLAS_OP_N, C, B*T, OC, &one, weight, C, dout, OC, &zero, dinp, C));
|
||||
// backward to weight
|
||||
cublasCheck(cublasSgemm(cublas_handle, CUBLAS_OP_N, CUBLAS_OP_T, C, OC, B*T, &alpha, inp, C, dout, OC, &beta, dweight, C));
|
||||
cublasCheck(cublasSgemm(cublas_handle, CUBLAS_OP_N, CUBLAS_OP_T, C, OC, B*T, &one, inp, C, dout, OC, &one, dweight, C));
|
||||
// backward to bias, if given
|
||||
if (dbias != NULL) {
|
||||
const int block_size=512;
|
||||
|
|
@ -983,7 +983,7 @@ void matmul_backward(float* dinp, float* dweight, float* dbias,
|
|||
}
|
||||
|
||||
void layernorm_backward(float* dinp, float* dweight, float* dbias,
|
||||
float* dout, float* inp, float* weight, float* mean, float* rstd,
|
||||
const float* dout, const float* inp, const float* weight, const float* mean, const float* rstd,
|
||||
int B, int T, int C) {
|
||||
const int block_size = 256;
|
||||
const int N = B * T;
|
||||
|
|
@ -1154,7 +1154,7 @@ typedef struct {
|
|||
float* v_accum; // (L, B, T, C)
|
||||
} ActivationTensors;
|
||||
|
||||
float* malloc_and_point_activations(ActivationTensors* acts, size_t* act_sizes) {
|
||||
float* malloc_and_point_activations(ActivationTensors* acts, const size_t* act_sizes) {
|
||||
size_t num_activations = 0;
|
||||
for (size_t i = 0; i < NUM_ACTIVATION_TENSORS; i++) {
|
||||
num_activations += act_sizes[i];
|
||||
|
|
@ -1288,6 +1288,39 @@ void gpt2_build_from_checkpoint(GPT2 *model, const char* checkpoint_path) {
|
|||
model->mean_loss = -1.0f; // -1.0f will designate no loss
|
||||
}
|
||||
|
||||
void fill_in_activation_sizes(size_t* act_sizes, int B, int T, GPT2Config config) {
|
||||
int V = config.vocab_size;
|
||||
int L = config.num_layers;
|
||||
int NH = config.num_heads;
|
||||
int C = config.channels;
|
||||
|
||||
act_sizes[0] = B * T * C; // encoded
|
||||
act_sizes[1] = L * B * T * C; // ln1
|
||||
act_sizes[2] = L * B * T; // ln1_mean
|
||||
act_sizes[3] = L * B * T; // ln1_rstd
|
||||
act_sizes[4] = L * B * T * 3*C; // qkv
|
||||
act_sizes[5] = L * B * T * C; // atty
|
||||
act_sizes[6] = L * B * NH * T * T; // preatt
|
||||
act_sizes[7] = L * B * NH * T * T; // att
|
||||
act_sizes[8] = L * B * T * C; // attproj
|
||||
act_sizes[9] = L * B * T * C; // residual2
|
||||
act_sizes[10] = L * B * T * C; // ln2
|
||||
act_sizes[11] = L * B * T; // ln2_mean
|
||||
act_sizes[12] = L * B * T; // ln2_rstd
|
||||
act_sizes[13] = L * B * T * 4*C; // fch
|
||||
act_sizes[14] = L * B * T * 4*C; // fch_gelu
|
||||
act_sizes[15] = L * B * T * C; // fcproj
|
||||
act_sizes[16] = L * B * T * C; // residual3
|
||||
act_sizes[17] = B * T * C; // lnf
|
||||
act_sizes[18] = B * T; // lnf_mean
|
||||
act_sizes[19] = B * T; // lnf_rstd
|
||||
act_sizes[20] = B * T * V; // logits
|
||||
act_sizes[21] = B * T * V; // probs
|
||||
act_sizes[22] = B * T; // losses
|
||||
act_sizes[23] = L * B * T * 3*C; // qkvr
|
||||
act_sizes[24] = L * B * T * C; // v_accum
|
||||
}
|
||||
|
||||
void gpt2_forward(GPT2 *model, int* inputs, int* targets, int B, int T) {
|
||||
// targets are optional and could be NULL
|
||||
|
||||
|
|
@ -1317,31 +1350,7 @@ void gpt2_forward(GPT2 *model, int* inputs, int* targets, int B, int T) {
|
|||
model->batch_size = B;
|
||||
model->seq_len = T;
|
||||
// and now allocate the space
|
||||
model->act_sizes[0] = B * T * C; // encoded
|
||||
model->act_sizes[1] = L * B * T * C; // ln1
|
||||
model->act_sizes[2] = L * B * T; // ln1_mean
|
||||
model->act_sizes[3] = L * B * T; // ln1_rstd
|
||||
model->act_sizes[4] = L * B * T * 3*C; // qkv
|
||||
model->act_sizes[5] = L * B * T * C; // atty
|
||||
model->act_sizes[6] = L * B * NH * T * T; // preatt
|
||||
model->act_sizes[7] = L * B * NH * T * T; // att
|
||||
model->act_sizes[8] = L * B * T * C; // attproj
|
||||
model->act_sizes[9] = L * B * T * C; // residual2
|
||||
model->act_sizes[10] = L * B * T * C; // ln2
|
||||
model->act_sizes[11] = L * B * T; // ln2_mean
|
||||
model->act_sizes[12] = L * B * T; // ln2_rstd
|
||||
model->act_sizes[13] = L * B * T * 4*C; // fch
|
||||
model->act_sizes[14] = L * B * T * 4*C; // fch_gelu
|
||||
model->act_sizes[15] = L * B * T * C; // fcproj
|
||||
model->act_sizes[16] = L * B * T * C; // residual3
|
||||
model->act_sizes[17] = B * T * C; // lnf
|
||||
model->act_sizes[18] = B * T; // lnf_mean
|
||||
model->act_sizes[19] = B * T; // lnf_rstd
|
||||
model->act_sizes[20] = B * T * V; // logits
|
||||
model->act_sizes[21] = B * T * V; // probs
|
||||
model->act_sizes[22] = B * T; // losses
|
||||
model->act_sizes[23] = L * B * T * 3*C; // qkvr
|
||||
model->act_sizes[24] = L * B * T * C; // v_accum
|
||||
fill_in_activation_sizes(model->act_sizes, B, T, model->config);
|
||||
size_t num_activations = 0;
|
||||
for (size_t i = 0; i < NUM_ACTIVATION_TENSORS; i++) {
|
||||
num_activations += model->act_sizes[i];
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue