← Все главы

38 · Версия материала 6

Один раз заполните кэши, затем декодируйте по одному токену

Разберитесь, как отдельный KV-кэш (кэш ключей и значений) для каждого блока и один проверенный сеанс для модели и кэша позволяют один раз обработать промпт, затем согласованно декодировать по одному токену, а затем сравнить логиты последней позиции и решения при генерации с эталонными расчётами по полному префиксу.

Предскажите длины обоих кэшей до запуска

В примере используются L=2L=2 блока декодера, H=2H=2 головы внимания в каждом блоке, ширина модели D=4D=4, ширина головы dh=2d_h=2 и ёмкость контекста C=4C=4. Сначала состояние KV-кэшей всей модели пусто, а промпт состоит из токенов с идентификаторами [0,1][0,1].

До просмотра трассировки предскажите четыре факта:

  1. При заполнении кэшей по промпту обе его позиции проходят через оба блока, поэтому кэш каждого блока достигает логической длины 22 и формы [1,2,2,2][1,2,2,2].
  2. Следующим поступает токен с идентификатором 22 в абсолютной позиции 22.
  3. При успешном декодировании оба кэша переходят от длины 22 к длине 33: один блок не может отстать от другого.
  4. Выбранный токен EOS возвращается вызывающей стороне, но не проходит через декодер, если последующие логиты не нужны.

Последние логиты промпта равны [1.768374438,0.208825256,1.056205728,0.451857108,0.388467944][1.768374438,0.208825256,1.056205728,-0.451857108,0.388467944]. После токена 22 следующие логиты равны [0.032908910,0.679583624,1.408381841,0.525525421,0.588014095][0.032908910,-0.679583624,1.408381841,0.525525421,-0.588014095]. После обработки промпта и после декодирования токена 22 максимальная абсолютная разность до округления между результатами с KV-кэшем и по полному префиксу не превышает 2×10122\times10^{-12}; на экране она равна 0.0000000000000.000000000000.

В модели, восстановленной из контрольной точки, ёмкость контекста равна 22. Генерация начинается с промпта [0][0] и состояния генератора псевдослучайных чисел 0x9e3779b97f4a7c38. Оба пути выбирают [4,4][4,4]. При преобразовании сгенерированных ID токенов обратно в текст получается строка 44. Оба пути используют одинаковые псевдослучайные числа при выборе токенов, завершаются при одинаковом состоянии генератора псевдослучайных чисел и останавливаются на границе контекста. Заполнение по промпту доводит длину кэша до 11; первый токен 44 проходит через декодер и увеличивает длину кэша до 22; второй токен 44 выбирается из полученных логитов и возвращается, после чего ограничение контекста останавливает генерацию до его декодирования. Если токен 44 задан как EOS, оба пути останавливаются после выбора [4][4] и не выполняют ни одного последующего вызова декодера для этого токена.

Считайте значения оценок внимания, а не полное время работы

Зафиксируем размер пакета, число слоёв и голов и ширину головы. При конечной длине сохранённого префикса TT два суммарных числа значений оценок внимания растут как

t=1Tt2Θ(T3),t=1TtΘ(T2).\sum_{t=1}^{T}t^2\in\Theta(T^3),\quad \sum_{t=1}^{T}t\in\Theta(T^2)\,.

При сохранённой длине tt повторный расчёт по полному префиксу заново строит плотную каузальную матрицу из t2t^2 оценок для каждой фиксированной комбинации элемента пакета, блока и головы. Шаг с KV-кэшем формирует только строку запроса для новой позиции: tt оценок для всех сохранённых ключей. Сумма таких величин на каждом шаге от t=1t=1 до TT даёт два указанных класса роста.

Область этого сравнения намеренно узка. Время расчёта внимания не становится постоянным: запрос для новой позиции по-прежнему просматривает префикс, длина которого растёт вместе с tt. Формула не учитывает вычисления проекций или MLP, обращения к памяти, выделение памяти, полное время работы или измеренное ускорение.

В заданном примере размер пакета 11, 22 блока и 22 головы дают фиксированный множитель 44. Последовательная обработка промпта и декодирование вычисляют 4(1+2+3)=244(1+2+3)=24 значения оценок внимания с KV-кэшем. Два независимых эталонных вызова при длинах 22 и 33 вычисляют 4(22+32)=524(2^2+3^2)=52 значения.

Последовательности вызовов в примере намеренно различаются. Путь с KV-кэшем обрабатывает префиксы длины 11, 22 и 33, а глава выполняет лишь две независимые проверки по полному префиксу — при длинах 22 и 33. Поэтому 2424 и 5252 — измеренные числа для фактически выполненных вызовов, а не две асимптотические суммы, вычисленные при T=3T=3.

Не смешивайте сохранённую и конечную длины

  • tt — текущая длина сохранённого префикса. Столько же ключей читает запрос для одной новой позиции при расчёте с KV-кэшем.
  • TT — конечная длина сохранённого префикса, охваченная сравнением.
  • t2t^2 — матрица оценок по полному префиксу, заново построенная при длине tt для каждой фиксированной комбинации элемента пакета, блока и головы.
  • t=1T\sum_{t=1}^{T} суммирует работу по вычислению оценок внимания от сохранённой длины 11 до длины TT.
  • Θ(T3)\Theta(T^3) — класс роста повторно вычисляемых матриц оценок по полному префиксу, когда опущенные множители остаются фиксированными.
  • Θ(T2)\Theta(T^2) — класс роста строк оценок запроса для новой позиции с KV-кэшем при тех же фиксированных множителях.

Сам кэш по-прежнему хранит в каждом блоке один логический префикс K/V формы [B,H,t,dh][B,H,t,d_h]. Одинаковая форма не делает два кэша взаимозаменяемыми: их строки получены на разной глубине стека декодера. DecoderKvCache записывает конфигурацию декодера, идентичность узла каждого параметра и версию его значения. Каждый вложенный кэш блока отдельно записывает геометрию внимания, конфигурацию RoPE и сведения о четырёх параметрах внимания. Эти записи служат данными для проверки совместимости, но сами по себе не создают действующую связь между моделью и кэшем.

Вызов cache.bind(&model) проверяет эти данные и создаёт DecoderKvSession для одной конкретной пары модели и кэша. Затем сеанс удерживает значения параметров доступными только для чтения. Одних формы и совпадения узлов недостаточно: обновление на месте может сохранить узлы, но изменить версии значений, а строки, вычисленные с двумя версиями параметров, нельзя объединять в один логический префикс.

От стека каузальных слоёв к обработке промпта и последовательному декодированию

Каузальный декодер Transformer может генерировать по одному токену, каждый раз повторно обрабатывая весь известный префикс. Однако такой интерфейс без состояния при каждом следующем вызове заново строит матрицы оценок внимания и проекции ключей и значений для прежних позиций.

Attention Is All You Need описывает каузальный стек. Васвани и соавторы описывают стек авторегрессионных слоёв декодера Transformer, в котором маскированное самовнимание не позволяет позиции обращаться к последующим позициям; их архитектура «кодировщик — декодер» также содержит перекрёстное внимание, но не задаёт API для KV-кэша.

Fast Transformer Decoding: One Write-Head is All You Need описывает явную передачу предыдущего состояния. Инкрементальное самовнимание Шейзира получает прежние тензоры ключей и значений, добавляет текущие спроецированные ключ и значение и возвращает обновлённое состояние; вклад статьи состоит в многозапросном внимании, а не в заявлении об изобретении KV-кэширования.

Efficient Memory Management for Large Language Model Serving with PagedAttention описывает современные системы генерации с помощью LLM. Квон и соавторы отделяют обработку промпта от последовательной генерации, описывают повторное использование сохранённых ключей и значений при вычислении только новой пары на последующих итерациях и учитывают состояние KV-кэша во всех слоях и головах Transformer.

При инкрементальном декодировании прежние тензоры ключей и значений стали явным состоянием. Позднее в системах обслуживания LLM обработку промпта отделили от последовательной генерации, сохраняя состояние ключей и значений во всех слоях и головах декодера.

При генерации с KV-кэшем модель один раз заполняет отдельный кэш каждого блока декодера, а затем обновляет состояние всех блоков лишь тогда, когда нужны следующие логиты. В заданных примерах логиты последней позиции совпадают с эталонными расчётами по полному префиксу в пределах допуска, а после восстановления контрольной точки совпадают выбранные токены, псевдослучайные числа, использованные при выборе токенов, конечное состояние генератора псевдослучайных чисел и причина остановки.

Программа из этой главы связывает этот путь с подсчётом работы внимания. В четырёх комбинациях «элемент пакета — слой — голова» последовательные строки с KV-кэшем при сохранённых длинах [1,2,3][1,2,3] содержат 2424 значения оценок внимания. Два вызова по полному префиксу при длинах [2,3][2,3] содержат 5252 значения, поэтому в этой заданной последовательности вызовов не вычисляются 2828 значений. Это числа элементов тензоров, а не измерение страничной организации памяти или времени работы.

Чтобы повторно используемое состояние оставалось согласованным, реализация из этой главы при создании сеанса проверяет все отношения между декодером и кэшем, сохраняет доступ ко всем значениям параметров только для чтения на всё время сеанса, подготавливает изменения каждого блока до записи любой строки в кэш и явно обрабатывает сброс, счётчики, типизированные ошибки и последовательную обработку промпта. Это требования данной реализации, а не правила, установленные цитируемыми статьями.

Страничная организация, совместное использование, вытеснение и управление памятью в промышленных системах остаются за пределами главы. Используемый здесь кэш с непрерывным размещением данных позволяет отдельно разобрать требования к корректности перед изучением этих задач обслуживания моделей.

Измерить тензоры оценок для путей с KV-кэшем и по полному префиксу при двух последовательностях вызовов из примера rust/demos/ch38-cached-generation/src/lib.rs#historical-cache-contrast
/// Measures complete-prefix replay against retained model-wide KV state.
pub fn historical_cache_contrast(
    config: DecoderModelConfig,
    cached_retained_lengths: &[usize],
    complete_prefix_lengths: &[usize],
    measured_cached_scores: usize,
    measured_complete_prefix_scores: usize,
) -> Result<HistoricalCacheContrast, FixtureError> {
    require(
        !cached_retained_lengths.is_empty() && !complete_prefix_lengths.is_empty(),
        "history evidence needs cached and complete-prefix calls",
    )?;
    let batch_layer_head_lanes = config
        .layers()
        .checked_mul(config.heads())
        .ok_or(FixtureError::Invariant("history lane count overflowed"))?;
    let cached_attention_score_values =
        cached_retained_lengths
            .iter()
            .try_fold(0usize, |total, &length| {
                let scores =
                    batch_layer_head_lanes
                        .checked_mul(length)
                        .ok_or(FixtureError::Invariant(
                            "cached history score count overflowed",
                        ))?;
                total.checked_add(scores).ok_or(FixtureError::Invariant(
                    "cached history score total overflowed",
                ))
            })?;
    let complete_prefix_attention_score_values =
        complete_prefix_attention_score_values(config, complete_prefix_lengths)?;
    require(
        cached_attention_score_values == measured_cached_scores
            && complete_prefix_attention_score_values == measured_complete_prefix_scores,
        "history score contrast disagrees with measured work",
    )?;
    let avoided_attention_score_values = complete_prefix_attention_score_values
        .checked_sub(cached_attention_score_values)
        .ok_or(FixtureError::Invariant(
            "cached history exceeds complete-prefix reference",
        ))?;
    Ok(HistoricalCacheContrast {
        batch_layer_head_lanes,
        cached_retained_lengths: cached_retained_lengths.to_vec(),
        cached_attention_score_values,
        complete_prefix_lengths: complete_prefix_lengths.to_vec(),
        complete_prefix_attention_score_values,
        avoided_attention_score_values,
    })
}

