Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 46 additions & 18 deletions src/llama-context.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3925,26 +3925,36 @@ void llama_context::opt_epoch_iter(
}
// 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]);
// in between too, under the same adapter: the same values). Run again before every
// retry below: a refused chunk graph is refused AFTER its ubatch was applied, so the
// chunk's cells already sit in the cache and the recurrent cell's position has moved
// to the chunk's end; retrying on top of that fails to prepare the ubatch at all
// (the 5090, 2026-10-10: "init_batch: failed to prepare attention ubatches" after
// "retrying at 512", and the epoch stopped instead of shrinking the horizon).
const auto rewind = [&]() -> bool {
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);
return false;
}
if (p_pop < c.c0 && !decode_span(p_pop, c.c0)) {
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)) {
return false;
}
}
return true;
};
if (!rewind()) {
ok = false;
break;
}
cparams.walk_exact = true;
cparams.walk_surrogate = c.surrogate;
Expand Down Expand Up @@ -3973,6 +3983,24 @@ void llama_context::opt_epoch_iter(
__func__, cc.c0, span, opt_walk_horizon);
opt_alloc_failed.store(false);
opt_stop_requested.store(false);
// the rewind's decode is a plain forward, as before the first try: no walk flags.
// walk_grad_from is not restored here because the loop's first statements recompute
// it from the shrunken horizon (cc.grad_from, then cparams.walk_grad_from) before the
// retried train_chunk runs.
const bool surrogate = cparams.walk_surrogate;
const bool state_surrogate = cparams.walk_state_surrogate;
cparams.walk_exact = false;
cparams.walk_grad_from = 0;
cparams.walk_surrogate = false;
cparams.walk_state_surrogate = false;
const bool rewound = rewind();
cparams.walk_exact = true;
cparams.walk_surrogate = surrogate;
cparams.walk_state_surrogate = state_surrogate;
if (!rewound) {
ok = false;
break;
}
}
opt_walk_horizon_used = opt_walk_horizon;
cparams.walk_exact = false;
Expand Down
99 changes: 93 additions & 6 deletions tests/test-walk-exact.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,10 @@

static const float LR = 1e-3f; // an SGD step is linear in the gradient at any lr
static const uint32_t WINDOW = 512;
// the window a step trains: WINDOW for every case but the device refusal (6), whose horizon must
// change the chunk graph's size, which a window this short does not
static const uint32_t REFUSAL_WINDOW = 2048;
static uint32_t g_window = WINDOW;

