| 7029 | } |
| 7030 | |
| 7031 | void ggml_graph_cpy(struct ggml_cgraph * src, struct ggml_cgraph * dst) { |
| 7032 | GGML_ASSERT(dst->size >= src->n_leafs); |
| 7033 | GGML_ASSERT(dst->size >= src->n_nodes); |
| 7034 | GGML_ASSERT(dst->visited_hash_set.size >= src->visited_hash_set.size); |
| 7035 | |
| 7036 | dst->n_leafs = src->n_leafs; |
| 7037 | dst->n_nodes = src->n_nodes; |
| 7038 | dst->order = src->order; |
| 7039 | |
| 7040 | for (int i = 0; i < src->n_leafs; ++i) { |
| 7041 | dst->leafs[i] = src->leafs[i]; |
| 7042 | } |
| 7043 | |
| 7044 | for (int i = 0; i < src->n_nodes; ++i) { |
| 7045 | dst->nodes[i] = src->nodes[i]; |
| 7046 | } |
| 7047 | |
| 7048 | for (size_t i = 0; i < src->visited_hash_set.size; ++i) { |
| 7049 | // copy all hashset keys (tensors) that are in use |
| 7050 | if (ggml_bitset_get(src->visited_hash_set.used, i)) { |
| 7051 | size_t new_hash_pos = ggml_hash_insert(&dst->visited_hash_set, src->visited_hash_set.keys[i]); |
| 7052 | dst->use_counts[new_hash_pos] = src->use_counts[i]; |
| 7053 | } |
| 7054 | } |
| 7055 | |
| 7056 | if (dst->grads) { |
| 7057 | memset(dst->grads, 0, dst->visited_hash_set.size*sizeof(struct ggml_tensor *)); |
| 7058 | memset(dst->grad_accs, 0, dst->visited_hash_set.size*sizeof(struct ggml_tensor *)); |
| 7059 | } |
| 7060 | if (src->grads) { |
| 7061 | GGML_ASSERT(dst->grads != NULL); |
| 7062 | GGML_ASSERT(dst->grad_accs != NULL); |
| 7063 | for (int i = 0; i < src->n_nodes; ++i) { |
| 7064 | const size_t igrad_src = ggml_hash_find(&src->visited_hash_set, src->nodes[i]); |
| 7065 | const size_t igrad_dst = ggml_hash_find(&dst->visited_hash_set, dst->nodes[i]); |
| 7066 | |
| 7067 | GGML_ASSERT(igrad_src != GGML_HASHSET_FULL); |
| 7068 | GGML_ASSERT(ggml_bitset_get(src->visited_hash_set.used, igrad_src)); |
| 7069 | GGML_ASSERT(igrad_dst != GGML_HASHSET_FULL); |
| 7070 | GGML_ASSERT(ggml_bitset_get(dst->visited_hash_set.used, igrad_dst)); |
| 7071 | |
| 7072 | dst->grads[igrad_dst] = src->grads[igrad_src]; |
| 7073 | dst->grad_accs[igrad_dst] = src->grad_accs[igrad_src]; |
| 7074 | } |
| 7075 | } |
| 7076 | } |
| 7077 | |
| 7078 | struct ggml_cgraph * ggml_graph_dup(struct ggml_context * ctx, struct ggml_cgraph * cgraph, bool force_grads) { |
| 7079 | struct ggml_cgraph * result = ggml_new_graph_custom(ctx, cgraph->size, cgraph->grads || force_grads); |