Подготовьте каждый блок, затем согласованно запишите изменения во все кэши

Операция одного слоя из главы 37 теперь разделена на два этапа. Сначала она вычисляет выход для новой позиции и две добавляемые строки: повёрнутый ключ и значение без поворота. Логическое состояние кэша при этом не меняется. Обычный публичный вызов сразу записывает подготовленную пару. Код для всей модели вместо этого сохраняет по одному подготовленному результату для каждого блока и записывает строки только после успешного вычисления всех блоков и итоговой проекции на словарь.

Подготовить добавляемую пару K/V и записать её только после успешного вычисления всей строки декодера rust/crates/llm-from-scratch/src/attention/incremental.rs#incremental-attention
/// A fallible incremental-attention buffer or tensor stage.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum IncrementalAttentionStage {
    Scores,
    HeadOutputs,
    HeadOutputLeaf,
}

impl fmt::Display for IncrementalAttentionStage {
    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
        formatter.write_str(match self {
            Self::Scores => "attention scores",
            Self::HeadOutputs => "weighted head outputs",
            Self::HeadOutputLeaf => "head-output tensor",
        })
    }
}

/// A rejected single-token input, cache pairing, or incremental forward stage.
#[derive(Clone, Debug, PartialEq)]
pub enum IncrementalAttentionError {
    InputRank {
        rank: usize,
    },
    SingleTokenRequired {
        tokens: usize,
    },
    InputBatchMismatch {
        cache: usize,
        input: usize,
    },
    InputWidthMismatch {
        expected: usize,
        actual: usize,
    },
    CacheModelWidthMismatch {
        layer: usize,
        cache: usize,
    },
    CacheHeadCountMismatch {
        layer: usize,
        cache: usize,
    },
    CacheHeadWidthMismatch {
        layer: usize,
        cache: usize,
    },
    CacheLayerMismatch,
    CacheLayerRevisionMismatch {
        parameter: usize,
        cache: u64,
        layer: u64,
    },
    CacheRopeMismatch {
        cache_feature_width: usize,
        layer_feature_width: usize,
        cache_max_positions: usize,
        layer_max_positions: usize,
        cache_base: f64,
        layer_base: f64,
    },
    BatchHeadOverflow {
        batch: usize,
        heads: usize,
    },
    BufferSizeOverflow {
        stage: IncrementalAttentionStage,
    },
    BufferAllocationFailed {
        stage: IncrementalAttentionStage,
        elements: usize,
    },
    Cache(LayerKvCacheError),
    QkvProjection(QkvError),
    HeadLayout {
        input: MultiHeadInput,
        source: HeadLayoutError,
    },
    Rotary {
        input: MultiHeadInput,
        source: RopeError,
    },
    Probability(ProbabilityError),
    Tensor {
        stage: IncrementalAttentionStage,
        source: TensorError,
    },
    Autodiff {
        stage: IncrementalAttentionStage,
        source: TensorAutodiffError,
    },
    MergeLayout(HeadLayoutError),
    OutputProjection(LinearError),
}

impl fmt::Display for IncrementalAttentionError {
    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
        match self {
            Self::InputRank { rank } => write!(
                formatter,
                "incremental attention input must have rank three [batch, 1, model_width], got rank {rank}"
            ),
            Self::SingleTokenRequired { tokens } => write!(
                formatter,
                "incremental attention needs exactly one new token, got {tokens}"
            ),
            Self::InputBatchMismatch { cache, input } => write!(
                formatter,
                "incremental attention input batch {input} must match cache batch {cache}"
            ),
            Self::InputWidthMismatch { expected, actual } => write!(
                formatter,
                "incremental attention input width must equal model width {expected}, got {actual}"
            ),
            Self::CacheModelWidthMismatch { layer, cache } => write!(
                formatter,
                "KV cache model width {cache} must match attention layer width {layer}"
            ),
            Self::CacheHeadCountMismatch { layer, cache } => write!(
                formatter,
                "KV cache head count {cache} must match attention layer head count {layer}"
            ),
            Self::CacheHeadWidthMismatch { layer, cache } => write!(
                formatter,
                "KV cache head width {cache} must match attention layer head width {layer}"
            ),
            Self::CacheLayerMismatch => formatter
                .write_str("KV cache parameter identity does not match this attention layer"),
            Self::CacheLayerRevisionMismatch {
                parameter,
                cache,
                layer,
            } => write!(
                formatter,
                "KV cache parameter revision {cache} at stable index {parameter} does not match layer revision {layer}"
            ),
            Self::CacheRopeMismatch {
                cache_feature_width,
                layer_feature_width,
                cache_max_positions,
                layer_max_positions,
                cache_base,
                layer_base,
            } => write!(
                formatter,
                "KV cache RoPE configuration ({cache_feature_width} features, {cache_max_positions} positions, base {cache_base:?}) does not match layer configuration ({layer_feature_width} features, {layer_max_positions} positions, base {layer_base:?})"
            ),
            Self::BatchHeadOverflow { batch, heads } => write!(
                formatter,
                "incremental attention lane count overflows for batch {batch} and {heads} heads"
            ),
            Self::BufferSizeOverflow { stage } => {
                write!(formatter, "incremental {stage} element count overflows")
            }
            Self::BufferAllocationFailed { stage, elements } => write!(
                formatter,
                "cannot allocate incremental {stage} buffer for {elements} f64 values"
            ),
            Self::Cache(source) => source.fmt(formatter),
            Self::QkvProjection(source) => {
                write!(formatter, "incremental Q/K/V projection: {source}")
            }
            Self::HeadLayout { input, source } => {
                write!(formatter, "incremental {input} head layout: {source}")
            }
            Self::Rotary { input, source } => {
                write!(formatter, "incremental {input} RoPE: {source}")
            }
            Self::Probability(source) => write!(formatter, "incremental softmax: {source}"),
            Self::Tensor { stage, source } => {
                write!(formatter, "incremental {stage}: {source}")
            }
            Self::Autodiff { stage, source } => {
                write!(formatter, "incremental {stage}: {source}")
            }
            Self::MergeLayout(source) => {
                write!(formatter, "incremental head output merge: {source}")
            }
            Self::OutputProjection(source) => {
                write!(formatter, "incremental output projection: {source}")
            }
        }
    }
}

impl Error for IncrementalAttentionError {
    fn source(&self) -> Option<&(dyn Error + 'static)> {
        match self {
            Self::Cache(source) => Some(source),
            Self::QkvProjection(source) => Some(source),
            Self::HeadLayout { source, .. } => Some(source),
            Self::Rotary { source, .. } => Some(source),
            Self::Probability(source) => Some(source),
            Self::Tensor { source, .. } => Some(source),
            Self::Autodiff { source, .. } => Some(source),
            Self::MergeLayout(source) => Some(source),
            Self::OutputProjection(source) => Some(source),
            _ => None,
        }
    }
}

impl From<LayerKvCacheError> for IncrementalAttentionError {
    fn from(source: LayerKvCacheError) -> Self {
        Self::Cache(source)
    }
}

/// Exact row counts for comparing full-prefix and cached projections.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct IncrementalAttentionWork {
    position: usize,
    full_prefix_rows_per_projection: usize,
    incremental_rows_per_projection: usize,
    reused_key_value_rows: usize,
}

impl IncrementalAttentionWork {
    pub const fn position(&self) -> usize {
        self.position
    }

    pub const fn full_prefix_rows_per_projection(&self) -> usize {
        self.full_prefix_rows_per_projection
    }

    pub const fn incremental_rows_per_projection(&self) -> usize {
        self.incremental_rows_per_projection
    }

    pub const fn reused_key_value_rows(&self) -> usize {
        self.reused_key_value_rows
    }
}

/// Inspectable graph-free evidence from one committed cache append.
#[derive(Clone, Debug)]
pub struct IncrementalAttentionForward {
    projected_query_heads: TensorValue,
    projected_key_heads: TensorValue,
    projected_value_heads: TensorValue,
    rotated_query_heads: TensorValue,
    rotated_key_heads: TensorValue,
    attention_weights: Tensor,
    head_outputs: TensorValue,
    merged: TensorValue,
    output: TensorValue,
    work: IncrementalAttentionWork,
    cache_len: usize,
}

impl IncrementalAttentionForward {
    pub fn projected_query_heads(&self) -> &TensorValue {
        &self.projected_query_heads
    }

    pub fn projected_key_heads(&self) -> &TensorValue {
        &self.projected_key_heads
    }

    pub fn projected_value_heads(&self) -> &TensorValue {
        &self.projected_value_heads
    }

    pub fn rotated_query_heads(&self) -> &TensorValue {
        &self.rotated_query_heads
    }

    pub fn rotated_key_heads(&self) -> &TensorValue {
        &self.rotated_key_heads
    }

    /// Returns `[batch, heads, 1, cache_len]` probabilities for the new query.
    pub fn attention_weights(&self) -> &Tensor {
        &self.attention_weights
    }

    pub fn head_outputs(&self) -> &TensorValue {
        &self.head_outputs
    }

    pub fn merged(&self) -> &TensorValue {
        &self.merged
    }

    pub fn output(&self) -> &TensorValue {
        &self.output
    }

    pub const fn work(&self) -> IncrementalAttentionWork {
        self.work
    }

    pub const fn cache_len(&self) -> usize {
        self.cache_len
    }

    pub fn into_output(self) -> TensorValue {
        self.output
    }
}

/// A crate-sealed incremental result whose candidate K/V row is not committed.
///
/// Chapter 38 prepares one ticket per decoder block, completes the later blocks
/// and tied vocabulary head, verifies every ticket still targets its original
/// cache, and only then commits the complete stack.
pub(crate) struct PreparedIncrementalAttention {
    forward: IncrementalAttentionForward,
    candidate_key: Tensor,
    candidate_value: Tensor,
    expected_len: usize,
    key_storage: *const f64,
    value_storage: *const f64,
}

impl PreparedIncrementalAttention {
    pub(crate) fn output(&self) -> &TensorValue {
        self.forward.output()
    }

    pub(crate) const fn cache_len(&self) -> usize {
        self.forward.cache_len()
    }

    pub(crate) fn attention_score_values(&self) -> usize {
        self.forward.attention_weights().len()
    }

    pub(crate) fn matches_cache(&self, cache: &LayerKvCache) -> bool {
        self.expected_len == cache.len()
            && std::ptr::eq(self.key_storage, cache.key_storage().as_ptr())
            && std::ptr::eq(self.value_storage, cache.value_storage().as_ptr())
    }

    pub(crate) fn commit(self, cache: &mut LayerKvCache) -> IncrementalAttentionForward {
        debug_assert!(self.matches_cache(cache));
        cache.append_prevalidated(&self.candidate_key, &self.candidate_value);
        self.forward
    }
}

impl MultiHeadAttention {
    /// Projects one new row, attends over retained K/V rows, and commits one append.
    ///
    /// The cache is changed only after projection, RoPE, stable softmax, value
    /// mixing, merge, and output projection all succeed.
    pub fn forward_incremental(
        &self,
        input: &TensorValue,
        cache: &mut LayerKvCache,
    ) -> Result<IncrementalAttentionForward, IncrementalAttentionError> {
        let prepared = self.prepare_incremental(input, cache)?;
        Ok(prepared.commit(cache))
    }

    /// Computes one incremental row without changing the layer cache.
    pub(crate) fn prepare_incremental(
        &self,
        input: &TensorValue,
        cache: &LayerKvCache,
    ) -> Result<PreparedIncrementalAttention, IncrementalAttentionError> {
        self.validate_incremental_request(input, cache)?;
        self.prepare_incremental_bound(input, cache)
    }

