spec : fuse the DFlash encoder into the KV cache injection (#27310)
* dflash : fuse the encoder into the KV injection decode The encoder is a single fc + norm, but running it as a separate llama_encode forced a device-to-host round trip of its output before the injection decode could re-upload it, plus a second graph build per round. Fold the encoder into the decoder's embd branch and feed the target features directly to one llama_decode. Assisted-by: Claude Fable * nit * Apply batched suggestions from code review Co-authored-by: Ruixiang Wang <wangruixiang07@outlook.com> * Fix missing references from renaming --------- Co-authored-by: Sigbjørn Skjæret <sigbjorn.skjaeret@huggingface.co> Co-authored-by: Ruixiang Wang <wangruixiang07@outlook.com>
This commit is contained in:
co-authored by
Ruixiang Wang
Sigbjørn Skjæret
parent
2cdae802e4
commit
662a0b0121
@@ -1660,7 +1660,9 @@ int llama_context::decode(const llama_batch & batch_inp) {
|
||||
|
||||
const int64_t n_vocab = vocab.n_tokens();
|
||||
const bool mtp_embd = cparams.ctx_type == LLAMA_CONTEXT_TYPE_MTP && batch_inp.embd;
|
||||
const int64_t n_embd = mtp_embd ? hparams.n_embd_out() : hparams.n_embd_inp();
|
||||
// DFlash embd batches carry the fused target features at the encoder input width
|
||||
const bool dflash_embd = model.arch == LLM_ARCH_DFLASH && batch_inp.embd;
|
||||
const int64_t n_embd = mtp_embd ? hparams.n_embd_out() : dflash_embd ? hparams.n_embd_inp_enc() : hparams.n_embd_inp();
|
||||
|
||||
// when computing embeddings, all tokens are output
|
||||
const bool output_all = cparams.embeddings;
|
||||
|
||||
Reference in New Issue
Block a user