diff --git a/examples/training/finetune-lora.cpp b/examples/training/finetune-lora.cpp index 5244b45f6da6..1f97b2bcde03 100644 --- a/examples/training/finetune-lora.cpp +++ b/examples/training/finetune-lora.cpp @@ -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); diff --git a/examples/training/finetune.cpp b/examples/training/finetune.cpp index 55d3c9dc9558..c19e5f82dd19 100644 --- a/examples/training/finetune.cpp +++ b/examples/training/finetune.cpp @@ -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); diff --git a/ggml/include/ggml-opt.h b/ggml/include/ggml-opt.h index 9dc27b7bf9be..c04ffba47f67 100644 --- a/ggml/include/ggml-opt.h +++ b/ggml/include/ggml-opt.h @@ -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. ## diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index f7dedfe22e3a..619b82b88e62 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -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 { @@ -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 diff --git a/ggml/src/ggml-opt.cpp b/ggml/src/ggml-opt.cpp index b709375b52d6..6379f996f4f3 100644 --- a/ggml/src/ggml-opt.cpp +++ b/ggml/src/ggml-opt.cpp @@ -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 param_grad_acc; + std::unordered_map param_m; + std::unordered_map 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 leaf_grads; std::vector grad_accs; - std::vector grad_m; - std::vector grad_v; int64_t iter = 1; int32_t opt_period = 1; @@ -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"); } @@ -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); @@ -606,32 +641,41 @@ 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); } } } @@ -639,6 +683,17 @@ static void ggml_opt_build(ggml_opt_context_t opt_ctx) { // 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()); } @@ -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); } @@ -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); @@ -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 { @@ -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); } } } @@ -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; diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 87b9e778aeaa..8f64ff6399bb 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -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: { @@ -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); @@ -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 @@ -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); diff --git a/include/llama.h b/include/llama.h index 231d00a6b5e2..f128e2b6ba7e 100644 --- a/include/llama.h +++ b/include/llama.h @@ -1714,6 +1714,14 @@ extern "C" { // per-layer residual, "l_out") instead of keeping every layer's alive: about one extra // forward, and one layer's attention scores in memory instead of all of them. bool recompute; + + // THE EXACT WALK (pure attention models): one optimizer step per training window, with + // each chunk's loss carried backward through the cached K/V of every chunk before it, so + // the step is the gradient of the whole window's loss, not of each chunk alone. false keeps + // the walk's stop-gradient at each chunk boundary. walk_horizon: how many cached positions + // before a chunk receive its gradient (0 = the whole window). + bool walk_exact; + uint32_t walk_horizon; }; LLAMA_API void llama_opt_init(struct llama_context * lctx, struct llama_model * model, struct llama_opt_params lopt_params); @@ -1747,6 +1755,20 @@ extern "C" { // batch): the footprint a run of this shape needs, from the run's own allocation preflight. LLAMA_API size_t llama_opt_graph_bytes(struct llama_context * lctx); + // Cap, in bytes, the HOST memory the exact walk keeps per training window (0 = no cap): the + // gradient accumulated on every cached K/V position of every attention layer, and a recurrent + // model's state snapshot at every chunk boundary. A window over the cap refuses the run by + // name before anything is decoded (llama_opt_failure). Set before llama_opt_epoch. + LLAMA_API void llama_opt_set_walk_host_budget(struct llama_context * lctx, size_t bytes); + + // The host bytes the exact walk's largest window so far kept (0 before one ran, or when the + // walk is not exact): what a run of this shape needs, from the run's own arithmetic. + LLAMA_API size_t llama_opt_walk_host_bytes(struct llama_context * lctx); + + // The gradient horizon the exact walk's last window trained at, in positions (0 = the whole + // window): the requested one, or smaller where a chunk's graph did not fit the device. + LLAMA_API uint32_t llama_opt_walk_horizon(struct llama_context * lctx); + LLAMA_API void llama_opt_epoch( struct llama_context * lctx, ggml_opt_dataset_t dataset, diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 80eaeb1336df..2ffe62e3d52b 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3335,6 +3335,9 @@ void llama_context::opt_init(struct llama_model * model, struct llama_opt_params // The training window lives on this context, not the shared model (a serving context may // be running on the same weights). opt_n_ctx_train = lopt_params.n_ctx_train > 0 ? lopt_params.n_ctx_train : n_ctx(); + opt_walk_exact = lopt_params.walk_exact; + opt_walk_horizon_req = lopt_params.walk_horizon; + opt_walk_horizon = lopt_params.walk_horizon; const uint32_t n_batch = std::min(this->n_batch(), opt_n_ctx_train); const uint32_t n_ubatch = std::min(this->n_ubatch(), n_batch); GGML_ASSERT(opt_n_ctx_train % n_batch == 0); @@ -3495,6 +3498,25 @@ void llama_context::opt_epoch_iter( return ok; }; + // THE EXACT WALK's state for the chunk being trained (null in the plain walk): its place in + // the window's one optimizer period, and the gradient accumulated on the window's cached K/V + struct exact_chunk { + uint32_t c0; + uint32_t n; + uint32_t n_labels; + uint32_t grad_from; // its GRAD leaves cover cached positions [grad_from, c0) + bool surrogate; // a later chunk's gradient sits on its own K/V + bool needed; + bool period_end; // the window's optimizer step + float loss_scale; // its labelled positions over the window's + }; + const exact_chunk * xc = nullptr; + // per layer, per position: dL/d(cached K) and dL/d(cached V), [position][n_embd_*_gqa] + std::vector> walk_gk, walk_gv; + // per recurrent state tensor: dL/d(the state the chunk just before leaves), from the chunk + // trained last (each chunk's exit state feeds exactly one chunk, the next) + std::unordered_map> walk_gstate; + // one training chunk [pos_ctx, pos_ctx + n_tokens): forward and backward over it alone, // attending to the cached context before it auto train_chunk = [&](uint32_t pos_ctx, uint32_t n_chunk) -> bool { @@ -3504,7 +3526,10 @@ void llama_context::opt_epoch_iter( batch.pos [pos_batch] = pos_ctx + pos_batch; batch.n_seq_id[pos_batch] = 1; batch.seq_id [pos_batch][0] = 0; - batch.logits [pos_batch] = labels_sparse[pos_ctx + pos_batch] >= 0; + // an exact walk's context chunk has no label but still trains (its K/V carry later + // chunks' gradient): one output row, weighted zero, so the graph has its loss + batch.logits [pos_batch] = labels_sparse[pos_ctx + pos_batch] >= 0 + || (xc != nullptr && xc->n_labels == 0 && pos_batch + 1 == n_chunk); } // output_all = false: the batch's own logits flags decide which rows exist @@ -3577,6 +3602,9 @@ void llama_context::opt_epoch_iter( ctx_compute_opt = ggml_init(params); } ggml_opt_prepare_alloc(opt_ctx, ctx_compute_opt, gf, res->get_inp_tokens(), res->get_logits()); + if (xc != nullptr) { + ggml_opt_set_next_step(opt_ctx, xc->period_end, xc->loss_scale, res->t_walk_surrogate); + } if (!ggml_opt_alloc(opt_ctx, train)) { const char * why = ggml_opt_refusal(opt_ctx); opt_failure = why; @@ -3603,13 +3631,44 @@ void llama_context::opt_epoch_iter( continue; } const uint32_t ilabel = pos_ctx + pos_batch + pos_ubatch; - GGML_ASSERT(labels_sparse[ilabel] >= 0 && labels_sparse[ilabel] < labels->ne[0]); - ggml_backend_tensor_set(labels, &onef, (row*labels->ne[0] + labels_sparse[ilabel])*sizeof(float), sizeof(float)); + // an exact walk's context chunk: its one row has no label (weighted zero) + GGML_ASSERT((xc != nullptr && labels_sparse[ilabel] < 0) || (labels_sparse[ilabel] >= 0 && labels_sparse[ilabel] < labels->ne[0])); + if (labels_sparse[ilabel] >= 0) { + ggml_backend_tensor_set(labels, &onef, (row*labels->ne[0] + labels_sparse[ilabel])*sizeof(float), sizeof(float)); + } ++row; } GGML_ASSERT(row == n_outputs); } - ggml_opt_eval(opt_ctx, result); + if (xc != nullptr) { + // the GRAD leaves start at zero; the surrogate inputs are what later chunks + // accumulated on this chunk's own K/V + for (const auto & io : res->t_walk) { + if (io.dk) { + ggml_backend_tensor_memset(io.dk, 0, 0, ggml_nbytes(io.dk)); + ggml_backend_tensor_memset(io.dv, 0, 0, ggml_nbytes(io.dv)); + } + if (io.gk) { + const size_t kd = io.gk->ne[0]*io.gk->ne[1]; + const size_t vd = io.gv->ne[0]*io.gv->ne[1]; + ggml_backend_tensor_set(io.gk, walk_gk[io.il].data() + (size_t) pos_ctx*kd, 0, ggml_nbytes(io.gk)); + ggml_backend_tensor_set(io.gv, walk_gv[io.il].data() + (size_t) pos_ctx*vd, 0, ggml_nbytes(io.gv)); + } + } + } + if (xc != nullptr) { + for (const auto & ws : res->t_walk_state) { + if (ws.ds) { + ggml_backend_tensor_memset(ws.ds, 0, 0, ggml_nbytes(ws.ds)); + } + if (ws.gs) { + const auto it = walk_gstate.find(ws.cache); + GGML_ASSERT(it != walk_gstate.end() && it->second.size() == (size_t) ggml_nelements(ws.gs)); + ggml_backend_tensor_set(ws.gs, it->second.data(), 0, ggml_nbytes(ws.gs)); + } + } + } + ggml_opt_eval(opt_ctx, xc != nullptr && xc->n_labels == 0 ? nullptr : result); if (const char * why = ggml_opt_refusal(opt_ctx); why[0] != '\0') { // the backend failed the graph (ggml_opt_eval): the run fails, as a refused graph does opt_failure = why; @@ -3619,7 +3678,40 @@ void llama_context::opt_epoch_iter( opt_stop_requested.store(true); return false; } - if (callback) { + if (xc != nullptr) { + // dL/d(cached K/V) of this chunk's prefix joins what later chunks put there + std::vector g; + for (const auto & io : res->t_walk) { + if (!io.dk) { + continue; + } + for (int kv = 0; kv < 2; ++kv) { + ggml_tensor * leaf = kv == 0 ? io.dk : io.dv; + ggml_tensor * grad = ggml_opt_leaf_grad(opt_ctx, leaf); + GGML_ASSERT(grad != nullptr && grad->type == GGML_TYPE_F32 && ggml_is_contiguous(grad)); + g.resize(ggml_nelements(grad)); + ggml_backend_tensor_get(grad, g.data(), 0, ggml_nbytes(grad)); + float * acc = (kv == 0 ? walk_gk : walk_gv)[io.il].data() + (size_t) io.grad_from*leaf->ne[0]*leaf->ne[1]; + for (size_t e = 0; e < g.size(); ++e) { + acc[e] += g[e]; + } + } + } + } + if (xc != nullptr) { + // dL/d(entry state): what the previous chunk's exit state takes in its surrogate + for (const auto & ws : res->t_walk_state) { + if (!ws.ds) { + continue; + } + ggml_tensor * grad = ggml_opt_leaf_grad(opt_ctx, ws.ds); + GGML_ASSERT(grad != nullptr && grad->type == GGML_TYPE_F32 && ggml_is_contiguous(grad)); + auto & g = walk_gstate[ws.cache]; + g.resize(ggml_nelements(grad)); + ggml_backend_tensor_get(grad, g.data(), 0, ggml_nbytes(grad)); + } + } + if (callback && !(xc != nullptr && xc->n_labels == 0)) { callback(train, opt_ctx, dataset, result, idata_in_loop + (pos_ctx + pos_batch)/n_batch + 1, ndata_in_loop, t_loop_start); } ggml_free(ctx_compute_opt); @@ -3639,12 +3731,243 @@ void llama_context::opt_epoch_iter( } else { recr = dynamic_cast(memory.get()); } - if (recr != nullptr && cparams.n_seq_max < 2) { + // the plain walk snapshots a recurrent state into a scratch sequence; the exact walk keeps host snapshots + if (recr != nullptr && cparams.n_seq_max < 2 && !(train && opt_walk_exact)) { LLAMA_LOG_ERROR("%s: a recurrent model trains with a scratch sequence for its state snapshot: create the training context with n_seq_max >= 2\n", __func__); opt_stop_requested.store(true); return; } + if (train && opt_walk_exact) { + // THE EXACT WALK. The plain walk below steps per chunk and stops the gradient at every + // chunk boundary: a reply learns only through its own chunk's K/V, never through what + // an earlier chunk computed for it. Here the window is decoded once under the adapter + // as it is, then trained in REVERSE: each chunk's graph attends to the cache, with its + // last walk_horizon prefix positions behind a GRAD leaf, and adds the surrogate + // + on its own K/V, where G is what every later chunk put there. Its + // backward then carries that gradient through itself into the adapter and into its own + // prefix's G. ONE optimizer step per window, on the gradient of the window's loss + // (exact within the horizon). Memory stays chunk x window; the cost is one decode plus + // one forward and backward per chunk, context chunks included. + // THE INVARIANT THAT MAKES THE REVERSE PASS VALID (Fable): the adapter does not change + // between the forward decode and the last chunk's backward. One optimizer step per window, + // taken by the chunk the reverse pass trains last, is what keeps each chunk's recomputed + // K/V (and recurrent state) equal to the cached values its successors attended to; a step + // inside the window would make every surrogate a gradient of a different function. + // + // Each window starts from the REQUESTED horizon: a halving that fit one window's device + // budget must not shrink every later window and epoch, which may be shorter (Fable). + opt_walk_horizon = opt_walk_horizon_req; + // A recurrent model's state is carried the same way, one chunk at a time: the state + // entering a chunk is a GRAD leaf, the state it leaves meets the next chunk's gradient + // (build_rs, build_walk_state_exit). The reverse pass restores each chunk's entry state + // from a host snapshot taken at its boundary in the forward pass. + std::vector chunks; + uint32_t n_labels_window = 0; + for (uint32_t p = 0; p <= (uint32_t) last_label; ) { + const bool labelled = labels_sparse[p] >= 0; + uint32_t end = p; + while (end <= (uint32_t) last_label && (labels_sparse[end] >= 0) == labelled) { + ++end; + } + for (uint32_t c0 = p; c0 < end; c0 += n_batch) { + const uint32_t c1 = std::min(c0 + n_batch, end); + chunks.push_back({ c0, c1 - c0, labelled ? c1 - c0 : 0, 0, false, false, false, 0.0f }); + n_labels_window += labelled ? c1 - c0 : 0; + } + p = end; + } + // which chunks train, and how far back each one's gradient reaches + uint32_t reach = UINT32_MAX; // the lowest position a later trained chunk's gradient reaches + int64_t first = -1; // the last chunk the reverse pass trains: the step + for (int64_t j = (int64_t) chunks.size() - 1; j >= 0; --j) { + exact_chunk & c = chunks[j]; + c.surrogate = reach < c.c0 + c.n; + // a recurrent state reaches every later chunk: a chunk before a trained one trains + c.needed = c.n_labels > 0 || c.surrogate || (recr != nullptr && first >= 0); + if (!c.needed) { + continue; + } + c.grad_from = opt_walk_horizon == 0 || c.c0 < opt_walk_horizon ? 0 : c.c0 - opt_walk_horizon; + c.loss_scale = (float) c.n_labels / (float) n_labels_window; + reach = std::min(reach, c.grad_from); + first = j; + } + GGML_ASSERT(first >= 0); + chunks[first].period_end = true; + + // THE HOST BUDGET (Fable): the accumulators are layers x positions x (k + v) floats, and a + // recurrent state is checkpointed every `stride` chunks (one sequence's state rows each; + // the first chunk's entry is the zero state and costs nothing). Both are known before + // anything runs. Under a budget the stride is the smallest that fits: the reverse pass + // rebuilds a chunk's entry state by decoding forward from its checkpoint, so a longer + // stride trades host memory for decode, never for exactness. + size_t stride = 1; + { + size_t kv_bytes = 0; + for (uint32_t il = 0; il < model.hparams.n_layer(); ++il) { + if (!model.hparams.is_recr(il)) { + kv_bytes += (size_t) (last_label + 1)*(model.hparams.n_embd_k_gqa(il) + model.hparams.n_embd_v_gqa(il))*sizeof(float); + } + } + size_t state = 0; + if (recr != nullptr) { + for (size_t il = 0; il < recr->r_l.size(); ++il) { + state += recr->r_l[il] ? recr->r_l[il]->nb[1] : 0; + state += recr->s_l[il] ? recr->s_l[il]->nb[1] : 0; + } + } + auto bytes_at = [&](size_t s) { return kv_bytes + ((chunks.size() + s - 1)/s - 1)*state; }; + if (opt_walk_host_budget > 0) { + while (stride < chunks.size() && bytes_at(stride) > opt_walk_host_budget) { + ++stride; + } + } + const size_t bytes = bytes_at(stride); + opt_walk_host_bytes = std::max(opt_walk_host_bytes, bytes); + if (opt_walk_host_budget > 0 && bytes > opt_walk_host_budget) { + char why[512]; + snprintf(why, sizeof(why), + "the exact walk needs %.1f MiB of host memory for this window (%lld positions, the K/V gradient alone), over its budget of %.1f MiB: " + "the K/V gradient grows with the window; neither the horizon nor the chunk size changes it", + bytes/1048576.0, (long long) (last_label + 1), opt_walk_host_budget/1048576.0); + opt_failure = why; + LLAMA_LOG_ERROR("%s: %s\n", __func__, opt_failure.c_str()); + opt_alloc_failed.store(true); + opt_stop_requested.store(true); + return; + } + if (stride > 1) { + LLAMA_LOG_INFO("%s: the exact walk checkpoints the recurrent state every %zu chunks to fit %.1f MiB of host memory (%.1f MiB)\n", + __func__, stride, opt_walk_host_budget/1048576.0, bytes/1048576.0); + } + } + walk_gk.assign(model.hparams.n_layer(), {}); + walk_gv.assign(model.hparams.n_layer(), {}); + for (uint32_t il = 0; il < model.hparams.n_layer(); ++il) { + if (model.hparams.is_recr(il)) { + continue; // a recurrent layer has no K/V: its state carries the chain + } + walk_gk[il].assign((size_t) (last_label + 1)*model.hparams.n_embd_k_gqa(il), 0.0f); + walk_gv[il].assign((size_t) (last_label + 1)*model.hparams.n_embd_v_gqa(il), 0.0f); + } + + // the window as the adapter now reads it, chunk by chunk: a recurrent model's state is + // snapshotted at each chunk's start (empty before the first: it starts from zero) + struct state_snapshot { + bool empty = true; + std::vector> r, s; // per layer + }; + std::vector snaps(recr != nullptr ? chunks.size() : 0); + auto state_rows = [&](bool write, state_snapshot & snap) { + const int32_t cell = recr->cells[0].tail; // the cell holding sequence 0's state + if (!write) { + snap.empty = cell < 0; + snap.r.assign(recr->r_l.size(), {}); + snap.s.assign(recr->s_l.size(), {}); + } + if (cell < 0) { + return; + } + for (size_t il = 0; il < recr->r_l.size(); ++il) { + for (int kind = 0; kind < 2; ++kind) { + ggml_tensor * t = kind == 0 ? recr->r_l[il] : recr->s_l[il]; + if (t == nullptr) { + continue; + } + GGML_ASSERT(t->type == GGML_TYPE_F32); + auto & row = kind == 0 ? snap.r[il] : snap.s[il]; + const size_t bytes = t->nb[1]; + if (write) { + ggml_backend_tensor_set(t, row.data(), (size_t) cell*bytes, bytes); + } else { + row.resize(t->ne[0]); + ggml_backend_tensor_get(t, row.data(), (size_t) cell*bytes, bytes); + } + } + } + }; + for (size_t j = 0; j < chunks.size(); ++j) { + if (recr != nullptr && j % stride == 0) { + state_rows(false, snaps[j]); + } + if (!decode_span(chunks[j].c0, chunks[j].c0 + chunks[j].n)) { + return; + } + } + bool ok = true; + for (int64_t j = (int64_t) chunks.size() - 1; j >= first && ok; --j) { + const exact_chunk & c = chunks[j]; + if (!c.needed) { + continue; + } + if (!yield_point()) { + ok = false; + break; + } + // attention: pop the cache to [0, c0); recurrent: restore the state at the chunk's + // checkpoint and decode forward to its start (that decode rebuilds the attention K/V + // in between too, under the same adapter: the same values) + const size_t cp = recr != nullptr ? (size_t) j - (size_t) j % stride : (size_t) j; + const uint32_t p_pop = chunks[cp].c0; + const bool popped = recr != nullptr ? (attn == nullptr || attn->seq_rm(0, p_pop, -1)) : memory->seq_rm(0, p_pop, -1); + if (!popped) { + LLAMA_LOG_ERROR("%s: could not pop the cache to [0, %u) for the reverse pass\n", __func__, p_pop); + opt_stop_requested.store(true); + ok = false; + break; + } + if (recr != nullptr) { + if (snaps[cp].empty) { + recr->seq_rm(0, -1, -1); // the first chunk starts from a zero state + } else { + state_rows(true, snaps[cp]); + } + if (p_pop < c.c0 && !decode_span(p_pop, c.c0)) { + ok = false; + break; + } + } + cparams.walk_exact = true; + cparams.walk_surrogate = c.surrogate; + cparams.walk_state_surrogate = recr != nullptr && j < (int64_t) chunks.size() - 1 && !walk_gstate.empty(); + // THE HORIZON FITS THE DEVICE. Its GRAD leaves are the chunk graph's largest growth + // with the window (layers x horizon x (k + v) floats, twice with their gradient); the + // graph preflight refuses a chunk that does not fit before allocating anything, so a + // memory refusal halves the horizon (whole chunks) and the chunk is tried again; every + // later chunk keeps the smaller one. The forward still attends to the whole window and + // the recurrent state is carried exactly; only attention gradient past the horizon is + // dropped. A refusal at a horizon of one chunk, or for a node the device cannot run, + // stands. + exact_chunk cc = c; + for (;;) { + cc.grad_from = opt_walk_horizon == 0 || cc.c0 < opt_walk_horizon ? 0 : cc.c0 - opt_walk_horizon; + cparams.walk_grad_from = cc.grad_from; + xc = &cc; + ok = train_chunk(cc.c0, cc.n); + xc = nullptr; + const uint32_t span = cc.c0 - cc.grad_from; + if (ok || !opt_failure.empty() || !opt_alloc_failed.load() || span <= n_batch) { + break; + } + opt_walk_horizon = std::max(n_batch, (span/2)/n_batch*n_batch); + LLAMA_LOG_WARN("%s: the chunk at %u did not fit with a gradient horizon of %u positions: retrying at %u\n", + __func__, cc.c0, span, opt_walk_horizon); + opt_alloc_failed.store(false); + opt_stop_requested.store(false); + } + opt_walk_horizon_used = opt_walk_horizon; + cparams.walk_exact = false; + cparams.walk_grad_from = 0; + cparams.walk_surrogate = false; + cparams.walk_state_surrogate = false; + } + walk_gk.clear(); + walk_gv.clear(); + walk_gstate.clear(); + return; + } + uint32_t pos = 0; while (pos <= (uint32_t) last_label && !opt_stop_requested.load(std::memory_order_relaxed)) { const bool labelled = labels_sparse[pos] >= 0; @@ -4464,6 +4787,18 @@ void llama_opt_set_step_callback(struct llama_context * ctx, llama_opt_step_call ctx->opt_step_callback_data = user_data; } +void llama_opt_set_walk_host_budget(struct llama_context * ctx, size_t bytes) { + ctx->opt_walk_host_budget = bytes; +} + +size_t llama_opt_walk_host_bytes(struct llama_context * ctx) { + return ctx->opt_walk_host_bytes; +} + +uint32_t llama_opt_walk_horizon(struct llama_context * ctx) { + return ctx->opt_walk_horizon_used; +} + void llama_opt_set_memory_budget(struct llama_context * ctx, size_t bytes) { ctx->opt_memory_budget = bytes; } diff --git a/src/llama-context.h b/src/llama-context.h index 405a74275898..1cd3e75d2cee 100644 --- a/src/llama-context.h +++ b/src/llama-context.h @@ -208,6 +208,10 @@ struct llama_context { std::string opt_failure; // what training may add per GPU device (llama_opt_set_memory_budget); 0 = no cap size_t opt_memory_budget = 0; + // the exact walk's host memory per window (llama_opt_set_walk_host_budget); 0 = no cap + size_t opt_walk_host_budget = 0; + size_t opt_walk_host_bytes = 0; // the largest window's, measured by its own arithmetic + uint32_t opt_walk_horizon_used = 0; // the gradient horizon the last exact window trained at (0 = all) // the largest training graph the allocation preflight measured on a GPU device (bytes) size_t opt_graph_bytes() const; @@ -369,6 +373,9 @@ struct llama_context { // model's hparams, because several contexts share one model (a serving context and a // training context on the same resident weights) uint32_t opt_n_ctx_train = 0; + bool opt_walk_exact = false; // llama_opt_params::walk_exact + uint32_t opt_walk_horizon_req = 0; // llama_opt_params::walk_horizon, as requested + uint32_t opt_walk_horizon = 0; // this window's: starts at the request, shrinks to fit the device ggml_threadpool_t threadpool = nullptr; ggml_threadpool_t threadpool_batch = nullptr; diff --git a/src/llama-cparams.h b/src/llama-cparams.h index 85fbde9c3c9a..5df493dec5da 100644 --- a/src/llama-cparams.h +++ b/src/llama-cparams.h @@ -12,6 +12,17 @@ struct llama_cparams { // set by llama_opt_init: the graph is built for a backward pass, so attention reads K/V // from this ubatch directly rather than through the KV cache (see build_attn) bool training = false; + + // THE EXACT WALK (llama_context::opt_epoch_iter, walk_exact): a chunk's training graph in + // the reverse pass. walk_grad_from: the first cached prefix position whose K/V carry a + // gradient (positions [walk_grad_from, n_past) attend through a GRAD leaf; [0, + // walk_grad_from) stay constants). walk_surrogate: later chunks accumulated a gradient on + // this chunk's own K/V, so its graph adds the surrogate term that carries it in. + bool walk_exact = false; + uint32_t walk_grad_from = 0; + bool walk_surrogate = false; + // a later chunk's gradient sits on the recurrent state this chunk leaves + bool walk_state_surrogate = false; uint32_t n_ctx_seq; // context for a single sequence uint32_t n_batch; uint32_t n_ubatch; diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index 7fef5d164082..698ec24bdaa7 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -1340,6 +1340,10 @@ void llm_graph_result::reset() { t_layer_inp.resize(LLAMA_MAX_LAYERS + 1); std::fill(t_layer_inp.begin(), t_layer_inp.end(), nullptr); + t_walk.clear(); + t_walk_state.clear(); + t_walk_surrogate = nullptr; + t_sampled.clear(); t_sampled_probs.clear(); t_sampled_logits.clear(); @@ -2849,29 +2853,88 @@ ggml_tensor * llm_graph_context::build_attn( const int64_t n_past = ubatch.pos[0]; k = k_cur; v = v_cur; - if (n_past > 0) { - ggml_tensor * k_all = mctx_cur->get_k(ctx0, il); - ggml_tensor * k_prev = ggml_view_4d(ctx0, k_all, k_all->ne[0], k_all->ne[1], n_past, 1, - k_all->nb[1], k_all->nb[2], k_all->nb[3], 0); + // the cached prefix positions [p0, p1) as this chunk's K / V type + ggml_tensor * k_all = n_past > 0 ? mctx_cur->get_k(ctx0, il) : nullptr; + ggml_tensor * v_all = n_past > 0 ? mctx_cur->get_v(ctx0, il) : nullptr; + auto k_span = [&](int64_t p0, int64_t p1) { + ggml_tensor * t = ggml_view_4d(ctx0, k_all, k_all->ne[0], k_all->ne[1], p1 - p0, 1, + k_all->nb[1], k_all->nb[2], k_all->nb[3], p0*k_all->nb[2]); // the cache may hold the constants at a smaller type than the chunk's own K/V - if (k_prev->type != k_cur->type) { - k_prev = ggml_cast(ctx0, k_prev, k_cur->type); - } - k = ggml_concat(ctx0, k_prev, k_cur, 2); - + return t->type != k_cur->type ? ggml_cast(ctx0, t, k_cur->type) : t; + }; + auto v_span = [&](int64_t p0, int64_t p1) { // a cache without flash attention stores V transposed: [n_kv, n_head_kv, n_embd_head_v] - ggml_tensor * v_all = mctx_cur->get_v(ctx0, il); - ggml_tensor * v_prev = v_all->nb[1] > v_all->nb[2] + ggml_tensor * t = v_all->nb[1] > v_all->nb[2] ? ggml_cont(ctx0, ggml_permute(ctx0, - ggml_view_4d(ctx0, v_all, n_past, v_all->ne[1], v_all->ne[2], 1, - v_all->nb[1], v_all->nb[2], v_all->nb[3], 0), + ggml_view_4d(ctx0, v_all, p1 - p0, v_all->ne[1], v_all->ne[2], 1, + v_all->nb[1], v_all->nb[2], v_all->nb[3], p0*v_all->nb[0]), 2, 1, 0, 3)) - : ggml_view_4d(ctx0, v_all, v_all->ne[0], v_all->ne[1], n_past, 1, - v_all->nb[1], v_all->nb[2], v_all->nb[3], 0); - if (v_prev->type != v_cur->type) { - v_prev = ggml_cast(ctx0, v_prev, v_cur->type); + : ggml_view_4d(ctx0, v_all, v_all->ne[0], v_all->ne[1], p1 - p0, 1, + v_all->nb[1], v_all->nb[2], v_all->nb[3], p0*v_all->nb[2]); + return t->type != v_cur->type ? ggml_cast(ctx0, t, v_cur->type) : t; + }; + if (cparams.walk_exact) { + // THE EXACT WALK's chunk graph. The prefix positions [grad_from, n_past) are the + // cached values PLUS a zero GRAD leaf, so the backward yields dL/d(cached K/V) for + // them without copying the cache; earlier positions stay constants. The chunk's own + // K/V meet the gradient later chunks put on them through the surrogate + // + : its gradient on k_cur / v_cur IS gk / gv, which the + // backward carries through this chunk into the adapter and into its own prefix. + // At a quantized or f16 cache the prefix's gradient is taken at the cache's values and + // applied to the chunk's own full-precision K/V: straight-through. Measured + // (test-walk-exact, q8_0 and f16 caches): cosine 0.998 against one graph, the same as + // at an F32 cache, so training reads the numbers serving runs on (Joel: align the bit + // depth to inference). + const int64_t g0 = std::min(cparams.walk_grad_from, n_past); + llm_graph_result::walk_layer io = { il, (uint32_t) g0, (uint32_t) n_past }; + // [const prefix][gradient prefix][this chunk]: the order the cells hold them in + std::vector k_parts, v_parts; + if (g0 > 0) { + k_parts.push_back(k_span(0, g0)); + v_parts.push_back(v_span(0, g0)); + } + if (n_past > g0) { + ggml_tensor * kc = ggml_cont(ctx0, ggml_cast(ctx0, k_span(g0, n_past), GGML_TYPE_F32)); + ggml_tensor * vc = ggml_cont(ctx0, ggml_cast(ctx0, v_span(g0, n_past), GGML_TYPE_F32)); + io.dk = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, kc->ne[0], kc->ne[1], kc->ne[2], 1); + io.dv = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, vc->ne[0], vc->ne[1], vc->ne[2], 1); + ggml_format_name(io.dk, "walk_dk-%d", il); + ggml_format_name(io.dv, "walk_dv-%d", il); + ggml_set_input(io.dk); + ggml_set_input(io.dv); + ggml_set_grad(io.dk); + ggml_set_grad(io.dv); + ggml_tensor * kg = ggml_add(ctx0, kc, io.dk); + ggml_tensor * vg = ggml_add(ctx0, vc, io.dv); + k_parts.push_back(kg->type != k_cur->type ? ggml_cast(ctx0, kg, k_cur->type) : kg); + v_parts.push_back(vg->type != v_cur->type ? ggml_cast(ctx0, vg, v_cur->type) : vg); + } + k_parts.push_back(k_cur); + v_parts.push_back(v_cur); + k = k_parts[0]; + v = v_parts[0]; + for (size_t i = 1; i < k_parts.size(); ++i) { + k = ggml_concat(ctx0, k, k_parts[i], 2); + v = ggml_concat(ctx0, v, v_parts[i], 2); + } + if (cparams.walk_surrogate) { + io.gk = ggml_new_tensor_3d(ctx0, GGML_TYPE_F32, k_cur->ne[0], k_cur->ne[1], k_cur->ne[2]); + io.gv = ggml_new_tensor_3d(ctx0, GGML_TYPE_F32, v_cur->ne[0], v_cur->ne[1], v_cur->ne[2]); + ggml_format_name(io.gk, "walk_gk-%d", il); + ggml_format_name(io.gv, "walk_gv-%d", il); + ggml_set_input(io.gk); + ggml_set_input(io.gv); + ggml_tensor * k32 = k_cur->type == GGML_TYPE_F32 ? k_cur : ggml_cast(ctx0, k_cur, GGML_TYPE_F32); + ggml_tensor * v32 = v_cur->type == GGML_TYPE_F32 ? v_cur : ggml_cast(ctx0, v_cur, GGML_TYPE_F32); + ggml_tensor * term = ggml_add(ctx0, + ggml_sum(ctx0, ggml_mul(ctx0, k32, io.gk)), + ggml_sum(ctx0, ggml_mul(ctx0, v32, io.gv))); + res->t_walk_surrogate = res->t_walk_surrogate ? ggml_add(ctx0, res->t_walk_surrogate, term) : term; } - v = ggml_concat(ctx0, v_prev, v_cur, 2); + res->t_walk.push_back(io); + } else if (n_past > 0) { + k = ggml_concat(ctx0, k_span(0, n_past), k_cur, 2); + v = ggml_concat(ctx0, v_span(0, n_past), v_cur, 2); } // the cache's cells hold the prefix at [0, n_past) and this ubatch right after it // (contiguous: soft_max_ext requires it, and a column slice of the padded mask is not) @@ -3494,6 +3557,17 @@ ggml_tensor * llm_graph_context::build_rs( // NOTE: assuming the copy destinations are ALL contained between rs_head and rs_head + n_rs // {state_size, rs_size} -> {state_size, n_seqs} ggml_tensor * output_states = get_state_rows(ctx0, states, state_copy_main); + if (cparams.training && cparams.walk_exact) { + // THE EXACT WALK: the state entering the chunk is the cached one PLUS a zero GRAD leaf, + // so the backward yields dL/d(entry state): the gradient the previous chunk's exit takes + GGML_ASSERT(output_states->type == GGML_TYPE_F32 && "the exact walk carries an F32 recurrent state"); + ggml_tensor * ds = ggml_new_tensor(ctx0, GGML_TYPE_F32, GGML_MAX_DIMS, output_states->ne); + ggml_format_name(ds, "walk_ds-%s", s->name); + ggml_set_input(ds); + ggml_set_grad(ds); + output_states = ggml_add(ctx0, output_states, ds); + res->walk_state_of(s).ds = ds; + } ggml_build_forward_expand(gf, output_states); // copy extra states which won't be changed further (between n_seqs and n_rs) @@ -3549,6 +3623,22 @@ ggml_tensor * llm_graph_context::build_rs( get_state_rows); } +void llm_graph_context::build_walk_state_exit(ggml_tensor * cache, ggml_tensor * exit) const { + if (!(cparams.training && cparams.walk_exact && cparams.walk_state_surrogate)) { + return; + } + ggml_tensor * flat = ggml_reshape_1d(ctx0, ggml_cont(ctx0, exit), ggml_nelements(exit)); + if (flat->type != GGML_TYPE_F32) { + flat = ggml_cast(ctx0, flat, GGML_TYPE_F32); + } + ggml_tensor * gs = ggml_new_tensor_1d(ctx0, GGML_TYPE_F32, ggml_nelements(exit)); + ggml_format_name(gs, "walk_gs-%s", cache->name); + ggml_set_input(gs); + ggml_tensor * term = ggml_sum(ctx0, ggml_mul(ctx0, flat, gs)); + res->t_walk_surrogate = res->t_walk_surrogate ? ggml_add(ctx0, res->t_walk_surrogate, term) : term; + res->walk_state_of(cache).gs = gs; +} + ggml_tensor * llm_graph_context::build_rwkv_token_shift_load( llm_graph_input_rs * inp, const llama_ubatch & ubatch, diff --git a/src/llama-graph.h b/src/llama-graph.h index b388e028cb53..5f924306f7f6 100644 --- a/src/llama-graph.h +++ b/src/llama-graph.h @@ -938,6 +938,37 @@ class llm_graph_result { std::vector t_layer_inp; + // THE EXACT WALK: per attention layer of a reverse-pass chunk graph, the tensors the walk + // fills before the step and reads after it (llama_context::opt_epoch_iter) + struct walk_layer { + int il; + uint32_t grad_from; // the GRAD leaves cover cached positions [grad_from, n_past) + uint32_t n_past; + ggml_tensor * dk = nullptr; // GRAD leaves, zero-filled: their gradient is dL/d(cached K/V) + ggml_tensor * dv = nullptr; // [n_embd_head, n_head_kv, n_past - grad_from], V untransposed + ggml_tensor * gk = nullptr; // inputs: the gradient later chunks accumulated on this chunk's + ggml_tensor * gv = nullptr; // own K/V, shaped like k_cur / v_cur (null without a surrogate) + }; + std::vector t_walk; + // the recurrent states of a reverse-pass chunk graph, one per state tensor of the memory + // (a layer's conv state r_l and its recurrent state s_l): the chain runs through them too + struct walk_state { + ggml_tensor * cache = nullptr; // the memory's state tensor this entry belongs to + ggml_tensor * ds = nullptr; // GRAD leaf, zero-filled, on the state ENTERING the chunk + ggml_tensor * gs = nullptr; // input: the gradient later chunks put on the state it LEAVES + }; + std::vector t_walk_state; + walk_state & walk_state_of(ggml_tensor * cache) { + for (auto & w : t_walk_state) { + if (w.cache == cache) { + return w; + } + } + t_walk_state.push_back({ cache, nullptr, nullptr }); + return t_walk_state.back(); + } + ggml_tensor * t_walk_surrogate = nullptr; // sum over layers of + + std::vector t_sampled; std::vector t_sampled_probs; std::vector t_sampled_logits; @@ -1326,6 +1357,10 @@ struct llm_graph_context { int32_t n_seqs, const llm_graph_get_rows_fn & get_state_rows = ggml_get_rows) const; + // THE EXACT WALK: the state a chunk leaves (`exit`, what is written back to `cache`) meets the + // gradient later chunks put on it, through the surrogate . A no-op outside the walk. + void build_walk_state_exit(ggml_tensor * cache, ggml_tensor * exit) const; + ggml_tensor * build_rwkv_token_shift_load( llm_graph_input_rs * inp, const llama_ubatch & ubatch, diff --git a/src/models/delta-net-base.cpp b/src/models/delta-net-base.cpp index ad6612647736..e0a6368bd2ed 100644 --- a/src/models/delta-net-base.cpp +++ b/src/models/delta-net-base.cpp @@ -494,7 +494,9 @@ ggml_tensor * llm_build_delta_net_base::build_conv_state( cb(conv_state_update, "conv_state_update", il); ggml_build_forward_expand(gf, ggml_cpy(ctx0, conv_state_last, conv_state_update)); + build_walk_state_exit(conv_states_all, conv_state_last); } else { + GGML_ASSERT(!cparams.walk_exact && "the exact walk keeps one recurrent state per sequence (n_rs_seq = 0)"); // [TAG_RECURRENT_ROLLBACK_SPLITS] // this logic assumes that the last (n_rs_seq + 1) tokens of a sequence in a batch are inside // the same ubatch, which `split_equal()` guarantees via its n_keep_tail argument @@ -556,6 +558,7 @@ ggml_tensor * llm_build_delta_net_base::build_recurrent_attn( ggml_cpy(ctx0, new_state, ggml_view_2d(ctx0, ssm_states_all, hparams.n_embd_s(), n_seqs, ssm_states_all->nb[1], kv_head * hparams.n_embd_s() * ggml_element_size(ssm_states_all)))); + build_walk_state_exit(ssm_states_all, new_state); return output; } diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index e04d5bf4cd17..96484141528c 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -150,6 +150,12 @@ endif () llama_build(test-recurrent-state-rollback.cpp) +# the exact walk trains on the window's whole gradient whatever the chunking (needs -m ) +llama_build(test-walk-exact.cpp) + +# GGML_TENSOR_FLAG_GRAD: a leaf whose gradient the backward computes, never optimized (the walk reads the cached prefix K/V gradient) +llama_build_and_test(test-grad-leaf.cpp) + # ggml-opt with per-step graphs (the llama training path) starts each optimizer period from zero accumulators llama_build_and_test(test-opt-dynamic-accum.cpp) diff --git a/tests/test-grad-leaf.cpp b/tests/test-grad-leaf.cpp new file mode 100644 index 000000000000..82d09e576e03 --- /dev/null +++ b/tests/test-grad-leaf.cpp @@ -0,0 +1,80 @@ +// what this catches: GGML_TENSOR_FLAG_GRAD, a leaf whose gradient the backward computes while +// no optimizer updates it (the training walk reads the gradient of the cached prefix K/V this +// way). The leaf must get its exact gradient beside a real parameter's, and a leaf WITHOUT the +// flag must stay a constant with no gradient. + +#include "ggml.h" +#include "ggml-alloc.h" +#include "ggml-backend.h" + +#include +#include +#include + +#define CHECK(cond) do { if (!(cond)) { fprintf(stderr, "FAILED %s:%d: %s\n", __FILE__, __LINE__, #cond); exit(1); } } while (0) + +int main() { + // by type, not a CPU-backend symbol: in a dynamically loaded backend build (CI) the CPU + // backend is its own library + ggml_backend_load_all(); + ggml_backend_t cpu = ggml_backend_init_by_type(GGML_BACKEND_DEVICE_TYPE_CPU, nullptr); + CHECK(cpu != nullptr); + + ggml_init_params params = { 64 * ggml_tensor_overhead() + ggml_graph_overhead_custom(GGML_DEFAULT_GRAPH_SIZE, true) * 2, nullptr, true }; + ggml_context * ctx = ggml_init(params); + + const int n = 4; + ggml_tensor * x = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, n); // the GRAD leaf + ggml_tensor * w = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, n); // a real parameter + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, n); // a plain constant + ggml_set_grad(x); + ggml_set_param(w); + + // loss = sum((w * x + c)^2): d/dx = 2 (w x + c) w, d/dw = 2 (w x + c) x + ggml_tensor * y = ggml_add(ctx, ggml_mul(ctx, w, x), c); + ggml_tensor * loss = ggml_sum(ctx, ggml_sqr(ctx, y)); + ggml_set_loss(loss); + + ggml_cgraph * gf = ggml_new_graph_custom(ctx, GGML_DEFAULT_GRAPH_SIZE, true); + ggml_build_forward_expand(gf, loss); + ggml_cgraph * gb = ggml_graph_dup(ctx, gf, true); + ggml_build_backward_expand(ctx, gb, nullptr); + + ggml_backend_buffer_t buf = ggml_backend_alloc_ctx_tensors(ctx, cpu); + CHECK(buf != nullptr); + for (int i = 0; i < n; ++i) { + const float xi = 0.5f + i, wi = 1.0f - 0.25f * i, ci = 2.0f; + ggml_backend_tensor_set(x, &xi, i * sizeof(float), sizeof(float)); + ggml_backend_tensor_set(w, &wi, i * sizeof(float), sizeof(float)); + ggml_backend_tensor_set(c, &ci, i * sizeof(float), sizeof(float)); + } + ggml_graph_reset(gb); // loss grad = 1, every other grad = 0 + CHECK(ggml_backend_graph_compute(cpu, gb) == GGML_STATUS_SUCCESS); + + ggml_tensor * gx = ggml_graph_get_grad(gb, x); + ggml_tensor * gw = ggml_graph_get_grad(gb, w); + CHECK(gx != nullptr && "a GRAD leaf gets a gradient"); + CHECK(gw != nullptr && "a PARAM still gets its gradient"); + CHECK(ggml_graph_get_grad(gb, c) == nullptr && "an unflagged constant gets none"); + for (int i = 0; i < n; ++i) { + const float xi = 0.5f + i, wi = 1.0f - 0.25f * i, yi = wi * xi + 2.0f; + float gxi, gwi; + ggml_backend_tensor_get(gx, &gxi, i * sizeof(float), sizeof(float)); + ggml_backend_tensor_get(gw, &gwi, i * sizeof(float), sizeof(float)); + CHECK(std::fabs(gxi - 2.0f * yi * wi) < 1e-5f); + CHECK(std::fabs(gwi - 2.0f * yi * xi) < 1e-5f); + } + CHECK(!(x->flags & GGML_TENSOR_FLAG_PARAM) && "the GRAD leaf is not a parameter"); + // regression: GRAD first shared COMPUTE's bit, so every computed node read as a GRAD leaf and + // the backward asked for gradients of the quantized weights (GGML_ASSERT in the 1.5B walk) + for (int i = 0; i < ggml_graph_n_nodes(gb); ++i) { + ggml_tensor * node = ggml_graph_node(gb, i); + CHECK((node == x || !(node->flags & GGML_TENSOR_FLAG_GRAD)) && "only the flagged leaf is a GRAD leaf"); + } + + ggml_backend_buffer_free(buf); + ggml_free(ctx); + ggml_backend_free(cpu); + printf("test-grad-leaf: OK\n"); + return 0; +} diff --git a/tests/test-opt-dynamic-accum.cpp b/tests/test-opt-dynamic-accum.cpp index 03f857d27596..ff5445fda605 100644 --- a/tests/test-opt-dynamic-accum.cpp +++ b/tests/test-opt-dynamic-accum.cpp @@ -86,10 +86,106 @@ static void run(int32_t opt_period) { ggml_backend_free(cpu); } +// what this catches: ggml_opt_set_next_step, the training walk's one-step-per-window period. +// Three graphs on loss = sum(w * x): the first two only accumulate, at weights 0.5 and 0.25, +// the second with an extra term 3 * w * x; the third (weight 1) takes the step. So the step +// is SGD on (0.5 + 0.25 + 3 + 1) * x. Each graph's x is a GRAD leaf, whose gradient (the +// weight times w, plus 3 w in the second) must read back after its eval; a following period +// starts from zero. +static void run_manual() { + // by type, not ggml_backend_cpu_init: in a dynamically loaded backend build (CI) the CPU + // backend is its own library + ggml_backend_t cpu = ggml_backend_init_by_type(GGML_BACKEND_DEVICE_TYPE_CPU, nullptr); + GGML_ASSERT(cpu != nullptr); + ggml_backend_t backends[] = { cpu }; + ggml_backend_sched_t sched = ggml_backend_sched_new(backends, nullptr, 1, GGML_DEFAULT_GRAPH_SIZE, false, true); + + ggml_init_params sp = { 8 * ggml_tensor_overhead(), nullptr, true }; + ggml_context * ctx_static = ggml_init(sp); + ggml_tensor * w = ggml_new_tensor_1d(ctx_static, GGML_TYPE_F32, 1); + ggml_set_param(w); + ggml_backend_buffer_t buf = ggml_backend_alloc_ctx_tensors(ctx_static, cpu); + float w_now = 1.0f; + ggml_backend_tensor_set(w, &w_now, 0, sizeof(float)); + + ggml_opt_params params = ggml_opt_default_params(sched, GGML_OPT_LOSS_TYPE_CROSS_ENTROPY); + params.optimizer = GGML_OPT_OPTIMIZER_TYPE_SGD; + params.get_opt_pars = sgd_pars; + ggml_opt_context_t opt_ctx = ggml_opt_init(params); + + // cross-entropy over 2 classes of logits [w x, 0] with label class 0: + // dCE/dlogit0 = softmax0 - 1, so d/dw = (s0 - 1) x and d/dx = (s0 - 1) w + const float scales[3] = { 0.5f, 0.25f, 1.0f }; + for (int period = 0; period < 2; ++period) { + float expected_grad_w = 0.0f; + for (int k = 0; k < 3; ++k) { + ggml_init_params cp = { GGML_DEFAULT_GRAPH_SIZE * ggml_tensor_overhead() + 4 * ggml_graph_overhead_custom(GGML_DEFAULT_GRAPH_SIZE, true), nullptr, true }; + ggml_context * ctx_compute = ggml_init(cp); + ggml_tensor * x = ggml_new_tensor_1d(ctx_compute, GGML_TYPE_F32, 1); + ggml_set_input(x); + ggml_set_grad(x); + ggml_tensor * z = ggml_new_tensor_2d(ctx_compute, GGML_TYPE_F32, 1, 1); + ggml_set_input(z); + ggml_tensor * wx = ggml_reshape_2d(ctx_compute, ggml_mul(ctx_compute, w, x), 1, 1); + ggml_tensor * logits = ggml_concat(ctx_compute, wx, z, 0); // [2, 1] + ggml_tensor * extra = k == 1 ? ggml_scale(ctx_compute, ggml_sum(ctx_compute, ggml_mul(ctx_compute, w, x)), 3.0f) : nullptr; + ggml_cgraph * gf = ggml_new_graph_custom(ctx_compute, GGML_DEFAULT_GRAPH_SIZE, true); + ggml_build_forward_expand(gf, logits); + if (extra) { + ggml_build_forward_expand(gf, extra); + } + ggml_opt_prepare_alloc(opt_ctx, ctx_compute, gf, x, logits); + ggml_opt_set_next_step(opt_ctx, /*period_end =*/ k == 2, scales[k], extra); + CHECK(ggml_opt_alloc(opt_ctx, true)); + ggml_backend_tensor_set(x, &X, 0, sizeof(float)); + const float zero = 0.0f, one = 1.0f; + ggml_backend_tensor_set(z, &zero, 0, sizeof(float)); + ggml_tensor * labels = ggml_opt_labels(opt_ctx); + ggml_backend_tensor_set(labels, &one, 0, sizeof(float)); + ggml_backend_tensor_set(labels, &zero, sizeof(float), sizeof(float)); + ggml_opt_eval(opt_ctx, nullptr); + + const float s0 = 1.0f / (1.0f + std::exp(-w_now * X)); + const float dlogit = scales[k] * (s0 - 1.0f); + const float dx = dlogit * w_now + (k == 1 ? 3.0f * w_now : 0.0f); + expected_grad_w += dlogit * X + (k == 1 ? 3.0f * X : 0.0f); + ggml_tensor * gx = ggml_opt_leaf_grad(opt_ctx, x); + CHECK(gx != nullptr && "a GRAD leaf's gradient reads back after the eval"); + float gx_val; + ggml_backend_tensor_get(gx, &gx_val, 0, sizeof(float)); + if (std::fabs(gx_val - dx) > 1e-5f) { + fprintf(stderr, "period %d graph %d: dL/dx %f, expected %f\n", period, k, gx_val, dx); + exit(1); + } + ggml_free(ctx_compute); + + float w_after; + ggml_backend_tensor_get(w, &w_after, 0, sizeof(float)); + if (k < 2) { + CHECK(std::fabs(w_after - w_now) < 1e-7f && "no step before the period ends"); + } else { + const float expected = w_now - LR * expected_grad_w; + if (std::fabs(w_after - expected) > 1e-5f) { + fprintf(stderr, "period %d: w moved to %f, expected %f (from %f)\n", period, w_after, expected, w_now); + exit(1); + } + w_now = w_after; + } + } + } + + ggml_opt_free(opt_ctx); + ggml_backend_buffer_free(buf); + ggml_free(ctx_static); + ggml_backend_sched_free(sched); + ggml_backend_free(cpu); +} + int main() { ggml_backend_load_all(); run(1); run(2); + run_manual(); printf("test-opt-dynamic-accum: OK\n"); return 0; } diff --git a/tests/test-walk-exact.cpp b/tests/test-walk-exact.cpp new file mode 100644 index 000000000000..958c14ae0063 --- /dev/null +++ b/tests/test-walk-exact.cpp @@ -0,0 +1,293 @@ +// what this catches: the exact walk (llama_opt_params::walk_exact) training on anything other +// than the gradient of the window's whole loss. The plain walk stops the gradient at every +// chunk boundary; the exact walk carries each chunk's gradient back through the cached K/V of +// every chunk before it, so how the window is CHUNKED must not change the step it takes. +// +// One SGD step from the same fresh adapter, read off the adapter itself: SGD moves every A and B +// by exactly -lr * gradient, so the steps compare gradients directly (Fable: compare gradients, +// not losses; a loss-change proxy was measured to be nonlinear at practical step sizes): +// 1. every position labelled: a chunk of the whole window is ONE graph, the true gradient; the +// exact walk in 4 chunks must match it, and the plain walk in 4 chunks must not +// 2. context, then a reply (a masked window, the walk's real case): the exact walk in 4 chunks +// must take the step it takes with one chunk per run, and the plain walk must not +// +// Run on a pure-attention model and on a hybrid (attention + gated delta-net): +// test-walk-exact -m Qwen2.5-Coder-1.5B-Instruct-Q4_K_M.gguf +// test-walk-exact -m Qwen3.5-0.8B-Q8_0.gguf + +#include "arg.h" +#include "common.h" +#include "llama.h" +#include "../src/llama-adapter.h" + +#include +#include +#include +#include +#include +#include + +static const float LR = 1e-3f; // an SGD step is linear in the gradient at any lr +static const uint32_t WINDOW = 512; + +static ggml_opt_optimizer_params sgd_pars(void *) { + ggml_opt_optimizer_params p = ggml_opt_get_default_optimizer_params(nullptr); + p.sgd.alpha = LR; + p.sgd.wd = 0.0f; + return p; +} + +// the cache the walk reads its prefix from: F32 for the exactness checks; the served types (f16, +// q8_0, in a flash-attention context that stores V untransposed) measure the straight-through +// approximation the walk makes there (Fable on #47: the gradient is taken at the cache's values) +static ggml_type g_cache_type = GGML_TYPE_F32; + +static llama_context * make_ctx(const common_params & params, llama_model * model, uint32_t chunk) { + auto cparams = common_context_params_to_llama(params); + cparams.n_ctx = WINDOW; + cparams.n_batch = chunk; + cparams.n_ubatch = chunk; + cparams.n_seq_max = 2; // the plain walk snapshots a recurrent state into a scratch sequence + cparams.kv_unified = true; // ...and the window keeps every cell + // a quantized V cache needs flash attention (it stores V untransposed); training graphs take + // the explicit path either way + cparams.flash_attn_type = g_cache_type == GGML_TYPE_F32 ? LLAMA_FLASH_ATTN_TYPE_DISABLED : LLAMA_FLASH_ATTN_TYPE_ENABLED; + cparams.type_k = g_cache_type; + cparams.type_v = g_cache_type; + return llama_init_from_model(model, cparams); +} + +// every A and B tensor of the adapter, in name order: SGD moves them by exactly -lr * gradient +static std::vector adapter_params(const llama_adapter_lora * adapter) { + std::map sorted; + for (const auto & [name, w] : adapter->ab_map) { + sorted[name] = &w; + } + std::vector out; + for (const auto & [name, w] : sorted) { + for (ggml_tensor * t : { w->a, w->b }) { + GGML_ASSERT(t->type == GGML_TYPE_F32); + const size_t n = out.size(); + out.resize(n + ggml_nelements(t)); + ggml_backend_tensor_get(t, out.data() + n, 0, ggml_nbytes(t)); + } + } + return out; +} + +// one SGD step from the fresh adapter: the step itself, -lr * the gradient it took +static std::vector step_delta(const common_params & params, llama_model * model, const std::string & init, + const std::vector & tokens, const std::vector & labelled, + uint32_t chunk, bool exact, size_t host_budget = 0, size_t * host_bytes = nullptr) { + llama_adapter_lora * adapter = llama_adapter_lora_init(model, init.c_str()); + GGML_ASSERT(adapter != nullptr); + const std::vector before = adapter_params(adapter); + + llama_context * ctx = make_ctx(params, model, chunk); + float scale = 1.0f; + GGML_ASSERT(llama_set_adapters_lora(ctx, &adapter, 1, &scale) == 0); + llama_opt_params lopt{}; + lopt.n_ctx_train = 0; + lopt.param_filter = llama_opt_param_filter_all; + lopt.get_opt_pars = sgd_pars; + lopt.optimizer_type = GGML_OPT_OPTIMIZER_TYPE_SGD; + lopt.adapter = adapter; + lopt.walk_exact = exact; + lopt.walk_horizon = 0; + llama_opt_init(ctx, model, lopt); + llama_opt_set_walk_host_budget(ctx, host_budget); + std::vector> seqs = { std::vector(tokens.begin(), tokens.begin() + WINDOW + 1) }; + std::vector> loss = { labelled }; + ggml_opt_dataset_t dataset = common_opt_dataset_init_masked(WINDOW, seqs, loss, tokens[0]); + ggml_opt_result_t result = ggml_opt_result_init(); + llama_opt_epoch(ctx, dataset, result, nullptr, /*idata_split =*/ 1, nullptr, nullptr); + GGML_ASSERT(!llama_opt_failed(ctx)); + { + double l = 0.0, unc = 0.0; + ggml_opt_result_loss(result, &l, &unc); + printf(" chunk %3u %-5s: training loss %.6f (the forward the step saw)\n", chunk, exact ? "exact" : "plain", l); + } + if (host_bytes) { + *host_bytes = llama_opt_walk_host_bytes(ctx); + } + ggml_opt_result_free(result); + ggml_opt_dataset_free(dataset); + llama_free(ctx); + + const std::vector after = adapter_params(adapter); + llama_adapter_lora_free(adapter); + std::vector d(before.size()); + for (size_t i = 0; i < d.size(); ++i) { + d[i] = after[i] - before[i]; + } + return d; +} + +struct agreement { + double cosine; + double norm_ratio; +}; + +// the angle and the size of a step against the reference's: a wrong direction (a misrouted +// gradient) shows in the cosine, a wrong scale (a misweighted loss) in the norm ratio +static agreement compare(const char * what, const std::vector & a, const std::vector & ref) { + double dot = 0.0, na = 0.0, nr = 0.0; + for (size_t i = 0; i < a.size(); ++i) { + dot += (double) a[i] * ref[i]; + na += (double) a[i] * a[i]; + nr += (double) ref[i] * ref[i]; + } + const agreement r = { dot / std::sqrt(na * nr), std::sqrt(na / nr) }; + printf(" %s: cosine %.4f, norm ratio %.4f\n", what, r.cosine, r.norm_ratio); + return r; +} + +// Measured on the 1.5B Q4_K_M (5090): the exact walk agrees with one graph at cosine 0.998 and +// norm 0.995 (what is left is the quantized matmuls' batch-size variance: a 128-row chunk and a +// 512-row graph take different kernels); the plain walk sits at cosine 0.81, norm 2.1. +// the served cache types' bar (see 1b). Measured on the 5090: f16 0.9986 (1.5B) / 0.9981 (hybrid), +// q8_0 0.9978 / 0.9979, against F32's 0.998: the straight-through approximation costs nothing +// measurable, so the bar sits just under F32's own agreement. +static const double STRAIGHT_THROUGH_COSINE = 0.99; + +static bool same_step(const agreement & r) { + return r.cosine > 0.995 && std::fabs(r.norm_ratio - 1.0) < 0.02; +} + +int main(int argc, char ** argv) { + common_params params; + if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_COMMON)) { + return 1; + } + params.no_extra_bufts = true; // OUT_PROD dequantizes the standard layout only + llama_backend_init(); + auto mparams = common_model_params_to_llama(params); + llama_model * model = llama_model_load_from_file(params.model.path.c_str(), mparams); + GGML_ASSERT(model != nullptr); + + const std::string init = "test-walk-exact.init.gguf"; + GGML_ASSERT(common_lora_write_fresh(model, init, 8, 16, { "attn_q", "attn_v" }, 42)); + + // words in a fixed pseudo-random order: a window the model cannot already predict, so its + // loss, and the gradient under test, is well above the kernels' batch-size noise + static const char * words[] = { "room", "card", "claim", "review", "patch", "lease", "board", "turn", + "engine", "window", "chunk", "cache", "gradient", "adapter", "verdict", "submission" }; + std::string text; + uint32_t lcg = 12345; + while (text.size() < 8 * WINDOW) { + lcg = lcg * 1664525u + 1013904223u; + text += words[(lcg >> 16) % 16]; + text += (lcg >> 8) % 7 == 0 ? ".\n" : " "; + } + llama_context * tok_ctx = make_ctx(params, model, WINDOW); + std::vector tokens = common_tokenize(tok_ctx, text, true); + llama_free(tok_ctx); + GGML_ASSERT(tokens.size() > WINDOW + 1); + + int failures = 0; + { + // 1. every position labelled: chunk = window is one graph, the true gradient + std::vector all(WINDOW + 1, 1); + all[0] = 0; + const auto ref = step_delta(params, model, init, tokens, all, WINDOW, false); + size_t full_bytes = 0; + const auto exact = step_delta(params, model, init, tokens, all, WINDOW / 4, true, 0, &full_bytes); + const auto plain = step_delta(params, model, init, tokens, all, WINDOW / 4, false); + if (!same_step(compare("all labelled: exact walk in 4 chunks vs one graph", exact, ref))) { + fprintf(stderr, "FAILED: the exact walk's step is not the window's gradient\n"); + ++failures; + } + if (same_step(compare("all labelled: plain walk in 4 chunks vs one graph", plain, ref))) { + fprintf(stderr, "FAILED: the plain walk matches one graph too: this window cannot tell the two apart\n"); + ++failures; + } + if (llama_model_is_hybrid(model)) { + // a host budget below every chunk's state snapshot: the walk checkpoints every few + // chunks and decodes forward from the checkpoint, and the step must not change + size_t strided_bytes = 0; + const auto strided = step_delta(params, model, init, tokens, all, WINDOW / 4, true, full_bytes * 3 / 4, &strided_bytes); + printf(" host memory: %.1f MiB with a snapshot per chunk, %.1f MiB checkpointed\n", full_bytes / 1048576.0, strided_bytes / 1048576.0); + if (!(strided_bytes < full_bytes) || !same_step(compare("all labelled: checkpointed state vs a snapshot per chunk", strided, exact))) { + fprintf(stderr, "FAILED: checkpointing the recurrent state changed the step (or saved nothing)\n"); + ++failures; + } + } + } + { + // 1b. the served cache types: the walk reads the prefix at the cache's precision and its + // gradient is taken there, then applied to the chunk's own full-precision K/V (straight- + // through). Measured against the same one-graph reference, which never reads the cache. + std::vector all(WINDOW + 1, 1); + all[0] = 0; + const auto ref = step_delta(params, model, init, tokens, all, WINDOW, false); + for (ggml_type t : { GGML_TYPE_F16, GGML_TYPE_Q8_0 }) { + g_cache_type = t; + const auto exact = step_delta(params, model, init, tokens, all, WINDOW / 4, true); + g_cache_type = GGML_TYPE_F32; + const agreement r = compare((std::string("all labelled, ") + ggml_type_name(t) + " cache: exact walk in 4 chunks vs one graph").c_str(), exact, ref); + if (!(r.cosine > STRAIGHT_THROUGH_COSINE)) { + fprintf(stderr, "FAILED: at a %s cache the exact walk's step strays from the window's gradient (cosine %.4f)\n", ggml_type_name(t), r.cosine); + ++failures; + } + } + } + { + // 2. context, then a reply (the walk's real case): the exact walk is invariant to how the + // window is chunked, context chunks included (they train through the surrogate alone) + std::vector reply(WINDOW + 1, 0); + for (uint32_t i = WINDOW / 2 + 37; i <= WINDOW; ++i) { + reply[i] = 1; + } + const auto runs = step_delta(params, model, init, tokens, reply, WINDOW, true); // one chunk per run + const auto exact4 = step_delta(params, model, init, tokens, reply, WINDOW / 4, true); + const auto plain4 = step_delta(params, model, init, tokens, reply, WINDOW / 4, false); + if (!same_step(compare("context + reply: exact walk in 4 chunks vs one chunk per run", exact4, runs))) { + fprintf(stderr, "FAILED: the exact walk's step depends on the chunking\n"); + ++failures; + } + if (same_step(compare("context + reply: plain walk in 4 chunks vs the exact walk", plain4, runs))) { + fprintf(stderr, "FAILED: the plain walk matches the exact one: this window cannot tell the two apart\n"); + ++failures; + } + } + + { + // 3. the host budget: a window whose accumulators and snapshots exceed it refuses by + // name before anything runs, and the adapter is untouched + llama_adapter_lora * adapter = llama_adapter_lora_init(model, init.c_str()); + const std::vector before = adapter_params(adapter); + llama_context * ctx = make_ctx(params, model, WINDOW / 4); + float scale = 1.0f; + GGML_ASSERT(llama_set_adapters_lora(ctx, &adapter, 1, &scale) == 0); + llama_opt_params lopt{}; + lopt.param_filter = llama_opt_param_filter_all; + lopt.get_opt_pars = sgd_pars; + lopt.optimizer_type = GGML_OPT_OPTIMIZER_TYPE_SGD; + lopt.adapter = adapter; + lopt.walk_exact = true; + llama_opt_init(ctx, model, lopt); + llama_opt_set_walk_host_budget(ctx, 1); + std::vector> seqs = { std::vector(tokens.begin(), tokens.begin() + WINDOW + 1) }; + std::vector> loss = { std::vector(WINDOW + 1, 1) }; + ggml_opt_dataset_t dataset = common_opt_dataset_init_masked(WINDOW, seqs, loss, tokens[0]); + llama_opt_epoch(ctx, dataset, nullptr, nullptr, /*idata_split =*/ 1, nullptr, nullptr); + const std::string why = llama_opt_failure(ctx); + printf(" host budget of 1 byte: failed=%d, \"%s\", needs %zu MiB\n", (int) llama_opt_failed(ctx), why.c_str(), llama_opt_walk_host_bytes(ctx) >> 20); + if (!llama_opt_failed(ctx) || why.find("host memory") == std::string::npos || adapter_params(adapter) != before) { + fprintf(stderr, "FAILED: a window over the host budget was not refused by name before it ran\n"); + ++failures; + } + ggml_opt_dataset_free(dataset); + llama_free(ctx); + llama_adapter_lora_free(adapter); + } + + std::remove(init.c_str()); + llama_model_free(model); + llama_backend_free(); + if (failures) { + return 1; + } + printf("test-walk-exact: OK\n"); + return 0; +} diff --git a/tools/server/server-train.cpp b/tools/server/server-train.cpp index 10e472fcb5ed..9328d433ff17 100644 --- a/tools/server/server-train.cpp +++ b/tools/server/server-train.cpp @@ -314,7 +314,7 @@ json server_trainer::start(const json & body_in) { // assert in the training path would take the server down (Cormac on #14). { auto num = [&](const char * key, double def) { return body.contains(key) && body.at(key).is_number() ? body.at(key).get() : def; }; - for (const char * key : {"rank", "alpha", "window", "epochs", "lr", "val_split", "seed", "memory_budget_mib", "top_layers", "share_ppm", "max_slowdown_ppm", "chunk"}) { + for (const char * key : {"rank", "alpha", "window", "epochs", "lr", "val_split", "seed", "memory_budget_mib", "top_layers", "share_ppm", "max_slowdown_ppm", "chunk", "walk_horizon", "walk_host_budget_mib"}) { if (body.contains(key) && !body.at(key).is_number()) { return json::object({{"ok", false}, {"error", std::string("\"") + key + "\" must be a number"}}); } @@ -338,6 +338,12 @@ json server_trainer::start(const json & body_in) { else if (!(val >= 0 && val < 1)) why = "val_split must be in [0, 1)"; else if (body.contains("chunk") && !(num("chunk", 0) >= 256 && num("chunk", 0) == (int64_t) num("chunk", 0) && (int64_t) num("chunk", 0) % 256 == 0)) why = "chunk must be a multiple of 256, at least 256: the most tokens one training step holds in one graph"; + else if (body.contains("exact") && !body.at("exact").is_boolean()) + why = "exact must be true or false: the exact walk (one step per window on its whole loss) or the per-chunk walk"; + else if (body.contains("walk_host_budget_mib") && !(num("walk_host_budget_mib", 0) >= 1 && num("walk_host_budget_mib", 0) == (int64_t) num("walk_host_budget_mib", 0))) + why = "walk_host_budget_mib must be an integer >= 1: the host memory the exact walk may keep per window"; + else if (body.contains("walk_horizon") && !(num("walk_horizon", 0) >= 0 && num("walk_horizon", 0) == (int64_t) num("walk_horizon", 0) && num("walk_horizon", 0) <= 4294967295.0)) + why = "walk_horizon must be an integer >= 0: how many cached positions before a chunk receive its gradient (0 = the whole window)"; else if (body.contains("memory_budget_mib") && !(num("memory_budget_mib", 0) >= 1)) why = "memory_budget_mib must be >= 1"; else if (body.contains("share_ppm") && !(num("share_ppm", 0) >= 1 && num("share_ppm", 0) <= 1000000 && num("share_ppm", 0) == (int64_t) num("share_ppm", 0))) @@ -861,8 +867,14 @@ void server_trainer::run(json req, examples_data ex) { // with it on is measured beside #39's receipt; Cormac): one layer's // attention scores alive in the backward pass instead of every layer's /*recompute =*/ req.value("recompute", false), + // the exact walk (OFF unless the request says "exact": true, until it is measured against + // the single-context gradient): one step per window on the window's whole loss + /*walk_exact =*/ req.value("exact", false), + /*walk_horizon =*/ (uint32_t) req.value("walk_horizon", (int64_t) 0), }; llama_opt_set_memory_budget(ctx, budget); + // the exact walk's host memory per window: refused by name over the caller's cap + llama_opt_set_walk_host_budget(ctx, (size_t) req.value("walk_host_budget_mib", (int64_t) 0) << 20); llama_opt_init(ctx, model, lopt); llama_opt_set_step_callback(ctx, &server_trainer::before_window, this); { @@ -966,6 +978,8 @@ void server_trainer::run(json req, examples_data ex) { // copied before llama_free: the refusal names a node the device cannot run, when that was it const std::string refusal = llama_opt_failure(ctx); const double graph_mib = llama_opt_graph_bytes(ctx) / 1048576.0; + const double walk_host_mib = llama_opt_walk_host_bytes(ctx) / 1048576.0; + const int64_t walk_horizon = llama_opt_walk_horizon(ctx); const int32_t saved = (cancelled || no_fit) ? 0 : llama_adapter_lora_save(adapter, out.c_str()); llama_adapter_lora_free(adapter); llama_free(ctx); @@ -974,6 +988,13 @@ void server_trainer::run(json req, examples_data ex) { // what the training graph needed on the GPU, measured by its own preflight (also on a refusal: // it is the number to size the next attempt, or a caller's admission, by) state["graph_mib"] = graph_mib; + // what the exact walk kept on the host per window (its K/V gradient and state snapshots) + state["walk_host_mib"] = walk_host_mib; + // the exact walk's gradient horizon as trained (0 = the whole window), smaller where the device + // could not hold the requested one + if (req.value("exact", false)) { + state["walk_horizon"] = walk_horizon; + } if (no_fit) { state["state"] = "error"; state["error"] = !refusal.empty() ? refusal + " (window " + std::to_string(window) + ")" :