    fn validate_incremental_request(
        &self,
        input: &TensorValue,
        cache: &LayerKvCache,
    ) -> Result<(), IncrementalAttentionError> {
        let shape = input.shape();
        if shape.len() != 3 {
            return Err(IncrementalAttentionError::InputRank { rank: shape.len() });
        }
        if shape[1] != 1 {
            return Err(IncrementalAttentionError::SingleTokenRequired { tokens: shape[1] });
        }
        if shape[0] != cache.batch_size() {
            return Err(IncrementalAttentionError::InputBatchMismatch {
                cache: cache.batch_size(),
                input: shape[0],
            });
        }
        if shape[2] != self.model_width() {
            return Err(IncrementalAttentionError::InputWidthMismatch {
                expected: self.model_width(),
                actual: shape[2],
            });
        }
        self.validate_incremental_cache_binding(cache)?;
        if cache.is_full() {
            return Err(LayerKvCacheError::Full {
                capacity: cache.capacity(),
            }
            .into());
        }
        Ok(())
    }

    /// Checks the persistent relationship between one layer and one cache.
    ///
    /// A model-wide session calls this once while binding its complete cache.
    /// The standalone entry calls it for every arbitrary layer/cache pairing.
    pub(crate) fn validate_incremental_cache_binding(
        &self,
        cache: &LayerKvCache,
    ) -> Result<(), IncrementalAttentionError> {
        if cache.model_width() != self.model_width() {
            return Err(IncrementalAttentionError::CacheModelWidthMismatch {
                layer: self.model_width(),
                cache: cache.model_width(),
            });
        }
        if cache.heads() != self.heads() {
            return Err(IncrementalAttentionError::CacheHeadCountMismatch {
                layer: self.heads(),
                cache: cache.heads(),
            });
        }
        if cache.head_width() != self.head_width() {
            return Err(IncrementalAttentionError::CacheHeadWidthMismatch {
                layer: self.head_width(),
                cache: cache.head_width(),
            });
        }
        if !self
            .parameters()
            .iter()
            .zip(&cache.parameter_bindings)
            .all(|(parameter, cached)| cached.node_matches(parameter.tensor()))
        {
            return Err(IncrementalAttentionError::CacheLayerMismatch);
        }
        if let Some((parameter, (cached, layer))) = cache
            .parameter_bindings
            .iter()
            .zip(self.parameters())
            .enumerate()
            .find(|(_, (cached, parameter))| !cached.revision_matches(parameter.tensor()))
        {
            return Err(IncrementalAttentionError::CacheLayerRevisionMismatch {
                parameter,
                cache: cached.revision(),
                layer: layer.tensor().value_revision(),
            });
        }
        if cache.rope_feature_width != self.rope().feature_width()
            || cache.rope_max_positions != self.rope().max_positions()
            || cache.rope_base_bits != self.rope().base().to_bits()
        {
            return Err(IncrementalAttentionError::CacheRopeMismatch {
                cache_feature_width: cache.rope_feature_width,
                layer_feature_width: self.rope().feature_width(),
                cache_max_positions: cache.rope_max_positions,
                layer_max_positions: self.rope().max_positions(),
                cache_base: f64::from_bits(cache.rope_base_bits),
                layer_base: self.rope().base(),
            });
        }
        Ok(())
    }

    /// Prepares one row after its crate-private caller establishes every precondition.
    ///
    /// A model-wide bind establishes the persistent layer/cache relationship.
    /// The current session operation separately guarantees the one-row input
    /// shape and remaining capacity. The caller must preserve that exact
    /// layer/cache pairing until every prepared row either commits or is
    /// discarded. Keeping this entry crate-private lets Chapter 38 reuse the one
    /// attention implementation without creating an unchecked public path.
    pub(crate) fn prepare_incremental_bound(
        &self,
        input: &TensorValue,
        cache: &LayerKvCache,
    ) -> Result<PreparedIncrementalAttention, IncrementalAttentionError> {
        no_grad(|| {
            let position = cache.len();
            let projected = self
                .qkv()
                .forward(input)
                .map_err(IncrementalAttentionError::QkvProjection)?;
            let projected_query_heads =
                split_heads(projected.query(), self.heads()).map_err(|source| {
                    IncrementalAttentionError::HeadLayout {
                        input: MultiHeadInput::Query,
                        source,
                    }
                })?;
            let projected_key_heads =
                split_heads(projected.key(), self.heads()).map_err(|source| {
                    IncrementalAttentionError::HeadLayout {
                        input: MultiHeadInput::Key,
                        source,
                    }
                })?;
            let projected_value_heads =
                split_heads(projected.value(), self.heads()).map_err(|source| {
                    IncrementalAttentionError::HeadLayout {
                        input: MultiHeadInput::Value,
                        source,
                    }
                })?;
            let rotated_query_heads = self
                .rope()
                .rotate(&projected_query_heads, position)
                .map_err(|source| IncrementalAttentionError::Rotary {
                    input: MultiHeadInput::Query,
                    source,
                })?;
            let rotated_key_heads =
                self.rope()
                    .rotate(&projected_key_heads, position)
                    .map_err(|source| IncrementalAttentionError::Rotary {
                        input: MultiHeadInput::Key,
                        source,
                    })?;
            let candidate_key = rotated_key_heads.value_snapshot();
            let candidate_value = projected_value_heads.value_snapshot();
            let (attention_weights, head_output_tensor) = incremental_mixture(
                &rotated_query_heads.value(),
                &candidate_key,
                &candidate_value,
                cache,
            )?;
            let head_outputs = TensorValue::constant(head_output_tensor).map_err(|source| {
                IncrementalAttentionError::Autodiff {
                    stage: IncrementalAttentionStage::HeadOutputLeaf,
                    source,
                }
            })?;
            let merged =
                merge_heads(&head_outputs).map_err(IncrementalAttentionError::MergeLayout)?;
            let output = self
                .output_projection()
                .forward(&merged)
                .map_err(IncrementalAttentionError::OutputProjection)?;
            let cache_len = position + 1;
            let result = IncrementalAttentionForward {
                projected_query_heads,
                projected_key_heads,
                projected_value_heads,
                rotated_query_heads,
                rotated_key_heads,
                attention_weights,
                head_outputs,
                merged,
                output,
                work: IncrementalAttentionWork {
                    position,
                    full_prefix_rows_per_projection: cache_len,
                    incremental_rows_per_projection: 1,
                    reused_key_value_rows: position,
                },
                cache_len,
            };
            cache.validate_append(&candidate_key, &candidate_value)?;
            Ok(PreparedIncrementalAttention {
                forward: result,
                candidate_key,
                candidate_value,
                expected_len: position,
                key_storage: cache.key_storage().as_ptr(),
                value_storage: cache.value_storage().as_ptr(),
            })
        })
    }
}

fn incremental_mixture(
    query: &Tensor,
    candidate_key: &Tensor,
    candidate_value: &Tensor,
    cache: &LayerKvCache,
) -> Result<(Tensor, Tensor), IncrementalAttentionError> {
    let lanes = cache.batch_size().checked_mul(cache.heads()).ok_or(
        IncrementalAttentionError::BatchHeadOverflow {
            batch: cache.batch_size(),
            heads: cache.heads(),
        },
    )?;
    let prefix = cache.len() + 1;
    let score_elements =
        lanes
            .checked_mul(prefix)
            .ok_or(IncrementalAttentionError::BufferSizeOverflow {
                stage: IncrementalAttentionStage::Scores,
            })?;
    let mut scores = reserved_buffer(score_elements, IncrementalAttentionStage::Scores)?;
    let scale = 1.0 / (cache.head_width() as f64).sqrt();

    for batch in 0..cache.batch_size() {
        for head in 0..cache.heads() {
            let lane = batch * cache.heads() + head;
            let query_start = lane * cache.head_width();
            for position in 0..prefix {
                let key_start = if position == cache.len() {
                    lane * cache.head_width()
                } else {
                    (lane * cache.capacity() + position) * cache.head_width()
                };
                let key_values = if position == cache.len() {
                    candidate_key.as_slice()
                } else {
                    cache.key_storage()
                };
                let mut dot = 0.0;
                for feature in 0..cache.head_width() {
                    dot +=
                        query.as_slice()[query_start + feature] * key_values[key_start + feature];
                }
                scores[lane * prefix + position] = dot * scale;
            }
        }
    }

    let score_tensor = Tensor::from_vec(vec![cache.batch_size(), cache.heads(), 1, prefix], scores)
        .map_err(|source| IncrementalAttentionError::Tensor {
            stage: IncrementalAttentionStage::Scores,
            source,
        })?;
    let weights =
        softmax(&score_tensor.view(), 3).map_err(IncrementalAttentionError::Probability)?;
    let output_elements = lanes.checked_mul(cache.head_width()).ok_or(
        IncrementalAttentionError::BufferSizeOverflow {
            stage: IncrementalAttentionStage::HeadOutputs,
        },
    )?;
    let mut outputs = reserved_buffer(output_elements, IncrementalAttentionStage::HeadOutputs)?;
    for batch in 0..cache.batch_size() {
        for head in 0..cache.heads() {
            let lane = batch * cache.heads() + head;
            for feature in 0..cache.head_width() {
                let mut mixture = 0.0;
                for position in 0..prefix {
                    let value_start = if position == cache.len() {
                        lane * cache.head_width()
                    } else {
                        (lane * cache.capacity() + position) * cache.head_width()
                    };
                    let value_values = if position == cache.len() {
                        candidate_value.as_slice()
                    } else {
                        cache.value_storage()
                    };
                    mixture += weights.as_slice()[lane * prefix + position]
                        * value_values[value_start + feature];
                }
                outputs[lane * cache.head_width() + feature] = mixture;
            }
        }
    }
    let outputs = Tensor::from_vec(
        vec![cache.batch_size(), cache.heads(), 1, cache.head_width()],
        outputs,
    )
    .map_err(|source| IncrementalAttentionError::Tensor {
        stage: IncrementalAttentionStage::HeadOutputs,
        source,
    })?;
    Ok((weights, outputs))
}

fn reserved_buffer(
    elements: usize,
    stage: IncrementalAttentionStage,
) -> Result<Vec<f64>, IncrementalAttentionError> {
    let mut values = Vec::new();
    values
        .try_reserve_exact(elements)
        .map_err(|_| IncrementalAttentionError::BufferAllocationFailed { stage, elements })?;
    values.resize(elements, 0.0);
    Ok(values)
}

DecoderKvCache::new принимает конкретный декодер. Он выделяет по одному LayerKvCache фиксированной ёмкости для каждого блока и записывает точную конфигурацию декодера. Для каждого параметра модели кэш запоминает идентичность узла и текущую версию значения. Кэш владеет повторно используемым хранилищем K/V и данными для проверки совместимости, но не копирует веса и не сохраняет заимствование декодера.

Перед генерацией вызов cache.bind(&model) проверяет эти данные и возвращает DecoderKvSession для одной конкретной пары модели и кэша. При создании сеанса проверяются конфигурация декодера; число параметров, порядок идентичностей их узлов и версий значений; число кэшей блоков и одинаковая логическая длина этих кэшей; а также размер пакета, ёмкость, геометрия внимания, параметры внимания и конфигурация RoPE каждого кэша блока. Слово «один раз» здесь означает один раз для каждого нового сеанса, а не один раз за всё время существования кэша.

После успешных проверок сеанс сохраняет активные неизменяемые заимствования значений всех параметров. Это доступ к исходным значениям только для чтения, а не копии весов. Операции декодера продолжают читать эти значения, но AdamW не может получить исключительный доступ для их изменения. Так значение параметра не может измениться между проверкой его версии и последующим использованием.

Кроме того, только сеанс заимствует кэш с правом изменения. Его внутренний переход для одной строки и открытый метод reset изменяют длины всех блоков согласованно, поэтому проверенное при создании сеанса равенство длин сохраняется.

