diff --git a/src/llama-batch.cpp b/src/llama-batch.cpp index 89a1f3f37c..7a17b3a4e9 100644 --- a/src/llama-batch.cpp +++ b/src/llama-batch.cpp @@ -800,6 +800,7 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector & idxs, u udata->seq_idx .resize(LLAMA_MAX_SEQ, -1); udata->output .resize(n_tokens); + udata->batch_idxs = idxs; udata->seq_id_data.reserve(n_tokens); seq_set_t seq_set_unq; diff --git a/src/llama-batch.h b/src/llama-batch.h index 201d48cce1..18e8144c80 100644 --- a/src/llama-batch.h +++ b/src/llama-batch.h @@ -62,6 +62,7 @@ struct llama_ubatch { std::vector seq_id_unq; std::vector seq_idx; std::vector output; + std::vector batch_idxs; // original batch index for each token std::vector seq_id_data; }; diff --git a/src/llama-context.cpp b/src/llama-context.cpp index ba89ca794a..4f04db0ee6 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -1799,6 +1799,7 @@ int llama_context::decode(const llama_batch_ext & batch_inp) { n_queued_tokens += n_tokens_all; output_swaps.clear(); + embd_batch_idxs.clear(); sched_reserve(); @@ -2008,7 +2009,8 @@ int llama_context::decode(const llama_batch_ext & batch_inp) { } } - extract_layer_inputs(res, n_tokens_prev, ubatch.n_tokens); + // [TAG_EXTRACT_TARGET_EMBEDDINGS] + bool extract_all_idxs = extract_layer_inputs(res, n_tokens_prev, ubatch.n_tokens); // extract nextn embeddings before // only meaningful in LLAMA_POOLING_TYPE_NONE (per-token); other pooling modes are ignored. @@ -2026,9 +2028,17 @@ int llama_context::decode(const llama_batch_ext & batch_inp) { GGML_ASSERT((offset + n_rows)*n_embd <= (int64_t) embd_nextn.size); ggml_backend_tensor_get_async(backend_h, t_h_nextn, embd_nextn_out, 0, n_rows*n_embd*sizeof(float)); + extract_all_idxs = extract_all_idxs || !masked; } } + if (extract_all_idxs) { + GGML_ASSERT(ubatch.data && ubatch.data->batch_idxs.size() == ubatch.n_tokens); + GGML_ASSERT(embd_batch_idxs.size() == (size_t) n_tokens_prev); + const auto & batch_idxs = ubatch.data->batch_idxs; + embd_batch_idxs.insert(embd_batch_idxs.end(), batch_idxs.begin(), batch_idxs.end()); + } + if (has_samplers) { const auto stride = n_vocab; @@ -2265,7 +2275,8 @@ uint32_t llama_context::output_reserve(int32_t n_outputs) { return n_outputs_max; } -void llama_context::extract_layer_inputs(const llm_graph_result * res, size_t token_offset, size_t n_tokens) { +bool llama_context::extract_layer_inputs(const llm_graph_result * res, size_t token_offset, size_t n_tokens) { + bool extracted = false; for (uint32_t il = 0; il < cparams.embeddings_layer_inp.size(); ++il) { if (!cparams.embeddings_layer_inp[il]) { continue; @@ -2284,13 +2295,17 @@ void llama_context::extract_layer_inputs(const llm_graph_result * res, size_t to GGML_ASSERT(nfloats % n_tokens == 0); const size_t row_floats = nfloats / n_tokens; + GGML_ASSERT(row_floats == model.hparams.n_embd); const size_t dst_offset = token_offset * row_floats; GGML_ASSERT(dst_offset + nfloats <= embd_layer_inp[il].size); ggml_backend_t backend = ggml_backend_sched_get_tensor_backend(sched.get(), t); GGML_ASSERT(backend != nullptr); + // Tensor-split backends require a zero source offset. ggml_backend_tensor_get_async(backend, t, embd_layer_inp[il].data + dst_offset, 0, nbytes); + extracted = true; } + return extracted; } void llama_context::output_reorder() { @@ -2314,22 +2329,12 @@ void llama_context::output_reorder() { } } - if (embd_nextn.size > 0) { + if (embd_nextn.size > 0 && cparams.embeddings_nextn_masked) { for (uint64_t k = 0; k < n_embd_out; k++) { std::swap(embd_nextn.data[i0*n_embd_out + k], embd_nextn.data[i1*n_embd_out + k]); } } - if (embd_layer_inp.size() > 0) { - for (int lid = 0; lid < (int) embd_layer_inp.size(); ++lid) { - if (embd_layer_inp[lid].size > 0) { - for (uint64_t k = 0; k < n_embd; ++k) { - std::swap(embd_layer_inp[lid].data[i0*n_embd + k], embd_layer_inp[lid].data[i1*n_embd + k]); - } - } - } - } - if (!sampling.samplers.empty()) { assert(sampling.logits.size > 0); assert(sampling.probs.size > 0); @@ -2359,6 +2364,29 @@ void llama_context::output_reorder() { } output_swaps.clear(); + + // [TAG_EXTRACT_TARGET_EMBEDDINGS] + // Layer inputs and unmasked NextN embeddings contain all token rows, independent of logits selection. + for (size_t i = 0; i < embd_batch_idxs.size(); ++i) { + while (embd_batch_idxs[i] != (int32_t) i) { + const int32_t j = embd_batch_idxs[i]; + GGML_ASSERT(j >= 0 && (size_t) j < embd_batch_idxs.size()); + if (embd_nextn.has_data() && !cparams.embeddings_nextn_masked) { + for (size_t k = 0; k < n_embd_out; ++k) { + std::swap(embd_nextn.data[i*n_embd_out + k], embd_nextn.data[j*n_embd_out + k]); + } + } + for (auto & layer : embd_layer_inp) { + if (layer.has_data()) { + for (size_t k = 0; k < n_embd; ++k) { + std::swap(layer.data[i*n_embd + k], layer.data[j*n_embd + k]); + } + } + } + std::swap(embd_batch_idxs[i], embd_batch_idxs[j]); + } + } + embd_batch_idxs.clear(); } // diff --git a/src/llama-context.h b/src/llama-context.h index 35a7071eda..95b7226288 100644 --- a/src/llama-context.h +++ b/src/llama-context.h @@ -239,7 +239,7 @@ private: // async-copy enabled layer-input tensors (per cparams.output_layer_inp) // from backend into host-side embd_layer_inp buffers - void extract_layer_inputs(const llm_graph_result * res, size_t token_offset, size_t n_tokens); + bool extract_layer_inputs(const llm_graph_result * res, size_t token_offset, size_t n_tokens); // // graph @@ -312,6 +312,7 @@ private: // host buffers for output layer input embeddings, per layer // populated when cparams.output_layer_inp[il] is true std::vector> embd_layer_inp; + std::vector embd_batch_idxs; // extracted index -> original batch index struct sampling_info { // !samplers.empty() to check if any samplers are active