| 6242 | } |
| 6243 | |
| 6244 | static 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], |
no test coverage detected