Методы DecoderKvSession::prefill и DecoderKvSession::decode не принимают модель отдельным аргументом, потому что используют модель, уже связанную с сеансом. Это не означает, что они работают без модели или без проверок. При каждом вызове по-прежнему проверяются условия, которые могут измениться: корректность промпта, текущая фаза, допустимость токена, свободная ёмкость, арифметика счётчиков и то, что каждое подготовленное изменение всё ещё относится к тому же хранилищу K/V при ожидаемой логической длине. Методы не обходят заново всю модель и кэши перед каждой операцией и не повторяют для каждой строки проверки неизменных условий совместимости. prefill принимает непустой проверенный промпт, только пока состояние пусто. В этой учебной реализации строки промпта последовательно проходят через один и тот же путь для одной строки. Так переход состояния остаётся наглядным; это не оптимизированная параллельная реализация обработки промпта.

Для каждой строки декодер вычисляет эмбеддинг, затем по порядку применяет в каждом блоке подслои внимания и сети прямого распространения с предварительной нормализацией и остаточными связями, после чего выполняет итоговую RMSNorm и проекцию на словарь с общими весами матрицы эмбеддингов. Каждый блок использует одну общую реализацию вычисления внимания из главы 37, чтобы подготовить две строки — повёрнутый ключ и значение без поворота, — не меняя кэш. Публичный путь из главы 37 выполняет полный набор проверок для самостоятельного вызова слоя и кэша. Сеанс вызывает доступный только коду этого крейта внутренний путь после того, как при создании сеанса проверены условия совместимости, которые остаются неизменными в течение этого сеанса. Это не второй алгоритм внимания и не открытый путь в обход проверок. Кэши блоков, общая длина, счётчики фаз и число значений оценок внимания обновляются вместе только после успешного вычисления оставшихся блоков, итоговой нормализации и проекции на словарь.

Связать один декодер с общим состоянием его кэшей и согласованно обновлять каждый блок rust/crates/llm-from-scratch/src/generation/kv_cache.rs#decoder-kv-cache
/// A checked model-wide work counter that could not be represented as `usize`.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DecoderKvCacheCounter {
    TokenForwards,
    PrefillTokens,
    DecodeTokens,
    CacheAppends,
    QkvProjectionRows,
    AttentionScoreValues,
    CompletePrefixAttentionScoreValues,
}

impl fmt::Display for DecoderKvCacheCounter {
    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
        formatter.write_str(match self {
            Self::TokenForwards => "token forwards",
            Self::PrefillTokens => "prefill tokens",
            Self::DecodeTokens => "decode tokens",
            Self::CacheAppends => "cache appends",
            Self::QkvProjectionRows => "Q/K/V projection rows",
            Self::AttentionScoreValues => "cached attention score values",
            Self::CompletePrefixAttentionScoreValues => "complete-prefix attention score values",
        })
    }
}

/// A graph-free cached-decoder stage that rejected a request or computation.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CachedDecoderStage {
    TiedWeightTranspose,
    TiedVocabularyProjection,
}

impl fmt::Display for CachedDecoderStage {
    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
        formatter.write_str(match self {
            Self::TiedWeightTranspose => "transpose tied embedding weight",
            Self::TiedVocabularyProjection => "project tied vocabulary logits",
        })
    }
}

/// A rejected model/cache pairing, phase transition, token, or decoder stage.
#[derive(Clone, Debug, PartialEq)]
pub enum DecoderKvCacheError {
    LayerAllocationFailed {
        layers: usize,
    },
    ParameterAllocationFailed {
        parameters: usize,
    },
    LayerCache {
        layer: usize,
        source: LayerKvCacheError,
    },
    ModelConfigMismatch,
    ModelParameterCountMismatch {
        cache: usize,
        model: usize,
    },
    ModelParameterMismatch {
        index: usize,
    },
    ModelParameterRevisionMismatch {
        index: usize,
        cache: u64,
        model: u64,
    },
    LayerCountMismatch {
        cache: usize,
        model: usize,
    },
    LayerBatchSizeMismatch {
        layer: usize,
        expected: usize,
        actual: usize,
    },
    LayerCapacityMismatch {
        layer: usize,
        expected: usize,
        actual: usize,
    },
    EmptyPrompt,
    PromptTooLong {
        tokens: usize,
        capacity: usize,
    },
    PromptTokenOutOfBounds {
        position: usize,
        token_id: u32,
        vocabulary_size: usize,
    },
    PrefillRequiresEmpty {
        len: usize,
    },
    DecodeRequiresPrefill,
    DecodeTokenOutOfBounds {
        token_id: u32,
        vocabulary_size: usize,
    },
    Full {
        capacity: usize,
    },
    LayerLengthInvariant {
        layer: usize,
        expected: usize,
        actual: usize,
    },
    PreparedCacheChanged {
        layer: usize,
    },
    WorkOverflow {
        counter: DecoderKvCacheCounter,
    },
    Embedding(EmbeddingError),
    AttentionNorm {
        layer: usize,
        source: RmsNormError,
    },
    IncrementalAttention {
        layer: usize,
        source: IncrementalAttentionError,
    },
    AttentionResidual {
        layer: usize,
        source: ResidualError,
    },
    FeedForwardNorm {
        layer: usize,
        source: RmsNormError,
    },
    FeedForward {
        layer: usize,
        source: SwiGluError,
    },
    FeedForwardResidual {
        layer: usize,
        source: ResidualError,
    },
    FinalNorm(RmsNormError),
    Autodiff {
        stage: CachedDecoderStage,
        source: TensorAutodiffError,
    },
}

impl fmt::Display for DecoderKvCacheError {
    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
        match self {
            Self::LayerAllocationFailed { layers } => {
                write!(formatter, "cannot allocate {layers} decoder layer caches")
            }
            Self::ParameterAllocationFailed { parameters } => write!(
                formatter,
                "cannot retain {parameters} decoder parameter bindings"
            ),
            Self::LayerCache { layer, source } => {
                write!(formatter, "decoder layer {layer} cache: {source}")
            }
            Self::ModelConfigMismatch => formatter
                .write_str("decoder cache configuration does not match this decoder model exactly"),
            Self::ModelParameterCountMismatch { cache, model } => write!(
                formatter,
                "decoder cache binds {cache} parameter nodes, but the model exposes {model}"
            ),
            Self::ModelParameterMismatch { index } => write!(
                formatter,
                "decoder cache parameter identity differs at stable index {index}"
            ),
            Self::ModelParameterRevisionMismatch {
                index,
                cache,
                model,
            } => write!(
                formatter,
                "decoder cache parameter revision {cache} at stable index {index} differs from model revision {model}"
            ),
            Self::LayerCountMismatch { cache, model } => write!(
                formatter,
                "decoder cache owns {cache} layer caches, but the model exposes {model} blocks"
            ),
            Self::LayerBatchSizeMismatch {
                layer,
                expected,
                actual,
            } => write!(
                formatter,
                "decoder layer {layer} cache batch size must be {expected}, got {actual}"
            ),
            Self::LayerCapacityMismatch {
                layer,
                expected,
                actual,
            } => write!(
                formatter,
                "decoder layer {layer} cache capacity must be {expected}, got {actual}"
            ),
            Self::EmptyPrompt => formatter.write_str("cached prefill needs a nonempty prompt"),
            Self::PromptTooLong { tokens, capacity } => write!(
                formatter,
                "cached prefill has {tokens} prompt tokens, exceeding capacity {capacity}"
            ),
            Self::PromptTokenOutOfBounds {
                position,
                token_id,
                vocabulary_size,
            } => write!(
                formatter,
                "cached prefill token {token_id} at position {position} is out of bounds for vocabulary {vocabulary_size}"
            ),
            Self::PrefillRequiresEmpty { len } => write!(
                formatter,
                "cached prefill requires empty state, but the cache length is {len}"
            ),
            Self::DecodeRequiresPrefill => {
                formatter.write_str("cached decode requires one completed nonempty prompt prefill")
            }
            Self::DecodeTokenOutOfBounds {
                token_id,
                vocabulary_size,
            } => write!(
                formatter,
                "cached decode token {token_id} is out of bounds for vocabulary {vocabulary_size}"
            ),
            Self::Full { capacity } => {
                write!(formatter, "decoder KV cache is full at capacity {capacity}")
            }
            Self::LayerLengthInvariant {
                layer,
                expected,
                actual,
            } => write!(
                formatter,
                "decoder layer {layer} cache length must be {expected}, got {actual}"
            ),
            Self::PreparedCacheChanged { layer } => write!(
                formatter,
                "decoder layer {layer} cache changed after its row was prepared"
            ),
            Self::WorkOverflow { counter } => {
                write!(formatter, "decoder cache {counter} counter overflows")
            }
            Self::Embedding(source) => write!(formatter, "cached token embedding: {source}"),
            Self::AttentionNorm { layer, source } => {
                write!(
                    formatter,
                    "decoder layer {layer} attention RMSNorm: {source}"
                )
            }
            Self::IncrementalAttention { layer, source } => write!(
                formatter,
                "decoder layer {layer} incremental attention: {source}"
            ),
            Self::AttentionResidual { layer, source } => write!(
                formatter,
                "decoder layer {layer} attention residual merge: {source}"
            ),
            Self::FeedForwardNorm { layer, source } => write!(
                formatter,
                "decoder layer {layer} feed-forward RMSNorm: {source}"
            ),
            Self::FeedForward { layer, source } => {
                write!(formatter, "decoder layer {layer} SwiGLU: {source}")
            }
            Self::FeedForwardResidual { layer, source } => write!(
                formatter,
                "decoder layer {layer} feed-forward residual merge: {source}"
            ),
            Self::FinalNorm(source) => write!(formatter, "cached final RMSNorm: {source}"),
            Self::Autodiff { stage, source } => write!(formatter, "cached {stage}: {source}"),
        }
    }
}

impl Error for DecoderKvCacheError {
    fn source(&self) -> Option<&(dyn Error + 'static)> {
        match self {
            Self::LayerCache { source, .. } => Some(source),
            Self::Embedding(source) => Some(source),
            Self::AttentionNorm { source, .. } | Self::FeedForwardNorm { source, .. } => {
                Some(source)
            }
            Self::IncrementalAttention { source, .. } => Some(source),
            Self::AttentionResidual { source, .. } | Self::FeedForwardResidual { source, .. } => {
                Some(source)
            }
            Self::FeedForward { source, .. } => Some(source),
            Self::FinalNorm(source) => Some(source),
            Self::Autodiff { source, .. } => Some(source),
            _ => None,
        }
    }
}

/// Exact model work committed to one decoder KV state.
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct DecoderKvCacheWork {
    token_forwards: usize,
    prefill_tokens: usize,
    decode_tokens: usize,
    cache_appends: usize,
    qkv_projection_rows: usize,
    attention_score_values: usize,
}

impl DecoderKvCacheWork {
    pub const fn token_forwards(self) -> usize {
        self.token_forwards
    }

    pub const fn prefill_tokens(self) -> usize {
        self.prefill_tokens
    }

    pub const fn decode_tokens(self) -> usize {
        self.decode_tokens
    }

    pub const fn cache_appends(self) -> usize {
        self.cache_appends
    }

    pub const fn qkv_projection_rows(self) -> usize {
        self.qkv_projection_rows
    }

    pub const fn attention_score_values(self) -> usize {
        self.attention_score_values
    }
}

/// The newest graph-free vocabulary logits after one committed cached row.
#[derive(Clone, Debug)]
pub struct CachedDecoderOutput {
    logits: TensorValue,
    position: usize,
    cache_len: usize,
    attention_score_values: usize,
}

impl CachedDecoderOutput {
    /// Returns logits shaped `[1, 1, vocabulary_size]`.
    pub fn logits(&self) -> &TensorValue {
        &self.logits
    }