static ggml_opt_optimizer_params sgd_pars(void *) {
ggml_opt_optimizer_params p = ggml_opt_get_default_optimizer_params(nullptr);
Expand All @@ -44,23 +48,44 @@ static ggml_opt_optimizer_params sgd_pars(void *) {
static ggml_type g_cache_type = GGML_TYPE_F32;
// recurrent rollback slots in the context (serving keeps some for speculative decoding)
static uint32_t g_n_rs_seq = 0;
// the gradient horizon a step asks for (0 = the whole window) and the device budget its graphs
// may add (0 = the device's own free figure), and what the last step measured: its largest
// graph and the horizon it ended at
static uint32_t g_walk_horizon = 0;
static size_t g_device_budget = 0;
static size_t g_last_graph_bytes = 0;
static uint32_t g_last_horizon = 0;

// counts the recurrent memory's "non-consecutive token position" warnings: a state restore that
// leaves the cell's position at the window's end makes every next chunk read as non-consecutive
// (the 5090 log, 2026-10-07: "position 65640 after 66345"), passing everything else through
static int g_nonconsecutive = 0;
// and, for the device refusal (6): the horizon retries, the ubatches that failed to prepare, and
// the chunk graphs the device budget refused
static int g_retries = 0;
static int g_ubatch_failures = 0;
static int g_refusals = 0;
static void count_log(ggml_log_level level, const char * text, void * ud) {
if (text && strstr(text, "non-consecutive token position")) {
++g_nonconsecutive;
}
if (text && strstr(text, "did not fit with a gradient horizon")) {
++g_retries;
}
if (text && strstr(text, "failed to prepare attention ubatches")) {
++g_ubatch_failures;
}
if (text && strstr(text, "it may add: not allocating it")) {
++g_refusals;
}
GGML_UNUSED(level);
GGML_UNUSED(ud);
fputs(text, stderr);
}

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_ctx = g_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
Expand Down Expand Up @@ -110,12 +135,13 @@ static std::vector<float> step_delta(const common_params & params, llama_model *
lopt.optimizer_type = GGML_OPT_OPTIMIZER_TYPE_SGD;
lopt.adapter = adapter;
lopt.walk_exact = exact;
lopt.walk_horizon = 0;
lopt.walk_horizon = g_walk_horizon;
llama_opt_set_memory_budget(ctx, g_device_budget);
llama_opt_init(ctx, model, lopt);
llama_opt_set_walk_host_budget(ctx, host_budget);
std::vector<std::vector<llama_token>> seqs = { std::vector<llama_token>(tokens.begin(), tokens.begin() + WINDOW + 1) };
std::vector<std::vector<llama_token>> seqs = { std::vector<llama_token>(tokens.begin(), tokens.begin() + g_window + 1) };
std::vector<std::vector<uint8_t>> loss = { labelled };
ggml_opt_dataset_t dataset = common_opt_dataset_init_masked(WINDOW, seqs, loss, tokens[0]);
ggml_opt_dataset_t dataset = common_opt_dataset_init_masked(g_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));
Expand All @@ -127,6 +153,8 @@ static std::vector<float> step_delta(const common_params & params, llama_model *
if (host_bytes) {
*host_bytes = llama_opt_walk_host_bytes(ctx);
}
g_last_graph_bytes = llama_opt_graph_bytes(ctx);
g_last_horizon = llama_opt_walk_horizon(ctx);
ggml_opt_result_free(result);
ggml_opt_dataset_free(dataset);
llama_free(ctx);
Expand Down Expand Up @@ -191,15 +219,15 @@ int main(int argc, char ** argv) {
"engine", "window", "chunk", "cache", "gradient", "adapter", "verdict", "submission" };
std::string text;
uint32_t lcg = 12345;
while (text.size() < 8 * WINDOW) {
while (text.size() < 8 * REFUSAL_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<llama_token> tokens = common_tokenize(tok_ctx, text, true);
llama_free(tok_ctx);
GGML_ASSERT(tokens.size() > WINDOW + 1);
GGML_ASSERT(tokens.size() > REFUSAL_WINDOW + 1);

int failures = 0;
{
Expand Down Expand Up @@ -348,6 +376,65 @@ int main(int argc, char ** argv) {
llama_adapter_lora_free(adapter);
}

{
// 6. regression for the 5090 (2026-10-10 04:34Z): after the device budget refuses a chunk
// graph, the walk shrinks the gradient horizon and tries the chunk again. The refusal comes
// AFTER the chunk's ubatch was applied, so the retry used to run on top of the chunk's own
// cells and fail to prepare its ubatch at all ("init_batch: failed to prepare attention
// ubatches"), stopping the epoch instead of reaching a horizon that fits. A budget just
// under the measured graph refuses EVERY horizon, so this pins the retry itself: each retry
// must reach the device preflight again (and be refused there), never fail in init_batch,
// and the epoch then refuses by name with the adapter untouched. Needs a device: a CPU run
// measures no graph, and says so.
g_window = REFUSAL_WINDOW;
std::vector<uint8_t> all(g_window + 1, 1);
all[0] = 0;
step_delta(params, model, init, tokens, all, WINDOW / 4, true);
const size_t graph_bytes = g_last_graph_bytes;
printf(" device graph: %.1f MiB per chunk\n", graph_bytes / 1048576.0);
if (graph_bytes == 0) {
printf(" device refusal: skipped (no device graph measured: a CPU-only run)\n");
} else {
llama_adapter_lora * adapter = llama_adapter_lora_init(model, init.c_str());
const std::vector<float> 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;
// what a graph may add is the budget less the allocator's 512 MiB margin (ggml-opt), and
// "add" is growth past the buffers the walk's own forward decodes already sized: one
// byte past the margin lets a chunk graph grow by nothing, so every horizon is refused
llama_opt_set_memory_budget(ctx, 512u*1024*1024 + 1);
llama_opt_init(ctx, model, lopt);
std::vector<std::vector<llama_token>> seqs = { std::vector<llama_token>(tokens.begin(), tokens.begin() + g_window + 1) };
std::vector<std::vector<uint8_t>> loss = { all };
ggml_opt_dataset_t dataset = common_opt_dataset_init_masked(g_window, seqs, loss, tokens[0]);
g_retries = 0;
g_ubatch_failures = 0;
g_refusals = 0;
llama_log_set(count_log, nullptr);
llama_opt_epoch(ctx, dataset, nullptr, nullptr, /*idata_split =*/ 1, nullptr, nullptr);
llama_log_set(nullptr, nullptr); // the default logger again
const std::string why = llama_opt_failure(ctx);
printf(" device refusal: %d retries, %d graph refusals, %d ubatch failures, failed=%d \"%s\"\n",
g_retries, g_refusals, g_ubatch_failures, (int) llama_opt_failed(ctx), why.c_str());
if (g_retries < 1 || g_refusals < 2 || g_ubatch_failures != 0 ||
!llama_opt_failed(ctx) || adapter_params(adapter) != before) {
fprintf(stderr, "FAILED: a refused chunk graph was not retried cleanly at a smaller horizon\n");
++failures;
}
ggml_opt_dataset_free(dataset);
llama_free(ctx);
llama_adapter_lora_free(adapter);
}
g_window = WINDOW;
}

std::remove(init.c_str());
llama_model_free(model);
llama_backend_free();
Expand Down