Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
14 commits
Select commit Hold shift + click to select a range
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions examples/training/finetune-lora.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,8 @@ int main(int argc, char ** argv) {
/*optimizer_type =*/params.optimizer,
/*adapter =*/adapter,
/*recompute =*/false, // keep every layer's intermediates, as before
/*walk_exact =*/false, // the per-chunk walk, as before
/*walk_horizon =*/0,
};
llama_opt_init(ctx, model, lopt_params);

Expand Down
2 changes: 2 additions & 0 deletions examples/training/finetune.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,8 @@ int main(int argc, char ** argv) {
/*optimizer_type =*/params.optimizer,
/*adapter =*/nullptr, // full-model training: every tensor the filter allows
/*recompute =*/false, // keep every layer's intermediates, as before
/*walk_exact =*/false, // the per-chunk walk, as before
/*walk_horizon =*/0,
};
llama_opt_init(ctx, model, lopt_params);

Expand Down
19 changes: 19 additions & 0 deletions ggml/include/ggml-opt.h
Original file line number Diff line number Diff line change
Expand Up @@ -211,6 +211,25 @@ extern "C" {
// do forward pass, increment result if not NULL, do backward pass if allocated
GGML_API void ggml_opt_eval(ggml_opt_context_t opt_ctx, ggml_opt_result_t result);

// A caller-driven optimizer period, for graphs built per step (the training walk's reverse
// pass: one graph per chunk, ONE optimizer step per window). Applies to the next
// ggml_opt_alloc + ggml_opt_eval only:
// period_end: true = this graph runs the optimizer step, false = it only accumulates
// loss_scale: the weight of this graph's loss in the period's total (replaces 1/opt_period)
// extra_loss: NULL, or a scalar F32 node of the forward graph added to the loss the backward
// differentiates (a surrogate carrying a later graph's gradient into this one);
// the loss the result reports stays the unweighted loss of the outputs
// A period starts from zero gradients after the step that ended the previous one.
GGML_API void ggml_opt_set_next_step(
ggml_opt_context_t opt_ctx,
bool period_end,
float loss_scale,
struct ggml_tensor * extra_loss);

// After ggml_opt_eval and until the next ggml_opt_alloc: the gradient the backward computed
// for a GRAD leaf (ggml_set_grad) of the evaluated graph, or NULL when the graph has none.
GGML_API struct ggml_tensor * ggml_opt_leaf_grad(ggml_opt_context_t opt_ctx, struct ggml_tensor * leaf);

// ############################################################################
// ## The high-level functions start here. They do not depend on any private ##
// ## functions or structs and can be copied to and adapted for user code. ##
Expand Down
4 changes: 4 additions & 0 deletions ggml/include/ggml.h
Original file line number Diff line number Diff line change
Expand Up @@ -654,6 +654,7 @@ extern "C" {
GGML_TENSOR_FLAG_PARAM = 4, // ...contains trainable parameters
GGML_TENSOR_FLAG_LOSS = 8, // ...defines loss for numerical optimization (multiple loss tensors add up)
GGML_TENSOR_FLAG_COMPUTE = 16, // ...must be computed
GGML_TENSOR_FLAG_GRAD = 32, // ...needs its gradient computed, but no optimizer updates it (a leaf whose gradient is read back)
};

enum ggml_tri_type {
Expand Down Expand Up @@ -883,6 +884,9 @@ extern "C" {
GGML_API void ggml_set_output(struct ggml_tensor * tensor);
GGML_API void ggml_set_param(struct ggml_tensor * tensor);
GGML_API void ggml_set_loss(struct ggml_tensor * tensor);
// the backward computes this leaf's gradient (readable after compute), and no optimizer
// step touches it: an input whose sensitivity is needed, never a trainable parameter
GGML_API void ggml_set_grad(struct ggml_tensor * tensor);

//
// operations on tensors with backpropagation
Expand Down
146 changes: 110 additions & 36 deletions ggml/src/ggml-opt.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -67,9 +67,26 @@ struct ggml_opt_context {
size_t alloc_budget = 0; // bytes a graph may add on a non-CPU device; 0 = the device's own free figure
size_t peak_graph_bytes = 0; // largest graph measured by the preflight on a non-CPU device
bool eval_ready = false;
// Accumulators and momenta belong to a PARAMETER (or to the loss), never to a node index: a
// graph built per step may differ in topology from the first one (the training walk's chunks
// with and without a prefix, its reverse pass with GRAD leaves and a surrogate term), and an
// index would bind one tensor's accumulator to whatever node took that index. grad_accs is
// the current graph's view of them, rebuilt by every ggml_opt_build.
std::unordered_map<const struct ggml_tensor *, struct ggml_tensor *> param_grad_acc;
std::unordered_map<const struct ggml_tensor *, struct ggml_tensor *> param_m;
std::unordered_map<const struct ggml_tensor *, struct ggml_tensor *> param_v;
struct ggml_tensor * loss_grad_acc = nullptr;
bool accumulators_created = false;

// ggml_opt_set_next_step: a caller-driven period, for the next graph only
bool next_manual = false;
bool next_period_end = false;
float next_loss_scale = 1.0f;
struct ggml_tensor * next_extra_loss = nullptr;
bool period_fresh = true; // the next backward starts a period (zeroed gradients)
// the gradients of the evaluated graph's GRAD leaves, readable until the next alloc
std::unordered_map<const struct ggml_tensor *, struct ggml_tensor *> leaf_grads;
std::vector<struct ggml_tensor *> grad_accs;
std::vector<struct ggml_tensor *> grad_m;
std::vector<struct ggml_tensor *> grad_v;

int64_t iter = 1;
int32_t opt_period = 1;
Expand Down Expand Up @@ -556,7 +573,9 @@ static void ggml_opt_build(ggml_opt_context_t opt_ctx) {
ggml_set_name(opt_ctx->labels, "labels");
opt_ctx->loss = ggml_cross_entropy_loss(ctx_results, opt_ctx->outputs, opt_ctx->labels);
ggml_set_name(opt_ctx->loss, "loss_cross_entropy");
if (opt_ctx->opt_period > 1) {
if (opt_ctx->next_manual) {
// weighted by the caller below, beside its extra term
} else if (opt_ctx->opt_period > 1) {
opt_ctx->loss = ggml_scale(ctx_results, opt_ctx->loss, 1.0f / opt_ctx->opt_period);
ggml_set_name(opt_ctx->loss, "loss_cross_entropy_scaled");
}
Expand All @@ -581,8 +600,24 @@ static void ggml_opt_build(ggml_opt_context_t opt_ctx) {
}
}
ggml_set_output(opt_ctx->loss);
ggml_set_loss(opt_ctx->loss);
ggml_build_forward_expand(opt_ctx->gf, opt_ctx->loss);
if (opt_ctx->next_manual) {
GGML_ASSERT(!opt_ctx->static_graphs && "a caller-driven period needs graphs built per step");
GGML_ASSERT(opt_ctx->loss_type == GGML_OPT_LOSS_TYPE_CROSS_ENTROPY);
// the loss the backward differentiates: the caller's weight of this graph's loss, plus
// its extra term; opt_ctx->loss stays the unweighted loss the result reports
struct ggml_tensor * total = ggml_scale(ctx_results, opt_ctx->loss, opt_ctx->next_loss_scale);
if (opt_ctx->next_extra_loss) {
GGML_ASSERT(ggml_is_scalar(opt_ctx->next_extra_loss) && opt_ctx->next_extra_loss->type == GGML_TYPE_F32);
total = ggml_add(ctx_results, total, opt_ctx->next_extra_loss);
}
ggml_set_name(total, "loss_total");
ggml_set_loss(total);
ggml_build_forward_expand(opt_ctx->gf, opt_ctx->loss);
ggml_build_forward_expand(opt_ctx->gf, total);
} else {
ggml_set_loss(opt_ctx->loss);
ggml_build_forward_expand(opt_ctx->gf, opt_ctx->loss);
}

if (opt_ctx->loss_type == GGML_OPT_LOSS_TYPE_CROSS_ENTROPY) {
opt_ctx->pred = ggml_argmax(ctx_results, opt_ctx->outputs);
Expand All @@ -606,39 +641,59 @@ static void ggml_opt_build(ggml_opt_context_t opt_ctx) {
return;
}

if (opt_ctx->grad_accs.empty()) {
if (!opt_ctx->accumulators_created) {
GGML_ASSERT(opt_ctx->build_type_alloc >= GGML_OPT_BUILD_TYPE_GRAD);
opt_ctx->accumulators_created = true;

const int n_nodes = opt_ctx->gf->n_nodes;
opt_ctx->grad_accs.resize(n_nodes);
for (int i = 0; i < n_nodes; ++i) {
// created once, in ctx_static, for the parameters of the first graph: every later graph
// trains the same parameters (asserted below), whatever else its topology holds
for (int i = 0; i < opt_ctx->gf->n_nodes; ++i) {
ggml_tensor * node = opt_ctx->gf->nodes[i];
if ((accumulate && (node->flags & GGML_TENSOR_FLAG_PARAM)) || (node->flags & GGML_TENSOR_FLAG_LOSS)) {
opt_ctx->grad_accs[i] = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne);
} else {
opt_ctx->grad_accs[i] = nullptr;
if (node->flags & GGML_TENSOR_FLAG_LOSS) {
opt_ctx->loss_grad_acc = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne);
}
if (!(node->flags & GGML_TENSOR_FLAG_PARAM)) {
continue;
}
if (accumulate) {
opt_ctx->param_grad_acc[node] = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne);
}
if (need_momenta && opt_ctx->build_type_alloc >= GGML_OPT_BUILD_TYPE_OPT) {
opt_ctx->param_m[node] = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne);
opt_ctx->param_v[node] = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne);
}
}
}

if (need_momenta && opt_ctx->build_type_alloc >= GGML_OPT_BUILD_TYPE_OPT) {
opt_ctx->grad_m.resize(n_nodes);
opt_ctx->grad_v.resize(n_nodes);
for (int i = 0; i < n_nodes; ++i) {
ggml_tensor * node = opt_ctx->gf->nodes[i];
if (node->flags & GGML_TENSOR_FLAG_PARAM) {
opt_ctx->grad_m[i] = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne);
opt_ctx->grad_v[i] = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne);
} else {
opt_ctx->grad_m[i] = nullptr;
opt_ctx->grad_v[i] = nullptr;
}
// this graph's view of the accumulators, by tensor
opt_ctx->grad_accs.assign(opt_ctx->gf->n_nodes, nullptr);
for (int i = 0; i < opt_ctx->gf->n_nodes; ++i) {
ggml_tensor * node = opt_ctx->gf->nodes[i];
if (node->flags & GGML_TENSOR_FLAG_LOSS) {
opt_ctx->grad_accs[i] = opt_ctx->loss_grad_acc;
} else if (node->flags & GGML_TENSOR_FLAG_PARAM) {
GGML_ASSERT((!accumulate || opt_ctx->param_grad_acc.count(node)) &&
"a graph trains a parameter the optimizer context was not built with");
if (accumulate) {
opt_ctx->grad_accs[i] = opt_ctx->param_grad_acc.at(node);
}
}
}

// gb_grad == graph backward gradients, forward pass, then backward pass to calculate gradients.
opt_ctx->gb_grad = ggml_graph_dup(opt_ctx->ctx_compute, opt_ctx->gf, /*force_grads =*/ true);
ggml_build_backward_expand(opt_ctx->ctx_compute, opt_ctx->gb_grad, opt_ctx->grad_accs.data());
opt_ctx->leaf_grads.clear();
for (int i = 0; i < opt_ctx->gf->n_nodes; ++i) {
ggml_tensor * node = opt_ctx->gf->nodes[i];
if (node->flags & GGML_TENSOR_FLAG_GRAD) {
ggml_tensor * grad = ggml_graph_get_grad(opt_ctx->gb_grad, node);
if (grad) {
ggml_set_output(grad); // read back after the step: the allocator must not reuse it
opt_ctx->leaf_grads[node] = grad;
}
}
}
if (!opt_ctx->checkpoint_prefix.empty()) {
ggml_opt_checkpoint(opt_ctx->ctx_compute, opt_ctx->gb_grad, opt_ctx->gf->n_nodes, opt_ctx->checkpoint_prefix.c_str());
}
Expand Down Expand Up @@ -670,8 +725,8 @@ static void ggml_opt_build(ggml_opt_context_t opt_ctx) {
struct ggml_tensor * m = nullptr;
struct ggml_tensor * v = nullptr;
if (need_momenta) {
m = opt_ctx->grad_m[i];
v = opt_ctx->grad_v[i];
m = opt_ctx->param_m.at(node);
v = opt_ctx->param_v.at(node);
ggml_format_name(m, "AdamW m for %s", node->name);
ggml_format_name(v, "AdamW v for %s", node->name);
}
Expand Down Expand Up @@ -748,6 +803,19 @@ void ggml_opt_free(ggml_opt_context_t opt_ctx) {
delete opt_ctx;
}

void ggml_opt_set_next_step(ggml_opt_context_t opt_ctx, bool period_end, float loss_scale, struct ggml_tensor * extra_loss) {
GGML_ASSERT(!opt_ctx->eval_ready && "set the next step before ggml_opt_alloc");
opt_ctx->next_manual = true;
opt_ctx->next_period_end = period_end;
opt_ctx->next_loss_scale = loss_scale;
opt_ctx->next_extra_loss = extra_loss;
}

struct ggml_tensor * ggml_opt_leaf_grad(ggml_opt_context_t opt_ctx, struct ggml_tensor * leaf) {
const auto it = opt_ctx->leaf_grads.find(leaf);
return it == opt_ctx->leaf_grads.end() ? nullptr : it->second;
}

void ggml_opt_reset(ggml_opt_context_t opt_ctx, bool optimizer) {
if (optimizer) {
ggml_graph_reset(opt_ctx->gb_opt);
Expand Down Expand Up @@ -897,7 +965,9 @@ bool ggml_opt_alloc(ggml_opt_context_t opt_ctx, bool backward) {
if (opt_ctx->build_type == GGML_OPT_BUILD_TYPE_OPT && opt_ctx->opt_period > 1 && opt_ctx->opt_i == 0) {
ggml_graph_reset(opt_ctx->gb_grad);
}
if (backward) {
if (backward && opt_ctx->next_manual) {
opt_ctx->build_type = opt_ctx->next_period_end ? GGML_OPT_BUILD_TYPE_OPT : GGML_OPT_BUILD_TYPE_GRAD;
} else if (backward) {
const int32_t opt_i_next = (opt_ctx->opt_i + 1) % opt_ctx->opt_period;
opt_ctx->build_type = opt_i_next == 0 ? GGML_OPT_BUILD_TYPE_OPT : GGML_OPT_BUILD_TYPE_GRAD;
} else {
Expand All @@ -911,14 +981,9 @@ bool ggml_opt_alloc(ggml_opt_context_t opt_ctx, bool backward) {
// and the reset above found no graph to reset (gb_grad is rebuilt every step). Without
// this, a period's step applied the SUM of every gradient since the run began, since
// each backward adds into the accumulator in place (test-opt-dynamic-accum).
if (backward && opt_ctx->opt_i == 0) {
// grad_accs is indexed by the FIRST graph's nodes: a later graph may hold fewer (Fable)
const size_t n = std::min(opt_ctx->grad_accs.size(), (size_t) opt_ctx->gf->n_nodes);
for (size_t i = 0; i < n; ++i) {
ggml_tensor * acc = opt_ctx->grad_accs[i];
if (acc && (opt_ctx->gf->nodes[i]->flags & GGML_TENSOR_FLAG_PARAM)) {
ggml_set_zero(acc);
}
if (backward && (opt_ctx->next_manual ? opt_ctx->period_fresh : opt_ctx->opt_i == 0)) {
for (auto & [param, acc] : opt_ctx->param_grad_acc) {
ggml_set_zero(acc);
}
}
}
Expand Down Expand Up @@ -1115,6 +1180,15 @@ void ggml_opt_eval(ggml_opt_context_t opt_ctx, ggml_opt_result_t result) {
}
opt_ctx->iter += opt_ctx->allocated_graph == opt_ctx->gb_opt;
opt_ctx->opt_i = (opt_ctx->opt_i + 1) % opt_ctx->opt_period;
if (opt_ctx->allocated_graph == opt_ctx->gb_opt) {
opt_ctx->period_fresh = true;
} else if (opt_ctx->allocated_graph == opt_ctx->gb_grad) {
opt_ctx->period_fresh = false;
}
opt_ctx->next_manual = false;
opt_ctx->next_period_end = false;
opt_ctx->next_loss_scale = 1.0f;
opt_ctx->next_extra_loss = nullptr;

if (!opt_ctx->static_graphs) {
opt_ctx->gf = nullptr;
Expand Down
20 changes: 15 additions & 5 deletions ggml/src/ggml.c
Original file line number Diff line number Diff line change
Expand Up @@ -7115,18 +7115,23 @@ static void ggml_compute_backward(
// layer concatenates the cached conv state (a leaf) with this ubatch's qkv, so the
// second slice is the one that carries a gradient back into the model.
const int dim = ggml_get_op_params_i32(tensor, 0);
// A view's first stride is the element size (ggml_view_4d takes nb[1..3] only), so the
// gradient is sliced from a CONTIGUOUS copy: one that arrives through a transpose (the
// training attention transposes V for its KQV product) has a first stride that is not,
// and slicing it in place read the right number of elements from the wrong places.
struct ggml_tensor * g = ggml_is_contiguous(grad) ? grad : ggml_cont(ctx, grad);
size_t offset = 0;
for (int j = 0; j < 2; ++j) {
struct ggml_tensor * src = tensor->src[j];
const size_t isrc = j == 0 ? isrc0 : isrc1;
const bool needs = j == 0 ? src0_needs_grads : src1_needs_grads;
if (needs) {
struct ggml_tensor * slice = ggml_view_4d(ctx, grad,
struct ggml_tensor * slice = ggml_view_4d(ctx, g,
src->ne[0], src->ne[1], src->ne[2], src->ne[3],
grad->nb[1], grad->nb[2], grad->nb[3], offset);
g->nb[1], g->nb[2], g->nb[3], offset);
ggml_add_or_set(ctx, cgraph, isrc, ggml_cont(ctx, slice));
}
offset += src->ne[dim] * grad->nb[dim];
offset += src->ne[dim] * g->nb[dim];
}
} break;
case GGML_OP_SSM_CONV: {
Expand Down Expand Up @@ -7348,7 +7353,7 @@ static size_t ggml_visit_parents_graph(struct ggml_cgraph * cgraph, struct ggml_
}
}

if (node->op == GGML_OP_NONE && !(node->flags & GGML_TENSOR_FLAG_PARAM)) {
if (node->op == GGML_OP_NONE && !(node->flags & (GGML_TENSOR_FLAG_PARAM | GGML_TENSOR_FLAG_GRAD))) {
// reached a leaf node, not part of the gradient graph (e.g. a constant)
GGML_ASSERT(cgraph->n_leafs < cgraph->size);

Expand Down Expand Up @@ -7446,7 +7451,7 @@ void ggml_build_backward_expand(
continue;
}

bool node_needs_grad = (node->flags & GGML_TENSOR_FLAG_PARAM) || (node->flags & GGML_TENSOR_FLAG_LOSS);
bool node_needs_grad = (node->flags & (GGML_TENSOR_FLAG_PARAM | GGML_TENSOR_FLAG_LOSS | GGML_TENSOR_FLAG_GRAD)) != 0;
bool ignore_src[GGML_MAX_SRC] = {false};
switch (node->op) {
// gradients in node->src[0] for one reason or another have no effect on output gradients
Expand Down Expand Up @@ -8077,6 +8082,11 @@ void ggml_set_param(struct ggml_tensor * tensor) {
tensor->flags |= GGML_TENSOR_FLAG_PARAM;
}

void ggml_set_grad(struct ggml_tensor * tensor) {
GGML_ASSERT(tensor->op == GGML_OP_NONE);
tensor->flags |= GGML_TENSOR_FLAG_GRAD;
}

void ggml_set_loss(struct ggml_tensor * tensor) {
GGML_ASSERT(ggml_is_scalar(tensor));
GGML_ASSERT(tensor->type == GGML_TYPE_F32);
Expand Down
Loading
Loading