    pub const fn position(&self) -> usize {
        self.position
    }

    pub const fn cache_len(&self) -> usize {
        self.cache_len
    }

    pub const fn attention_score_values(&self) -> usize {
        self.attention_score_values
    }
}

#[derive(Clone, Copy, Debug)]
enum CachedPhase {
    Prefill,
    Decode,
}

/// Reusable per-block cache storage plus coherent model-wide sequence state.
#[derive(Clone, Debug)]
pub struct DecoderKvCache {
    config: DecoderModelConfig,
    parameter_bindings: Vec<TensorValueBinding>,
    layers: Vec<LayerKvCache>,
    len: usize,
    prefill_complete: bool,
    work: DecoderKvCacheWork,
}

/// One decoder and one mutable model-wide cache bound for one checked session.
///
/// The retained parameter-value borrows prevent an optimizer from changing the
/// decoder while any K/V rows from that decoder may still be used.
pub struct DecoderKvSession<'model, 'cache> {
    model: &'model DecoderModel,
    cache: &'cache mut DecoderKvCache,
    _parameter_value_guards: Vec<Ref<'model, Tensor>>,
}

impl PartialEq for DecoderKvCache {
    fn eq(&self, other: &Self) -> bool {
        same_config(self.config, other.config)
            && self.layers == other.layers
            && self.len == other.len
            && self.prefill_complete == other.prefill_complete
            && self.work == other.work
            && self.parameter_bindings.len() == other.parameter_bindings.len()
            && self
                .parameter_bindings
                .iter()
                .zip(&other.parameter_bindings)
                .all(|(left, right)| left.same_binding(right))
    }
}

impl DecoderKvCache {
    /// Allocates one fixed-capacity cache per block and captures compatibility evidence.
    pub fn new(model: &DecoderModel) -> Result<Self, DecoderKvCacheError> {
        let config = model.config();
        let mut parameter_bindings = Vec::new();
        parameter_bindings
            .try_reserve_exact(model.parameters().len())
            .map_err(|_| DecoderKvCacheError::ParameterAllocationFailed {
                parameters: model.parameters().len(),
            })?;
        parameter_bindings.extend(
            model
                .parameters()
                .iter()
                .map(|parameter| TensorValueBinding::capture(parameter.tensor())),
        );

        let mut layers = Vec::new();
        layers
            .try_reserve_exact(model.blocks().len())
            .map_err(|_| DecoderKvCacheError::LayerAllocationFailed {
                layers: model.blocks().len(),
            })?;
        for (layer, block) in model.blocks().iter().enumerate() {
            layers.push(
                LayerKvCache::new(block.attention(), 1, config.max_positions())
                    .map_err(|source| DecoderKvCacheError::LayerCache { layer, source })?,
            );
        }
        Ok(Self {
            config,
            parameter_bindings,
            layers,
            len: 0,
            prefill_complete: false,
            work: DecoderKvCacheWork::default(),
        })
    }

    /// Validates one exact pairing and retains read borrows of the model's parameter values.
    pub fn bind<'model, 'cache>(
        &'cache mut self,
        model: &'model DecoderModel,
    ) -> Result<DecoderKvSession<'model, 'cache>, DecoderKvCacheError> {
        self.validate_binding(model)?;
        let mut parameter_value_guards: Vec<Ref<'model, Tensor>> = Vec::new();
        parameter_value_guards
            .try_reserve_exact(model.parameters().len())
            .map_err(|_| DecoderKvCacheError::ParameterAllocationFailed {
                parameters: model.parameters().len(),
            })?;
        parameter_value_guards.extend(
            model
                .parameters()
                .iter()
                .map(|parameter| parameter.tensor().value()),
        );
        Ok(DecoderKvSession {
            model,
            cache: self,
            _parameter_value_guards: parameter_value_guards,
        })
    }

    fn next_work(
        &self,
        phase: CachedPhase,
        score_values: usize,
    ) -> Result<DecoderKvCacheWork, DecoderKvCacheError> {
        let layers = self.layers.len();
        let qkv_rows = checked_mul(3, layers, DecoderKvCacheCounter::QkvProjectionRows)?;
        Ok(DecoderKvCacheWork {
            token_forwards: checked_add(
                self.work.token_forwards,
                1,
                DecoderKvCacheCounter::TokenForwards,
            )?,
            prefill_tokens: checked_add(
                self.work.prefill_tokens,
                usize::from(matches!(phase, CachedPhase::Prefill)),
                DecoderKvCacheCounter::PrefillTokens,
            )?,
            decode_tokens: checked_add(
                self.work.decode_tokens,
                usize::from(matches!(phase, CachedPhase::Decode)),
                DecoderKvCacheCounter::DecodeTokens,
            )?,
            cache_appends: checked_add(
                self.work.cache_appends,
                layers,
                DecoderKvCacheCounter::CacheAppends,
            )?,
            qkv_projection_rows: checked_add(
                self.work.qkv_projection_rows,
                qkv_rows,
                DecoderKvCacheCounter::QkvProjectionRows,
            )?,
            attention_score_values: checked_add(
                self.work.attention_score_values,
                score_values,
                DecoderKvCacheCounter::AttentionScoreValues,
            )?,
        })
    }

    fn validate_binding(&self, model: &DecoderModel) -> Result<(), DecoderKvCacheError> {
        if !same_config(self.config, model.config()) {
            return Err(DecoderKvCacheError::ModelConfigMismatch);
        }
        if self.parameter_bindings.len() != model.parameters().len() {
            return Err(DecoderKvCacheError::ModelParameterCountMismatch {
                cache: self.parameter_bindings.len(),
                model: model.parameters().len(),
            });
        }
        if let Some(index) = self
            .parameter_bindings
            .iter()
            .zip(model.parameters())
            .position(|(cached, parameter)| !cached.node_matches(parameter.tensor()))
        {
            return Err(DecoderKvCacheError::ModelParameterMismatch { index });
        }
        if let Some((index, (cached, parameter))) = self
            .parameter_bindings
            .iter()
            .zip(model.parameters())
            .enumerate()
            .find(|(_, (cached, parameter))| !cached.revision_matches(parameter.tensor()))
        {
            return Err(DecoderKvCacheError::ModelParameterRevisionMismatch {
                index,
                cache: cached.revision(),
                model: parameter.tensor().value_revision(),
            });
        }
        if self.layers.len() != model.blocks().len() {
            return Err(DecoderKvCacheError::LayerCountMismatch {
                cache: self.layers.len(),
                model: model.blocks().len(),
            });
        }
        self.validate_layer_lengths()?;
        for (layer, (block, cache)) in model.blocks().iter().zip(&self.layers).enumerate() {
            if cache.batch_size() != 1 {
                return Err(DecoderKvCacheError::LayerBatchSizeMismatch {
                    layer,
                    expected: 1,
                    actual: cache.batch_size(),
                });
            }
            if cache.capacity() != self.capacity() {
                return Err(DecoderKvCacheError::LayerCapacityMismatch {
                    layer,
                    expected: self.capacity(),
                    actual: cache.capacity(),
                });
            }
            block
                .attention()
                .validate_incremental_cache_binding(cache)
                .map_err(|source| DecoderKvCacheError::IncrementalAttention { layer, source })?;
        }
        Ok(())
    }

    fn validate_layer_lengths(&self) -> Result<(), DecoderKvCacheError> {
        for (layer, cache) in self.layers.iter().enumerate() {
            if cache.len() != self.len {
                return Err(DecoderKvCacheError::LayerLengthInvariant {
                    layer,
                    expected: self.len,
                    actual: cache.len(),
                });
            }
        }
        Ok(())
    }

    fn reset(&mut self) {
        for cache in &mut self.layers {
            cache.reset();
        }
        self.len = 0;
        self.prefill_complete = false;
        self.work = DecoderKvCacheWork::default();
    }

    pub const fn len(&self) -> usize {
        self.len
    }

    pub const fn capacity(&self) -> usize {
        self.config.max_positions()
    }

    pub fn layer_count(&self) -> usize {
        self.layers.len()
    }

    pub fn layer_len(&self, layer: usize) -> Option<usize> {
        self.layers.get(layer).map(LayerKvCache::len)
    }

    pub fn layer_cache(&self, layer: usize) -> Option<&LayerKvCache> {
        self.layers.get(layer)
    }

    pub const fn is_empty(&self) -> bool {
        self.len == 0
    }

    pub const fn is_full(&self) -> bool {
        self.len == self.capacity()
    }

    pub const fn work(&self) -> DecoderKvCacheWork {
        self.work
    }
}

impl DecoderKvSession<'_, '_> {
    /// Fills every block cache from one validated, nonempty prompt.
    ///
    /// Prompt rows advance serially through the shared one-row path. A later
    /// internal failure restores empty logical state while retaining allocation
    /// and the session's exact-model binding.
    pub fn prefill(&mut self, prompt: &[u32]) -> Result<CachedDecoderOutput, DecoderKvCacheError> {
        if prompt.is_empty() {
            return Err(DecoderKvCacheError::EmptyPrompt);
        }
        if self.cache.len != 0 || self.cache.prefill_complete {
            return Err(DecoderKvCacheError::PrefillRequiresEmpty {
                len: self.cache.len,
            });
        }
        if prompt.len() > self.cache.capacity() {
            return Err(DecoderKvCacheError::PromptTooLong {
                tokens: prompt.len(),
                capacity: self.cache.capacity(),
            });
        }
        for (position, &token_id) in prompt.iter().enumerate() {
            if !valid_token(token_id, self.cache.config.vocabulary_size()) {
                return Err(DecoderKvCacheError::PromptTokenOutOfBounds {
                    position,
                    token_id,
                    vocabulary_size: self.cache.config.vocabulary_size(),
                });
            }
        }

        let mut final_output = None;
        for &token_id in prompt {
            match self.forward_token(token_id, CachedPhase::Prefill) {
                Ok(output) => final_output = Some(output),
                Err(error) => {
                    self.cache.reset();
                    return Err(error);
                }
            }
        }
        self.cache.prefill_complete = true;
        Ok(final_output.expect("a validated prompt has at least one token"))
    }

    /// Appends one selected token and returns logits for the following choice.
    pub fn decode(&mut self, token_id: u32) -> Result<CachedDecoderOutput, DecoderKvCacheError> {
        if !self.cache.prefill_complete || self.cache.len == 0 {
            return Err(DecoderKvCacheError::DecodeRequiresPrefill);
        }
        if !valid_token(token_id, self.cache.config.vocabulary_size()) {
            return Err(DecoderKvCacheError::DecodeTokenOutOfBounds {
                token_id,
                vocabulary_size: self.cache.config.vocabulary_size(),
            });
        }
        if self.cache.is_full() {
            return Err(DecoderKvCacheError::Full {
                capacity: self.cache.capacity(),
            });
        }
        self.forward_token(token_id, CachedPhase::Decode)
    }

