| 6189 | } |
| 6190 | |
| 6191 | static void ggml_acc_or_set( |
| 6192 | struct ggml_context * ctx, |
| 6193 | struct ggml_cgraph * cgraph, |
| 6194 | size_t isrc, |
| 6195 | struct ggml_tensor * tensor, |
| 6196 | const size_t nb1, |
| 6197 | const size_t nb2, |
| 6198 | const size_t nb3, |
| 6199 | const size_t offset) { |
| 6200 | struct ggml_tensor * src = cgraph->visited_hash_set.keys[isrc]; |
| 6201 | GGML_ASSERT(src); |
| 6202 | if (cgraph->grads[isrc]) { |
| 6203 | cgraph->grads[isrc] = ggml_acc_impl(ctx, cgraph->grads[isrc], tensor, nb1, nb2, nb3, offset, cgraph->grad_accs[isrc]); |
| 6204 | } else { |
| 6205 | struct ggml_tensor * a_zero = ggml_scale(ctx, src, 0.0f); // FIXME this is going to produce NaN if a contains inf/NaN |
| 6206 | cgraph->grads[isrc] = ggml_acc_impl(ctx, a_zero, tensor, nb1, nb2, nb3, offset, false); |
| 6207 | } |
| 6208 | ggml_format_name(cgraph->grads[isrc], "grad for %s", cgraph->visited_hash_set.keys[isrc]->name); |
| 6209 | ggml_build_forward_expand(cgraph, cgraph->grads[isrc]); |
| 6210 | } |
| 6211 | |
| 6212 | static void ggml_add1_or_set( |
| 6213 | struct ggml_context * ctx, |
no test coverage detected