mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-11 07:20:33 +02:00
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:
+2
-16
@@ -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
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user