    fn forward_token(
        &mut self,
        token_id: u32,
        phase: CachedPhase,
    ) -> Result<CachedDecoderOutput, DecoderKvCacheError> {
        let position = self.cache.len;
        let layer_count = self.cache.layers.len();
        let (logits, prepared, score_values) = no_grad(|| {
            let embedding = self
                .model
                .embedding()
                .forward(&[token_id], &[1, 1])
                .map_err(DecoderKvCacheError::Embedding)?;
            let mut current = embedding;
            let mut prepared = Vec::new();
            prepared.try_reserve_exact(layer_count).map_err(|_| {
                DecoderKvCacheError::LayerAllocationFailed {
                    layers: layer_count,
                }
            })?;
            let mut score_values = 0usize;
            for (layer, (block, cache)) in self
                .model
                .blocks()
                .iter()
                .zip(&self.cache.layers)
                .enumerate()
            {
                let attention_norm = block
                    .attention_norm()
                    .forward(&current)
                    .map_err(|source| DecoderKvCacheError::AttentionNorm { layer, source })?;
                let ticket = block
                    .attention()
                    .prepare_incremental_bound(&attention_norm, cache)
                    .map_err(|source| DecoderKvCacheError::IncrementalAttention {
                        layer,
                        source,
                    })?;
                score_values = checked_add(
                    score_values,
                    ticket.attention_score_values(),
                    DecoderKvCacheCounter::AttentionScoreValues,
                )?;
                let after_attention = residual_add(&current, ticket.output())
                    .map_err(|source| DecoderKvCacheError::AttentionResidual { layer, source })?;
                let feed_forward_norm = block
                    .feed_forward_norm()
                    .forward(&after_attention)
                    .map_err(|source| DecoderKvCacheError::FeedForwardNorm { layer, source })?;
                let feed_forward = block
                    .feed_forward()
                    .forward(&feed_forward_norm)
                    .map_err(|source| DecoderKvCacheError::FeedForward { layer, source })?;
                current = residual_add(&after_attention, &feed_forward)
                    .map_err(|source| DecoderKvCacheError::FeedForwardResidual { layer, source })?;
                prepared.push(ticket);
            }
            let final_norm = self
                .model
                .final_norm()
                .forward(&current)
                .map_err(DecoderKvCacheError::FinalNorm)?;
            let tied_weight = self
                .model
                .tied_embedding()
                .tensor()
                .transpose(0, 1)
                .map_err(|source| DecoderKvCacheError::Autodiff {
                    stage: CachedDecoderStage::TiedWeightTranspose,
                    source,
                })?;
            let logits = final_norm.matmul(&tied_weight).map_err(|source| {
                DecoderKvCacheError::Autodiff {
                    stage: CachedDecoderStage::TiedVocabularyProjection,
                    source,
                }
            })?;
            Ok::<_, DecoderKvCacheError>((logits, prepared, score_values))
        })?;

        for (layer, (ticket, cache)) in prepared.iter().zip(&self.cache.layers).enumerate() {
            if ticket.cache_len() != position + 1 || !ticket.matches_cache(cache) {
                return Err(DecoderKvCacheError::PreparedCacheChanged { layer });
            }
        }
        let next_len = checked_add(position, 1, DecoderKvCacheCounter::TokenForwards)?;
        let next_work = self.cache.next_work(phase, score_values)?;

        for (ticket, cache) in prepared.into_iter().zip(&mut self.cache.layers) {
            let _ = ticket.commit(cache);
        }
        self.cache.len = next_len;
        self.cache.work = next_work;
        Ok(CachedDecoderOutput {
            logits,
            position,
            cache_len: next_len,
            attention_score_values: score_values,
        })
    }

    /// Clears sequence state while retaining this session's exact model binding.
    pub fn reset(&mut self) {
        self.cache.reset();
    }

    /// Exposes the bound cache for deterministic state and work inspection.
    pub fn cache(&self) -> &DecoderKvCache {
        self.cache
    }
}

fn valid_token(token_id: u32, vocabulary_size: usize) -> bool {
    usize::try_from(token_id)
        .ok()
        .is_some_and(|token| token < vocabulary_size)
}

fn same_config(left: DecoderModelConfig, right: DecoderModelConfig) -> bool {
    left.vocabulary_size() == right.vocabulary_size()
        && left.model_width() == right.model_width()
        && left.heads() == right.heads()
        && left.feed_forward_width() == right.feed_forward_width()
        && left.layers() == right.layers()
        && left.max_positions() == right.max_positions()
        && left.rope_base().to_bits() == right.rope_base().to_bits()
        && left.rms_epsilon().to_bits() == right.rms_epsilon().to_bits()
}

fn checked_add(
    left: usize,
    right: usize,
    counter: DecoderKvCacheCounter,
) -> Result<usize, DecoderKvCacheError> {
    left.checked_add(right)
        .ok_or(DecoderKvCacheError::WorkOverflow { counter })
}

fn checked_mul(
    left: usize,
    right: usize,
    counter: DecoderKvCacheCounter,
) -> Result<usize, DecoderKvCacheError> {
    left.checked_mul(right)
        .ok_or(DecoderKvCacheError::WorkOverflow { counter })
}

generate_cached повторно использует правило выбора из главы 36 и проверяет условия остановки в порядке EOS, ограничение числа токенов, затем ограничение контекста. Только если ни одно из них не выполнено, выбранный токен передаётся в decode, чтобы получить следующие логиты. Поэтому второй токен 44 в восстановленном примере возвращается на границе контекста, а полный кэш остаётся на длине 22.

Генерировать из заполненного состояния всей модели с прежними правилами сэмплирования и остановки rust/crates/llm-from-scratch/src/generation/kv_cache.rs#cached-generation
/// One selected token and the categorical evidence used by cached generation.
#[derive(Clone, Debug, PartialEq)]
pub struct CachedGenerationStep {
    prefix_length: usize,
    token_id: u32,
    unit_draw: Option<f64>,
    interval_start: f64,
    interval_end: f64,
}

impl CachedGenerationStep {
    pub const fn prefix_length(&self) -> usize {
        self.prefix_length
    }

    pub const fn token_id(&self) -> u32 {
        self.token_id
    }

    pub const fn unit_draw(&self) -> Option<f64> {
        self.unit_draw
    }

    pub const fn interval_start(&self) -> f64 {
        self.interval_start
    }

    pub const fn interval_end(&self) -> f64 {
        self.interval_end
    }
}

/// Exact cached work and its dense complete-prefix attention-score baseline.
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct CachedGenerationWork {
    prefill_tokens: usize,
    decode_tokens: usize,
    layer_cache_count: usize,
    cache_appends: usize,
    qkv_projection_rows: usize,
    cached_attention_score_values: usize,
    complete_prefix_attention_score_values: usize,
}

impl CachedGenerationWork {
    pub const fn prefill_tokens(self) -> usize {
        self.prefill_tokens
    }

    pub const fn decode_tokens(self) -> usize {
        self.decode_tokens
    }

    pub const fn layer_cache_count(self) -> usize {
        self.layer_cache_count
    }

    pub const fn cache_appends(self) -> usize {
        self.cache_appends
    }

    pub const fn qkv_projection_rows(self) -> usize {
        self.qkv_projection_rows
    }

    pub const fn cached_attention_score_values(self) -> usize {
        self.cached_attention_score_values
    }

    pub const fn complete_prefix_attention_score_values(self) -> usize {
        self.complete_prefix_attention_score_values
    }
}

/// Cached tokens, stops, final logical state, and exact work evidence.
#[derive(Clone, Debug, PartialEq)]
pub struct CachedGenerationResult {
    prompt: Vec<u32>,
    generated: Vec<u32>,
    steps: Vec<CachedGenerationStep>,
    stop: GenerationStop,
    final_cache_len: usize,
    work: CachedGenerationWork,
}

impl CachedGenerationResult {
    pub fn prompt(&self) -> &[u32] {
        &self.prompt
    }

    pub fn generated(&self) -> &[u32] {
        &self.generated
    }

    pub fn steps(&self) -> &[CachedGenerationStep] {
        &self.steps
    }

    pub const fn stop(&self) -> GenerationStop {
        self.stop
    }

    pub const fn final_cache_len(&self) -> usize {
        self.final_cache_len
    }

    pub const fn work(&self) -> CachedGenerationWork {
        self.work
    }
}

/// A cached-generation request, model step, sampling step, or work count failed.
#[derive(Debug, PartialEq)]
pub enum CachedGenerationError {
    Cache(DecoderKvCacheError),
    Sampling(SamplingError),
    EmptyPrompt,
    PromptTooLong {
        tokens: usize,
        max_positions: usize,
    },
    PromptTokenOutOfBounds {
        position: usize,
        token_id: u32,
        vocabulary_size: usize,
    },
    EosTokenOutOfBounds {
        token_id: u32,
        vocabulary_size: usize,
    },
    LogitCountMismatch {
        expected: usize,
        actual: usize,
    },
    AllocationFailed {
        values: usize,
    },
    WorkOverflow {
        counter: DecoderKvCacheCounter,
    },
}

impl fmt::Display for CachedGenerationError {
    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
        match self {
            Self::Cache(source) => source.fmt(formatter),
            Self::Sampling(source) => source.fmt(formatter),
            Self::EmptyPrompt => formatter.write_str("cached generation needs a nonempty prompt"),
            Self::PromptTooLong {
                tokens,
                max_positions,
            } => write!(
                formatter,
                "cached generation prompt has {tokens} tokens, exceeding context capacity {max_positions}"
            ),
            Self::PromptTokenOutOfBounds {
                position,
                token_id,
                vocabulary_size,
            } => write!(
                formatter,
                "cached generation prompt token {token_id} at position {position} is out of bounds for vocabulary {vocabulary_size}"
            ),
            Self::EosTokenOutOfBounds {
                token_id,
                vocabulary_size,
            } => write!(
                formatter,
                "cached generation EOS token {token_id} is out of bounds for vocabulary {vocabulary_size}"
            ),
            Self::LogitCountMismatch { expected, actual } => write!(
                formatter,
                "cached last-position logits need {expected} values, received {actual}"
            ),
            Self::AllocationFailed { values } => write!(
                formatter,
                "cannot allocate cached generation evidence for {values} values"
            ),
            Self::WorkOverflow { counter } => {
                write!(formatter, "cached generation {counter} counter overflows")
            }
        }
    }
}

impl Error for CachedGenerationError {
    fn source(&self) -> Option<&(dyn Error + 'static)> {
        match self {
            Self::Cache(source) => Some(source),
            Self::Sampling(source) => Some(source),
            _ => None,
        }
    }
}

impl From<DecoderKvCacheError> for CachedGenerationError {
    fn from(source: DecoderKvCacheError) -> Self {
        Self::Cache(source)
    }
}

impl From<SamplingError> for CachedGenerationError {
    fn from(source: SamplingError) -> Self {
        Self::Sampling(source)
    }
}

