MCPcopy Create free account
hub / github.com/appdevforall/CodeOnTheGo / ggml_compute_backward

Function ggml_compute_backward

subprojects/llama.cpp/ggml/src/ggml.c:6244–6727  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

6242}
6243
6244static void ggml_compute_backward(
6245 struct ggml_context * ctx, struct ggml_cgraph * cgraph, int i, const bool * grads_needed) {
6246 struct ggml_tensor * tensor = cgraph->nodes[i];
6247 struct ggml_tensor * grad = ggml_graph_get_grad(cgraph, tensor);
6248
6249 if (!grad) {
6250 return;
6251 }
6252
6253 struct ggml_tensor * src0 = tensor->src[0];
6254 struct ggml_tensor * src1 = tensor->src[1];
6255 struct ggml_tensor * src2 = tensor->src[2];
6256 struct ggml_hash_set * hash_set = &cgraph->visited_hash_set;
6257 const size_t isrc0 = src0 ? ggml_hash_find(hash_set, src0) : (size_t) -1;
6258 const size_t isrc1 = src1 ? ggml_hash_find(hash_set, src1) : (size_t) -1;
6259 const size_t isrc2 = src2 ? ggml_hash_find(hash_set, src2) : (size_t) -1;
6260 const bool src0_needs_grads = src0 && isrc0 != GGML_HASHSET_FULL && ggml_bitset_get(hash_set->used, isrc0) && grads_needed[isrc0];
6261 const bool src1_needs_grads = src1 && isrc1 != GGML_HASHSET_FULL && ggml_bitset_get(hash_set->used, isrc1) && grads_needed[isrc1];
6262 const bool src2_needs_grads = src2 && isrc2 != GGML_HASHSET_FULL && ggml_bitset_get(hash_set->used, isrc2) && grads_needed[isrc2];
6263
6264 switch (tensor->op) {
6265 case GGML_OP_DUP: {
6266 if (src0_needs_grads) {
6267 ggml_add_or_set(ctx, cgraph, isrc0, grad);
6268 }
6269 } break;
6270 case GGML_OP_ADD: {
6271 if (src0_needs_grads) {
6272 ggml_add_or_set(ctx, cgraph, isrc0, grad);
6273 }
6274 if (src1_needs_grads) {
6275 struct ggml_tensor * tmp = grad;
6276 if (!ggml_are_same_shape(src0, src1)) {
6277 tmp = ggml_repeat_back(ctx, tmp, src1);
6278 }
6279 ggml_add_or_set(ctx, cgraph, isrc1, tmp);
6280 }
6281 } break;
6282 case GGML_OP_ADD1: {
6283 if (src0_needs_grads) {
6284 ggml_add_or_set(ctx, cgraph, isrc0, grad);
6285 }
6286 if (src1_needs_grads) {
6287 ggml_add_or_set(ctx, cgraph, isrc1, ggml_mean(ctx, grad)); // TODO: should probably be sum instead of mean
6288 }
6289 } break;
6290 case GGML_OP_ACC: {
6291 if (src0_needs_grads) {
6292 ggml_add_or_set(ctx, cgraph, isrc0, grad);
6293 }
6294 if (src1_needs_grads) {
6295 const size_t nb1 = ((int32_t *) tensor->op_params)[0];
6296 const size_t nb2 = ((int32_t *) tensor->op_params)[1];
6297 const size_t nb3 = ((int32_t *) tensor->op_params)[2];
6298 const size_t offset = ((int32_t *) tensor->op_params)[3];
6299
6300 struct ggml_tensor * tensor_grad_view = ggml_view_4d(ctx,
6301 grad, src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3],

Callers 1

Calls 15

ggml_graph_get_gradFunction · 0.85
ggml_hash_findFunction · 0.85
ggml_bitset_getFunction · 0.85
ggml_add_or_setFunction · 0.85
ggml_are_same_shapeFunction · 0.85
ggml_repeat_backFunction · 0.85
ggml_meanFunction · 0.85
ggml_view_4dFunction · 0.85
ggml_reshapeFunction · 0.85
ggml_contFunction · 0.85
ggml_sub_or_setFunction · 0.85
ggml_mulFunction · 0.85

Tested by

no test coverage detected