Skip to content
212 changes: 178 additions & 34 deletions src/llama-context.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,9 @@
#include "llama-batch.h"
#include "llama-io.h"
#include "llama-memory.h"
#include "llama-memory-hybrid.h"
#include "llama-memory-recurrent.h"
#include "llama-kv-cache.h"
#include "llama-mmap.h"
#include "llama-model.h"
#include "llama-ext.h"
Expand Down Expand Up @@ -582,6 +585,13 @@ void llama_context::sched_reserve() {
if (!sched_need_reserve) {
return;
}
// A training context's scheduler was sized for the training graph in opt_init, and the
// optimizer holds it (ggml_opt_params.backend_sched). Re-creating it here, as the first
// context decode of opt_epoch_iter's walk would, leaves the optimizer on a freed scheduler.
if (opt_ctx != nullptr) {
sched_need_reserve = false;
return;
}

sched_need_reserve = false;

Expand Down Expand Up @@ -2335,7 +2345,13 @@ uint32_t llama_context::graph_max_nodes(uint32_t n_tokens) const {
if (n_sampling_outputs_max > 1) {
res += (n_sampling_outputs_max - 1) * n_sampling_nodes_max;
}
if (cparams.training) {
// a training context (an optimizer exists), not this graph's flag: the walk decodes context
// with cparams.training off, and a sched_reserve() run then must not shrink the scheduler
// and the graph result below what the next training chunk needs
if (cparams.training || opt_ctx != nullptr) {
// each attention layer reads the cached context in front of the chunk (build_attn:
// views of K and V, V's permute and cont, two concats, the mask's view and cont)
res += 16u * model.hparams.n_layer();
// ggml_opt duplicates the forward graph at its own size, then appends the backward
// pass and the optimizer step: room for the forward three times over.
res *= 3;
Expand Down Expand Up @@ -3324,10 +3340,14 @@ void llama_context::opt_init(struct llama_model * model, struct llama_opt_params
GGML_ASSERT(opt_n_ctx_train % n_batch == 0);
GGML_ASSERT(n_batch % n_ubatch == 0);
// A training graph cannot backprop through the KV cache (SET_ROWS writes in place and
// attention reads the cache as a leaf), so attention consumes this ubatch's K/V
// directly; that is exact only when one ubatch is the whole context.
GGML_ASSERT(n_ubatch == opt_n_ctx_train && "training needs one ubatch per context: set -ub = -b = -c");
GGML_ASSERT(!cparams.flash_attn && "training needs flash attention off: FLASH_ATTN_EXT has no backward");
// attention reads the cache as a leaf), so attention consumes this chunk's K/V directly
// and the context before it from the cache as a constant (build_attn); opt_epoch_iter
// walks the window decoding context and training each labelled chunk. One chunk is one
// batch is one ubatch: a chunk's backward needs all of its own K/V in one graph.
GGML_ASSERT(n_ubatch == n_batch && "training needs one ubatch per batch: set -ub = -b");
// FLASH_ATTN_EXT has no backward: training graphs build explicit attention whatever the
// context's flag (build_attn_mha checks cparams.training); flash attention, when the context
// has it, serves the walk's context decodes and lets the cache hold V un-transposed.
cparams.training = true;
// The graph result and the scheduler were sized at context creation: before training
// (no backward pass) and before any adapter attached since (its nodes uncounted).
Expand Down Expand Up @@ -3404,25 +3424,82 @@ void llama_context::opt_epoch_iter(
int64_t ndata_in_loop,
int64_t t_loop_start) {
GGML_ASSERT(opt_ctx);
const uint32_t n_ctx = opt_n_ctx_train;
const uint32_t n_batch = std::min(this->n_batch(), n_ctx);
const uint32_t n_ubatch = std::min(this->n_ubatch(), n_batch);
const uint32_t n_ctx = opt_n_ctx_train;
const uint32_t n_batch = std::min(this->n_batch(), n_ctx);

memory->clear(true);

// OUTPUTS ONLY WHERE THE LOSS IS. A position whose label is masked (< 0: the system and
// tool head, the user's turns, padding) carries no loss, so its logits are never read;
// computing them anyway cost n_vocab x window floats for the logits, the same again for
// the one-hot labels, and the same again for the logit gradients. At 151,936 x 256 that
// was the 148 MiB the memory gate named on the M5; at one of a citizen's 15k-token
// turns it was ~9 GB three times over, and "could not create the training context" on
// the 5090 (continuum card 36c3c00a). Decode already selects output rows per token
// (batch.logits); training now does the same, and the label tensor is sized to the
// rows that exist. Every window has at least one labelled position: the server refuses
// an example with none before it reaches here.
for (uint32_t pos_ctx = 0; pos_ctx < n_ctx; pos_ctx += n_batch) {
batch.n_tokens = n_batch;
for (uint32_t pos_batch = 0; pos_batch < n_batch; ++pos_batch) {
// THE WALK (Fable's design, continuum 2026-10-06): one pass over the window in one
// context. A run with no label is context she did not write: a plain decode into the
// cache, no backward. A labelled run is her reply: it trains in chunks of n_batch, each
// attending to everything before it read from the cache as a constant (build_attn,
// cparams.training), then the reply is decoded into the cache under the adapter as it
// now is, so the next reply sees it. Memory is chunk x window, never window x window, so
// her whole lived window trains. Nothing after the last label is computed.
int64_t last_label = -1;
for (int64_t i = (int64_t) n_ctx - 1; i >= 0; --i) {
if (labels_sparse[i] >= 0) {
last_label = i;
break;
}
}
if (last_label < 0) {
LLAMA_LOG_ERROR("%s: a training window with no labelled position: nothing to learn from it\n", __func__);
return;
}

// THE YIELD POINT IS THE CHUNK, not the window: a turn that arrives mid-walk waits for one
// chunk, never for the whole example (measured on the 5090: one callback per example made a
// turn wait for its entire walk, 1.2-1.8 s on a 1.5B at 15k and 13.5 s on Fable's run). The
// first unit is the window's own boundary, which opt_epoch has already offered.
bool first_unit = true;
auto yield_point = [&]() -> bool {
if (first_unit) {
first_unit = false;
} else if (opt_step_callback && !opt_step_callback(train, opt_step_callback_data)) {
opt_stop_requested.store(true);
}
return !opt_stop_requested.load(std::memory_order_relaxed);
};

// a plain inference decode of [p0, p1) into the cache: the training graph's attention
// never writes the cache, so context and finished replies reach it this way
auto decode_span = [&](uint32_t p0, uint32_t p1) -> bool {
const bool training = cparams.training;
cparams.training = false;
gf_res_prev->reset();
bool ok = true;
for (uint32_t c0 = p0; c0 < p1 && ok; c0 += n_batch) {
if (!yield_point()) {
cparams.training = training;
gf_res_prev->reset();
return false; // stopped or cancelled at a chunk boundary: not a decode failure
}
const uint32_t c1 = std::min(c0 + n_batch, p1);
batch.n_tokens = c1 - c0;
for (uint32_t i = 0; i < c1 - c0; ++i) {
batch.token [i] = tokens[c0 + i];
batch.pos [i] = c0 + i;
batch.n_seq_id[i] = 1;
batch.seq_id [i][0] = 0;
batch.logits [i] = i + 1 == c1 - c0;
}
const int rc = decode(batch);
ok = rc == 0;
}
cparams.training = training;
gf_res_prev->reset();
if (!ok) {
LLAMA_LOG_ERROR("%s: the context decode of [%u, %u) failed\n", __func__, p0, p1);
}
return ok;
};

// 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 {
batch.n_tokens = n_chunk;
for (uint32_t pos_batch = 0; pos_batch < n_chunk; ++pos_batch) {
batch.token [pos_batch] = tokens[pos_ctx + pos_batch];
batch.pos [pos_batch] = pos_ctx + pos_batch;
batch.n_seq_id[pos_batch] = 1;
Expand All @@ -3433,7 +3510,7 @@ void llama_context::opt_epoch_iter(
// output_all = false: the batch's own logits flags decide which rows exist
if (!balloc->init(batch, model.vocab, nullptr, model.hparams.n_embd_inp(), cparams.kv_unified ? LLAMA_MAX_SEQ : cparams.n_seq_max, false)) {
LLAMA_LOG_ERROR("%s: failed to initialize batch\n", __func__);
return;
return false;
}

const uint32_t n_tokens_all = balloc->get_n_tokens();
Expand All @@ -3443,21 +3520,17 @@ void llama_context::opt_epoch_iter(
embd_seq.clear();

const uint32_t n_outputs_all = balloc->get_n_outputs();
// Unreachable from /train: it sets n_ctx = n_batch = n_ubatch = window (one ubatch is
// the whole context, asserted in opt_init) and refuses an example with no labelled
// position before it gets here, so every batch has at least one. Kept for any other
// caller of opt_epoch, and it RETURNS from the whole iteration (a later batch of this
// context would need the earlier one's state; with one batch per context there is
// none to lose). (Cormac and BigMama on fork #32.)
// Unreachable: the walk below trains only chunks of a labelled run. It RETURNS from the
// whole window, since a later chunk would need this one in the cache.
if (n_outputs_all == 0) {
LLAMA_LOG_ERROR("%s: a training batch with no labelled position (every label masked): nothing to learn from it\n", __func__);
return;
return false;
}

auto mctx = memory->init_batch(*balloc, cparams.n_ubatch, true);
if (!mctx || mctx->get_status() != LLAMA_MEMORY_STATUS_SUCCESS) {
LLAMA_LOG_ERROR("%s: could not initialize batch\n", __func__);
break;
return false;
}

// reserve output buffer
Expand All @@ -3481,7 +3554,7 @@ void llama_context::opt_epoch_iter(

if (!mctx->apply()) {
LLAMA_LOG_ERROR("%s: failed to update the memory context\n", __func__);
break;
return false;
}

auto * res = gf_res_prev.get();
Expand Down Expand Up @@ -3513,7 +3586,7 @@ void llama_context::opt_epoch_iter(
ggml_free(ctx_compute_opt);
opt_alloc_failed.store(true);
opt_stop_requested.store(true);
return;
return false;
}

res->set_inputs(&ubatch);
Expand Down Expand Up @@ -3544,15 +3617,86 @@ void llama_context::opt_epoch_iter(
ggml_free(ctx_compute_opt);
opt_alloc_failed.store(true);
opt_stop_requested.store(true);
return;
return false;
}
if (callback) {
callback(train, opt_ctx, dataset, result, idata_in_loop + (pos_ctx + pos_batch)/n_ubatch + 1, ndata_in_loop, t_loop_start);
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);

pos_batch += ubatch.n_tokens;
} while (mctx->next());
return true;
};

// a recurrent part (pure recurrent or hybrid) is snapshotted around each training chunk;
// a hybrid's attention part removes the chunk's cells like a plain cache
llama_memory_recurrent * recr = nullptr;
llama_kv_cache * attn = nullptr;
if (auto * hybrid = dynamic_cast<llama_memory_hybrid *>(memory.get())) {
recr = hybrid->get_mem_recr();
attn = hybrid->get_mem_attn();
} else {
recr = dynamic_cast<llama_memory_recurrent *>(memory.get());
}
if (recr != nullptr && cparams.n_seq_max < 2) {
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;
}

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;
uint32_t end = pos;
while (end <= (uint32_t) last_label && (labels_sparse[end] >= 0) == labelled) {
++end;
}
if (!labelled) {
if (!decode_span(pos, end)) {
return;
}
} else {
for (uint32_t c0 = pos; c0 < end; c0 += n_batch) {
const uint32_t c1 = std::min(c0 + n_batch, end);
// The training forward advances a recurrent state in place, and a recurrent
// memory cannot remove a partial range: snapshot the state to the scratch
// sequence first (Fable), restore it after, then the decode below advances it
// from where the context left it. The state is per sequence and small.
const bool snapshot = recr != nullptr && c1 <= (uint32_t) last_label;
if (!yield_point()) {
return;
}
if (snapshot) {
recr->seq_cp(0, 1, -1, -1);
}
if (!train_chunk(c0, c1 - c0)) {
return;
}
if (snapshot) {
recr->seq_rm(0, -1, -1);
recr->seq_cp(1, 0, -1, -1);
recr->seq_rm(1, -1, -1);
}
// the chunk joins the context under the adapter as it now is: the training
// graph never writes the cache, so its cells hold nothing until this decode
if (c1 <= (uint32_t) last_label) {
// the chunk's cells hold no K/V (the training graph never writes the
// cache): drop them from the attention memory, then decode them for real
const bool removed = recr != nullptr ? (attn == nullptr || attn->seq_rm(0, c0, -1)) : memory->seq_rm(0, c0, -1);
if (!removed) {
LLAMA_LOG_ERROR("%s: could not remove the trained chunk [%u, %u) from the cache before decoding it\n", __func__, c0, c1);
opt_stop_requested.store(true);
return;
}
if (!decode_span(c0, c1)) {
opt_stop_requested.store(true);
return;
}
}
}
}
pos = end;
}
}

Expand Down
48 changes: 41 additions & 7 deletions src/llama-graph.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2574,7 +2574,10 @@ ggml_tensor * llm_graph_context::build_attn_mha(

ggml_tensor * cur;

const bool use_flash_attn = cparams.flash_attn && kq_b == nullptr;
// A training graph takes the explicit path, which has a backward (FLASH_ATTN_EXT has none),
// even in a context created with flash attention: there the walk's context decodes use
// flash attention and the cache stores V un-transposed, so it can be quantized.
const bool use_flash_attn = cparams.flash_attn && kq_b == nullptr && !cparams.training;
if (use_flash_attn) {
GGML_ASSERT(kq_b == nullptr && "Flash attention does not support KQ bias yet");

Expand Down Expand Up @@ -2836,14 +2839,45 @@ ggml_tensor * llm_graph_context::build_attn(
ggml_tensor * v;
if (cparams.training) {
// A backward pass cannot go through the cache: the write is an in-place SET_ROWS
// and the read is a view of a leaf, so K/V would get no gradient. Training runs one
// ubatch per context (llama_context::opt_init asserts it), so this ubatch's K/V ARE
// the whole context: attend to them directly, as the no-cache path does.
// and the read is a view of a leaf. So this ubatch's K/V are attended to directly,
// and carry the gradient. The context before it (positions [0, n_past), decoded into
// the cache by llama_context::opt_epoch_iter before this chunk trains) is read from
// the cache as a CONSTANT, concatenated in front: her reply sees the whole window,
// and memory is chunk x window instead of window x window. The prefix's own K/V get
// no gradient (a stop-gradient at the chunk boundary).
const int64_t n_tokens = k_cur->ne[2];
const int64_t n_past = ubatch.pos[0];
k = k_cur;
v = v_cur;
if (kq_mask->ne[0] > k_cur->ne[2]) {
kq_mask = ggml_view_4d(ctx0, kq_mask, k_cur->ne[2], kq_mask->ne[1], kq_mask->ne[2], kq_mask->ne[3],
kq_mask->nb[1], kq_mask->nb[2], kq_mask->nb[3], 0);
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 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);

// 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_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),
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);
}
v = ggml_concat(ctx0, v_prev, 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)
if (kq_mask->ne[0] > n_past + n_tokens) {
kq_mask = ggml_cont(ctx0, ggml_view_4d(ctx0, kq_mask, n_past + n_tokens, kq_mask->ne[1], kq_mask->ne[2], kq_mask->ne[3],
kq_mask->nb[1], kq_mask->nb[2], kq_mask->nb[3], 0));
}
} else {
// store to KV cache
Expand Down
Loading
Loading