/// Generates with one prompt prefill followed by only the needed one-token decodes.
pub fn generate_cached(
    model: &DecoderModel,
    prompt: &[u32],
    config: GenerationConfig,
    rng: &mut SplitMix64,
) -> Result<CachedGenerationResult, CachedGenerationError> {
    let model_config = model.config();
    let vocabulary_size = model_config.vocabulary_size();
    let max_positions = model_config.max_positions();
    validate_generation_request(vocabulary_size, max_positions, prompt, config)?;

    let planned_steps = config.max_new_tokens().min(
        max_positions
            .checked_sub(prompt.len())
            .and_then(|remaining| remaining.checked_add(1))
            .ok_or(CachedGenerationError::AllocationFailed { values: usize::MAX })?,
    );
    let mut prompt_copy = Vec::new();
    prompt_copy.try_reserve_exact(prompt.len()).map_err(|_| {
        CachedGenerationError::AllocationFailed {
            values: prompt.len(),
        }
    })?;
    prompt_copy.extend_from_slice(prompt);
    let mut generated = Vec::new();
    generated.try_reserve_exact(planned_steps).map_err(|_| {
        CachedGenerationError::AllocationFailed {
            values: planned_steps,
        }
    })?;
    let mut steps = Vec::new();
    steps.try_reserve_exact(planned_steps).map_err(|_| {
        CachedGenerationError::AllocationFailed {
            values: planned_steps,
        }
    })?;

    if config.max_new_tokens() == 0 {
        return Ok(CachedGenerationResult {
            prompt: prompt_copy,
            generated,
            steps,
            stop: GenerationStop::TokenLimit,
            final_cache_len: 0,
            work: CachedGenerationWork {
                layer_cache_count: model_config.layers(),
                ..CachedGenerationWork::default()
            },
        });
    }

    let mut cache = DecoderKvCache::new(model)?;
    let mut session = cache.bind(model)?;
    let mut current = session.prefill(prompt)?;
    let mut prefix_length = prompt.len();
    let mut complete_prefix_attention_score_values = 0usize;

    let stop = loop {
        complete_prefix_attention_score_values = checked_complete_prefix_scores(
            complete_prefix_attention_score_values,
            model_config.layers(),
            model_config.heads(),
            prefix_length,
        )?;
        let decision = {
            let logits = current.logits().value();
            if logits.len() != vocabulary_size {
                return Err(CachedGenerationError::LogitCountMismatch {
                    expected: vocabulary_size,
                    actual: logits.len(),
                });
            }
            sample_next_token(logits.as_slice(), config.mode(), rng)?
        };
        let token_id = decision.token_id();
        generated.push(token_id);
        steps.push(CachedGenerationStep {
            prefix_length,
            token_id,
            unit_draw: decision.unit_draw(),
            interval_start: decision.interval_start(),
            interval_end: decision.interval_end(),
        });
        prefix_length =
            prefix_length
                .checked_add(1)
                .ok_or(CachedGenerationError::WorkOverflow {
                    counter: DecoderKvCacheCounter::TokenForwards,
                })?;

        if config.eos_token() == Some(token_id) {
            break GenerationStop::Eos;
        }
        if generated.len() == config.max_new_tokens() {
            break GenerationStop::TokenLimit;
        }
        if prefix_length > max_positions {
            break GenerationStop::ContextLimit;
        }
        current = session.decode(token_id)?;
    };

    let cache_work = session.cache().work();
    Ok(CachedGenerationResult {
        prompt: prompt_copy,
        generated,
        steps,
        stop,
        final_cache_len: session.cache().len(),
        work: CachedGenerationWork {
            prefill_tokens: cache_work.prefill_tokens(),
            decode_tokens: cache_work.decode_tokens(),
            layer_cache_count: session.cache().layer_count(),
            cache_appends: cache_work.cache_appends(),
            qkv_projection_rows: cache_work.qkv_projection_rows(),
            cached_attention_score_values: cache_work.attention_score_values(),
            complete_prefix_attention_score_values,
        },
    })
}

fn validate_generation_request(
    vocabulary_size: usize,
    max_positions: usize,
    prompt: &[u32],
    config: GenerationConfig,
) -> Result<(), CachedGenerationError> {
    if prompt.is_empty() {
        return Err(CachedGenerationError::EmptyPrompt);
    }
    if prompt.len() > max_positions {
        return Err(CachedGenerationError::PromptTooLong {
            tokens: prompt.len(),
            max_positions,
        });
    }
    for (position, &token_id) in prompt.iter().enumerate() {
        if !valid_token(token_id, vocabulary_size) {
            return Err(CachedGenerationError::PromptTokenOutOfBounds {
                position,
                token_id,
                vocabulary_size,
            });
        }
    }
    if let Some(token_id) = config.eos_token()
        && !valid_token(token_id, vocabulary_size)
    {
        return Err(CachedGenerationError::EosTokenOutOfBounds {
            token_id,
            vocabulary_size,
        });
    }
    if let SamplingMode::TemperatureTopK { temperature, top_k } = config.mode() {
        if !temperature.is_finite() || temperature <= 0.0 {
            return Err(SamplingError::InvalidTemperature { value: temperature }.into());
        }
        if top_k == 0 || top_k > vocabulary_size {
            return Err(SamplingError::InvalidTopK {
                top_k,
                vocabulary_size,
            }
            .into());
        }
    }
    Ok(())
}

fn checked_complete_prefix_scores(
    current: usize,
    layers: usize,
    heads: usize,
    prefix_length: usize,
) -> Result<usize, CachedGenerationError> {
    let counter = DecoderKvCacheCounter::CompletePrefixAttentionScoreValues;
    let square = prefix_length
        .checked_mul(prefix_length)
        .ok_or(CachedGenerationError::WorkOverflow { counter })?;
    let values = layers
        .checked_mul(heads)
        .and_then(|factor| factor.checked_mul(square))
        .ok_or(CachedGenerationError::WorkOverflow { counter })?;
    current
        .checked_add(values)
        .ok_or(CachedGenerationError::WorkOverflow { counter })
}

DecoderKvSession::reset обнуляет логическую длину, фазу и счётчики работы. При этом сохраняются выделенная память, записанные значения K/V за пределами пустого логического префикса и текущая связь между моделью и кэшем. Декодирование до обработки промпта, повторная обработка промпта в непустом состоянии и переполнение возвращают типизированные ошибки операции, не меняя записанное состояние. Заново построенная модель или изменённая конфигурация отклоняются раньше — при попытке создать сеанс для модели и кэша.

Пока сеанс существует, шаг AdamW доходит до предусмотренной границы изменяемого доступа и возвращает ParameterValueBorrowed. Значения параметров, состояние оптимизатора и кэша остаются неизменными. Неудачное обновление не нарушает сеанс, поэтому после него можно выполнить любой допустимый вызов decode. После уничтожения значения DecoderKvSession заимствования только для чтения освобождаются. Теперь AdamW может обновить модель и увеличить номера версий её значений. Следующая попытка связать старый кэш с этой моделью возвращает ModelParameterRevisionMismatch: сброс не может сделать устаревший кэш совместимым. Перед созданием следующего сеанса для обновлённой модели нужно создать новый DecoderKvCache.

Демонстрационная программа главы 38 проверяет логиты последней позиции при обработке промпта и декодировании в пределах допуска, отмечает раздельность двух буферов блоков, подсчитывает значения в тензорах оценок внимания с KV-кэшем и в эталонном расчёте, сравнивает решения при генерации из восстановленной контрольной точки и поведение EOS, выполняет сброс с повторным запуском и проверяет, что после каждого отклонённого создания сеанса или вызова состояние совпадает со снимком, сделанным перед этим вызовом.

Собрать проверки совпадения в пределах допуска, подсчёта работы внимания, генерации, сброса и точных вариантов ошибок для всей модели rust/demos/ch38-cached-generation/src/lib.rs#learner-evidence
/// Checks model-wide cache coherence, loaded generation parity, reset, and errors.
pub fn learner_evidence() -> Result<LearnerEvidence, FixtureError> {
    let model = fixture_model()?;
    let config = model.config();
    let mut cache = DecoderKvCache::new(&model)?;
    let mut session = cache.bind(&model)?;
    let prefill_output = session.prefill(&PROMPT)?;
    let prefill = phase_evidence(&model, session.cache(), &prefill_output, &PROMPT, 0, 0)?;
    let cached_scores_after_prefill = session.cache().work().attention_score_values();
    let decode_output = session.decode(DECODE_TOKEN)?;
    let decode_prefix = [PROMPT[0], PROMPT[1], DECODE_TOKEN];
    let decode = phase_evidence(
        &model,
        session.cache(),
        &decode_output,
        &decode_prefix,
        PROMPT.len(),
        cached_scores_after_prefill,
    )?;
    let layer_storage_distinct = session
        .cache()
        .layer_cache(0)
        .zip(session.cache().layer_cache(1))
        .is_some_and(|(left, right)| {
            left.key_storage().as_ptr() != right.key_storage().as_ptr()
                && left.value_storage().as_ptr() != right.value_storage().as_ptr()
        });
    let work = session.cache().work();
    let complete_prefix_attention_score_values = prefill
        .complete_prefix_attention_score_values
        .checked_add(decode.complete_prefix_attention_score_values)
        .ok_or(FixtureError::Invariant(
            "measured complete-prefix score total overflowed",
        ))?;
    let loaded = loaded_generation_evidence()?;
    let reset = reset_evidence(&mut session, &decode)?;
    let errors = error_evidence(&model)?;
    let cached_retained_lengths = [1, prefill.cache_after, decode.cache_after];
    let complete_prefix_lengths = [prefill.cache_after, decode.cache_after];
    let history = historical_cache_contrast(
        config,
        &cached_retained_lengths,
        &complete_prefix_lengths,
        work.attention_score_values(),
        complete_prefix_attention_score_values,
    )?;

    require(
        prefill.max_abs_difference <= TOLERANCE,
        "prefill logits differ",
    )?;
    require(
        decode.max_abs_difference <= TOLERANCE,
        "decode logits differ",
    )?;
    require(layer_storage_distinct, "layer caches share storage")?;
    require(
        work.prefill_tokens() == 2
            && work.decode_tokens() == 1
            && work.cache_appends() == 6
            && work.qkv_projection_rows() == 18
            && work.attention_score_values() == 24
            && complete_prefix_attention_score_values == 52,
        "two-layer work counters changed",
    )?;
    require(
        reset.after == 0
            && reset.allocation_reused
            && reset.storage_unchanged
            && reset.work_zeroed
            && reset.replay_identical,
        "reset evidence changed",
    )?;
    require(errors.unchanged, "a rejected cache operation changed state")?;
    require(
        loaded.cached.work().cached_attention_score_values() == 6
            && loaded
                .cached
                .work()
                .complete_prefix_attention_score_values()
                == 10,
        "loaded score counts changed",
    )?;
    Ok(LearnerEvidence {
        config,
        prefill,
        decode,
        layer_storage_distinct,
        work,
        complete_prefix_attention_score_values,
        loaded,
        reset,
        errors,
        history,
    })
}

Выполните cargo run --quiet --locked -p ch38-cached-generation. Команда напечатает точный учебный отчёт:

chapter=38-cached-generation
config=layers:2 heads:2 model_width:4 context:4 tolerance:0.000000000002
prefill=prompt:[0,1] cache:0->2 layer_lengths:[2,2] shape:[1,2,2,2] cached_scores:12 complete_prefix_scores:16 max_abs_diff:0.000000000000 logits:[1.768374438,0.208825256,1.056205728,-0.451857108,0.388467944]
decode=token:2 position:2 cache:2->3 layer_lengths:[3,3] shape:[1,2,3,2] cached_scores:12 complete_prefix_scores:36 max_abs_diff:0.000000000000 logits:[0.032908910,-0.679583624,1.408381841,0.525525421,-0.588014095]
work=prefill_tokens:2 decode_tokens:1 layer_caches:2 cache_appends:6 qkv_rows:18 cached_scores:24 complete_prefix_scores:52 layer_storage_distinct:true
loaded=checkpoint_bytes:6330 context_capacity:2 rng_state:0x9e3779b97f4a7c38 prompt:[0] generated:[4,4] text:44 prefixes:[1,2] stop:context-limit final_cache:2 prefill_tokens:1 decode_tokens:1 cached_scores:6 complete_prefix_scores:10 tokens_match:true rng_match:true
eos=token:4 generated:[4] stop:eos final_cache:1 decode_tokens:0 tokens_match:true rng_match:true
reset=before:3 after:0 allocation_reused:true storage_unchanged:true work_zeroed:true replay_identical:true
errors=decode_before_prefill:true prefill_nonempty:true overflow:true rebuilt_model:true changed_config:true unchanged:true
history=lanes:4 cached_lengths:[1,2,3] cached_scores:24 complete_prefix_lengths:[2,3] complete_prefix_scores:52 avoided_scores:28
next=assemble the complete end-to-end LLM pipeline
Напечатать точный отчёт главы 38 о генерации с KV-кэшем rust/demos/ch38-cached-generation/src/main.rs
fn main() -> Result<(), Box<dyn std::error::Error>> {
    print!("{}", ch38_cached_generation::learner_report()?);
    Ok(())
}

Проследите путь от промпта до декодирования одного токена

