llama: keep the backend sampling graph static across ubatches (#30223)

The reserve builds n_outputs_max_per_seq sampling chains per sampler,
while a decode built one per output row, so the graph changed its
topology after the reserve and GGML_SCHED_NO_REALLOC builds aborted on
the next same sized graph. Every sampler now builds
n_outputs_max_per_seq chains, the ones without a row of the ubatch on
the padding row and not selected, and graph_max_nodes counts them.
This commit is contained in:
Pascal
2026-10-09 17:20:25 +03:00
committed by GitHub
parent 8ae386707b
commit a518119d30
2 changed files with 16 additions and 29 deletions
+2 -16
View File
@@ -2442,23 +2442,9 @@ uint32_t llama_context::graph_max_nodes(uint32_t n_tokens) const {
}
}
uint32_t n_sampling_nodes = 0;
uint32_t n_sampling_nodes_max = 0;
// every sampler builds n_outputs_max_per_seq chains, see llm_graph_context::build_sampling
for (const auto & [seq_id, sampler] : sampling.samplers) {
const uint32_t n_nodes = llama_sampler_backend_n_nodes(sampler);
n_sampling_nodes += n_nodes;
if (cparams.n_outputs_max_per_seq > 1) {
n_sampling_nodes_max = std::max(n_sampling_nodes_max, n_nodes);
}
}
const uint32_t n_sampling_outputs_max = std::min<uint64_t>(
std::min(n_tokens, cparams.n_outputs_max),
(uint64_t) cparams.n_seq_max * cparams.n_outputs_max_per_seq);
res += n_sampling_nodes;
if (n_sampling_outputs_max > 1) {
res += (n_sampling_outputs_max - 1) * n_sampling_nodes_max;
res += llama_sampler_backend_n_nodes(sampler) * cparams.n_outputs_max_per_seq;
}
if (cparams.training) {
+14 -13
View File
@@ -3965,8 +3965,7 @@ void llm_graph_context::build_sampling() const {
// res->t_logits will contain logits for all tokens that want the logits calculated (logits=1 or output=1)
GGML_ASSERT(res->t_logits != nullptr && "missing t_logits tensor");
// add a dummy row to keep the single-output graph static regardless of active samplers
// multi-output graphs can still vary with the number of output rows
// the padding row gives the chains without a row of the ubatch a valid input, even with no output at all
ggml_tensor * logits_t = ggml_pad(ctx0, res->t_logits, 0, 1, 0, 0);
for (const auto & entry : samplers) {
@@ -3975,18 +3974,20 @@ void llm_graph_context::build_sampling() const {
}
}
static const std::vector<uint32_t> dummy_row = { 0 };
static const std::vector<uint32_t> no_rows;
// every sampler builds n_outputs_max_per_seq chains, like the reserve, so the graph keeps its topology
// whatever rows the ubatch outputs: a chain without a row works on the first row and is not selected
for (const auto & [seq_id, sampler] : samplers) {
const auto it = sampling_rows.find(seq_id);
const auto & rows = it != sampling_rows.end() ? it->second : no_rows;
// inactive samplers always work on the first row
const bool active = it != sampling_rows.end();
const auto & rows = active ? it->second : dummy_row;
const int i_out = active ? 1 : 0;
for (uint32_t i = 0; i < cparams.n_outputs_max_per_seq; ++i) {
const bool active = i < rows.size();
const uint32_t row = active ? rows[i] : 0;
const int i_out = active ? 1 : 0;
for (uint32_t i = 0; i < rows.size(); ++i) {
ggml_tensor * logits_seq = ggml_view_1d(ctx0, logits_t, logits_t->ne[0], rows[i] * logits_t->nb[1]);
ggml_tensor * logits_seq = ggml_view_1d(ctx0, logits_t, logits_t->ne[0], row * logits_t->nb[1]);
ggml_format_name(logits_seq, "logits_seq_%d_%u", seq_id, i);
struct llama_sampler_data data = {
@@ -4001,7 +4002,7 @@ void llm_graph_context::build_sampling() const {
if (data.sampled != nullptr) {
if (active) {
res->t_sampled[rows[i]] = data.sampled;
res->t_sampled[row] = data.sampled;
}
outs[1] = data.sampled;
ggml_build_forward_select(gf, outs.data(), outs.size(), i_out);
@@ -4009,7 +4010,7 @@ void llm_graph_context::build_sampling() const {
if (data.probs != nullptr) {
if (active) {
res->t_sampled_probs[rows[i]] = data.probs;
res->t_sampled_probs[row] = data.probs;
}
outs[1] = data.probs;
ggml_build_forward_select(gf, outs.data(), outs.size(), i_out);
@@ -4017,7 +4018,7 @@ void llm_graph_context::build_sampling() const {
if (data.logits != nullptr) {
if (active) {
res->t_sampled_logits[rows[i]] = data.logits;
res->t_sampled_logits[row] = data.logits;
}
outs[1] = data.logits;
ggml_build_forward_select(gf, outs.data(), outs.size(), i_out);
@@ -4025,7 +4026,7 @@ void llm_graph_context::build_sampling() const {
if (data.candidates != nullptr) {
if (active) {
res->t_candidates[rows[i]] = data.candidates;
res->t_candidates[row] = data.candidates;
}
outs[1] = data.candidates;
ggml_build_forward_select(gf, outs.data(), outs.size(), i_out);