На рисунке сначала показаны позиции промпта 00 и 11 в двух отдельных кэшах блоков, затем — выбранный токен 22 в абсолютной позиции 22. Сводки блоков на этапе обработки промпта имеют сплошную рамку, а на этапе декодирования — двойную, поэтому фазы различимы без цвета. Каждая координата логитов с KV-кэшем расположена рядом с соответствующей координатой эталонного расчёта по полному префиксу; рядом указана максимальная абсолютная разность.

Карточки с подсчётом значений оценок внимания сопоставляют 2424 значения для пути с KV-кэшем и 5252 эталонных значения. Далее отдельно показаны две причины остановки. В запуске после восстановления контрольной точки заполнение по промпту доводит длину кэша до 11, первый выбранный токен 44 декодируется до длины 22, а второй выбранный токен 44 возвращается до остановки по ограничению контекста без ещё одного прохода декодера. В запуске с EOS первый выбранный токен также возвращается, а следующий проход декодера не выполняется.

Заполните кэши всех блоков по промпту, затем обновляйте их согласованно

Точная трассировка программы на Rust проводит промпт из двух токенов и один следующий токен через два отдельных кэша блоков декодера, проверяет логиты последней позиции в пределах допуска и сравнивает измеренное число оценок внимания, причины остановки и поведение при сбросе.

  • заполнение по промпту — сплошная рамка
  • новый шаг декодирования одного токена — двойная рамка
  • логиты последней позиции совпадают в пределах допуска

Заполнение кэшей всех блоков и декодирование одного токена

Промпт заполняет кэши обоих блоков в позициях ноль и один. Затем токен два поступает в позицию два, и длины обоих блоков согласованно увеличиваются.

заполнение по промпту — сплошная рамка

Идентификаторы токенов промпта: [0,1]

Абсолютная позиция: p=0p=0p=1p=1

Логическая длина кэша: 020\to2

Форма кэша слоя: [1,2,2,2]\left[1,2,2,2\right]

  1. Блок декодера =0\ell=0 Логическая длина кэша t=2t=2 [1,2,2,2]\left[1,2,2,2\right] отдельное хранилище K/V для каждого блока
  2. Блок декодера =1\ell=1 Логическая длина кэша t=2t=2 [1,2,2,2]\left[1,2,2,2\right] отдельное хранилище K/V для каждого блока

Путь с KV-кэшем: g0cache=1.768374438g^{\mathrm{cache}}_{0}=1.768374438g1cache=0.208825256g^{\mathrm{cache}}_{1}=0.208825256g2cache=1.056205728g^{\mathrm{cache}}_{2}=1.056205728g3cache=0.451857108g^{\mathrm{cache}}_{3}=-0.451857108g4cache=0.388467944g^{\mathrm{cache}}_{4}=0.388467944

Эталонный расчёт по полному префиксу: g0full=1.768374438g^{\mathrm{full}}_{0}=1.768374438g1full=0.208825256g^{\mathrm{full}}_{1}=0.208825256g2full=1.056205728g^{\mathrm{full}}_{2}=1.056205728g3full=0.451857108g^{\mathrm{full}}_{3}=-0.451857108g4full=0.388467944g^{\mathrm{full}}_{4}=0.388467944

логиты последней позиции совпадают в пределах допуска Максимальная абсолютная разность: Δmax=0.000000000000\Delta_{\mathrm{max}}=0.000000000000

новый шаг декодирования одного токена — двойная рамка

Выбранный токен: z=2z=2

Абсолютная позиция: p=2p=2

Логическая длина кэша: 232\to3

Форма кэша слоя: [1,2,3,2]\left[1,2,3,2\right]

  1. Блок декодера =0\ell=0 Логическая длина кэша t=3t=3 [1,2,3,2]\left[1,2,3,2\right] отдельное хранилище K/V для каждого блока
  2. Блок декодера =1\ell=1 Логическая длина кэша t=3t=3 [1,2,3,2]\left[1,2,3,2\right] отдельное хранилище K/V для каждого блока

Путь с KV-кэшем: g0cache=0.032908910g^{\mathrm{cache}}_{0}=0.032908910g1cache=0.679583624g^{\mathrm{cache}}_{1}=-0.679583624g2cache=1.408381841g^{\mathrm{cache}}_{2}=1.408381841g3cache=0.525525421g^{\mathrm{cache}}_{3}=0.525525421g4cache=0.588014095g^{\mathrm{cache}}_{4}=-0.588014095

Эталонный расчёт по полному префиксу: g0full=0.032908910g^{\mathrm{full}}_{0}=0.032908910g1full=0.679583624g^{\mathrm{full}}_{1}=-0.679583624g2full=1.408381841g^{\mathrm{full}}_{2}=1.408381841g3full=0.525525421g^{\mathrm{full}}_{3}=0.525525421g4full=0.588014095g^{\mathrm{full}}_{4}=-0.588014095

логиты последней позиции совпадают в пределах допуска Максимальная абсолютная разность: Δmax=0.000000000000\Delta_{\mathrm{max}}=0.000000000000

Считайте только значения оценок внимания

В примере подсчитываются значения, вычисленные внутри механизма внимания. Эти числа не показывают полное время работы и не являются аппаратными измерениями.

Путь с KV-кэшем
4×(1+2+3)=244\times(1+2+3)=24

Каждая из четырёх комбинаций «блок — голова» вычисляет одну строку оценок для каждой сохранённой длины.

cache_appends=6 qkv_rows=18
Эталонный расчёт по полному префиксу
4×(22+32)=524\times(2^2+3^2)=52

Два эталонных вызова заново строят квадратные матрицы оценок при длине промпта два и длине три после декодирования.

layer_caches=2

Решения при выборе токенов и причины остановки после восстановления контрольной точки

После восстановления контрольной точки оба пути выбирают одинаковые токены, используют одинаковые псевдослучайные числа и завершаются с одним состоянием генератора. Один запуск останавливается на границе контекста, а отдельный запуск с EOS — сразу после выбора токена конца последовательности.

Восстановленная контрольная точка: совпадающие решения и точная граница контекста

Идентификаторы токенов промпта: [0]

Идентификаторы и текст сгенерированных токенов: [4,4] -> 44

Причина остановки: context-limit

Логическая длина кэша: t=2t=2

Значения оценок внимания: Ncache=6N_{\mathrm{cache}}=6 Nfull=10N_{\mathrm{full}}=10

tokens_match=true rng_match=true
EOS: выберите токен и остановитесь до ненужного декодирования

Выбранный токен: zEOS=4z_{\mathrm{EOS}}=4

Идентификаторы и текст сгенерированных токенов: [4]

Причина остановки: eos

Логическая длина кэша: t=1t=1 ndecode=0n_{\mathrm{decode}}=0

tokens_match=true rng_match=true

Сброс и отклонённые вызовы без повреждения состояния

Сброс сохраняет выделенную память, но очищает логическое состояние и счётчики работы. Недопустимые для текущей фазы или ёмкости операции и попытки связать кэш с несовместимой моделью не меняют состояние, записанное после последней успешной операции.

логическое состояние возвращается к нулю

303\to0

allocation_reused=true storage_unchanged=true work_zeroed=true replay_identical=true
отклонённые вызовы не изменяют состояние
decode_before_prefill=true prefill_nonempty=true overflow=true rebuilt_model=true changed_config=true unchanged=true

Сначала предскажите, затем сверьтесь с результатами

  1. Сколько кэшей слоёв принадлежит декодеру из двух блоков?
  2. Какова логическая длина каждого кэша слоя после обработки промпта из двух токенов?
  3. Какую абсолютную позицию использует первый декодируемый токен?
  4. На каком этапе отклоняется кэш для заново построенной модели с теми же значениями весов и почему?
  5. Зачем успешный сеанс продолжает удерживать значения всех параметров доступными только для чтения после проверки зафиксированных версий?
  6. Сколько добавлений в кэш происходит для двух строк промпта и одной строки декодирования?
  7. Почему в примере вычисляются 2424 значения оценок внимания с KV-кэшем?
  8. Становится ли время расчёта внимания для новой позиции постоянным благодаря KV-кэшу?
  9. Нужно ли декодировать первый выбранный токен, если это EOS?
  10. Какие условия совместимости проверяются один раз при создании сеанса, какие условия проверяются при каждом вызове и какую связь сохраняет сброс?

Заблуждение: кэши одинаковой формы можно совместно использовать в разных блоках. На самом деле каждая сохранённая строка зависит от скрытого состояния на входе своего блока, поэтому блок 11 не может повторно использовать строки K/V блока 00. DecoderKvCache хранит данные для проверки совместимости состояния всей модели, а DecoderKvSession проверяет и сохраняет точную связь между декодером и кэшем. Каждый вложенный кэш блока также записывает параметры и геометрию внимания своего блока и его конфигурацию RoPE.

Проверьте десять ответов
  1. Декодеру принадлежат 22 кэша слоёв — по одному для каждого блока.
  2. Оба кэша достигают логической длины 22 и формы [1,2,2,2][1,2,2,2].
  3. Первый последующий токен использует абсолютную позицию 22, равную прежней длине кэша.
  4. Проверка идентичности выполняется в bind. Даже при тех же значениях, формах и конфигурации заново построенная модель содержит другие узлы параметров, поэтому создать для неё сеанс со старым кэшем нельзя.
  5. Доступ только для чтения не позволяет AdamW изменить значение после того, как bind проверил его версию, но до того, как очередная операция декодера прочитала это значение. Так проверенная связь сохраняется на всё время сеанса без копирования весов.
  6. 22 блока для 33 строк токенов дают 66 согласованных добавлений в кэши.
  7. Есть 1×2×2=41\times2\times2=4 комбинации «элемент пакета — блок — голова», и каждая вычисляет 1+2+3=61+2+3=6 значений оценок, что даёт 4×6=244\times6=24.
  8. Нет. При длине префикса tt запрос для новой позиции по-прежнему читает tt сохранённых ключей.
  9. Нет. EOS возвращается как выбранный токен, а последующие логиты не нужны.
  10. bind один раз для данного сеанса проверяет конфигурацию, идентичности узлов и версии параметров, длины и геометрию кэшей блоков, параметры внимания и RoPE. При каждой операции по-прежнему проверяются промпт, фаза, токен, свободная ёмкость, счётчики, а также хранилище и длина подготовленного изменения. Сброс обнуляет логическую длину, фазу и счётчики, сохраняя ту же связь модели и кэша, выделенную память и прежние значения K/V за пределами пустого логического префикса.

Соедините генерацию со всем процессом

Теперь полный декодер может на время одного сеанса связать совместимое состояние K/V без графа вычислений во всех блоках с конкретной моделью, обработать промпт и декодировать выбранные токены, не получая модель повторно. Сеанс продолжает использовать тот же декодер и удерживает значения его параметров доступными только для чтения. После завершения сеанса обновление весов делает старый кэш несовместимым, поэтому для обновлённой модели нужно создать новый кэш. В заданном примере решения при генерации по-прежнему совпадают с эталонным расчётом по полному префиксу. Глава 39 соединит этот способ генерации с полным процессом; в пределах одного запуска тестовые данные не смогут повлиять на выбранное состояние, а сохранённый в репозитории порядок потерь из главы 39, при котором потери декодера ниже, чем у биграммной модели, будет служить только регрессионной проверкой фиксированного примера.

В последней главе одна программа объединит весь пройденный путь: разделит данные, обучит и применит BPE, обучит и выберет декодер, выполнит локально изолированную оценку на тестовых данных, сохранит и снова загрузит контрольную точку, заполнит кэши по промпту, сгенерирует продолжение с состоянием KV-кэшей всей модели и декодирует полученные идентификаторы токенов обратно в текст. Эта граница внутри запуска не превращает сохранённый в главе 39 результат сравнения декодера с биграммной моделью в новую независимую оценку, когда в последующих запусках это сравнение повторяют.