← Все главы

12 · Версия материала 8

Преобразуйте большие по модулю логиты в устойчивые вероятности

Преобразуйте логиты словаря и внимания в устойчивые вероятности, логарифмы вероятностей, значения log-sum-exp и среднее NLL по целевым индексам на Rust без сторонних зависимостей.

Предскажите результат сдвига трёх строк

Начнём с трёх строк, в каждой из которых находятся оценки двух классов:

shape [3,2]
[[    0,     1],
 [ 1000,  1001],
 [-1001, -1000]]

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

[0,1]1=[1,0],[1000,1001]1001=[1,0],[1001,1000](1000)=[1,0].\begin{aligned} [0,1]-1 &= [-1,0], \\ [1000,1001]-1001 &= [-1,0], \\ [-1001,-1000]-(-1000) &= [-1,0]. \end{aligned}

После экспоненцирования каждая строка превращается в [0.367879441171,1]. Деление на сумму 1.367879441171 каждый раз даёт одну и ту же пару:

[0.268941421370, 0.731058578630]

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

Для целей [1,0,1] предскажите, какой элемент будет выбран в каждой строке. Отрицательные логарифмы выбранных вероятностей равны [0.313261687518,1.313261687518,0.313261687518], а их среднее — 0.646595020852 ната на одну цель.

Нормируйте после вычитания максимума

Формула численно устойчивого softmax:

pi=exp(im)jexp(jm),m=maxjjp_i=\frac{\exp(\ell_i-m)}{\sum_j\exp(\ell_j-m)}, \quad m=\max_j\ell_j

Наибольший сдвинутый логит в точности равен нулю. Поэтому хотя бы одна экспонента равна единице, а остальные не превышают её. Так мы устраняем переполнение при непосредственном вычислении exp(1000)\exp(1000) и не допускаем нулевого знаменателя для строки с большими по модулю отрицательными значениями.

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

Те же сдвинутые значения позволяют вычислить три величины в логарифмической шкале:

LSE(1,,n)=m+lnjexp(jm)\operatorname{LSE}(\ell_1,\ldots,\ell_n) =m+\ln\sum_j\exp(\ell_j-m) logpi=(im)lnjexp(jm)\log p_i=(\ell_i-m)-\ln\sum_j\exp(\ell_j-m) t=(mt)+lnjexp(jm)\mathcal{L}_t=(m-\ell_t)+\ln\sum_j\exp(\ell_j-m)

При вычислении log-sum-exp максимум возвращается после взятия натурального логарифма суммы сдвинутых экспонент. В log-softmax из сдвинутой разности вычитается этот логарифм. При вычислении NLL по индексу выбирается целевой класс, и его потеря вычисляется без предварительного округления обычной вероятности.

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

S=jexp(jm)S=\sum_j\exp(\ell_j-m)

Максимум mm, сумма сдвинутых экспонент SS и логарифм нормирующей суммы lnS\ln S используются для всех классов этой группы. Для softmax нужны mm и SS; для log-softmax — mm и lnS\ln S; для NLL по индексу — mm, выбранный целевой логит t\ell_t и lnS\ln S. Если во время обучения операция должна также сохранить вероятности для последующего вычисления градиента при обратном проходе, она может сформировать их из тех же трёх величин, не вычисляя статистики группы повторно.

Разберите обозначения вероятностей

СимволСмысл в вычислении
pip_iНормированная вероятность, присвоенная классу ii.
i\ell_iКонечный входной логит класса ii.
mmНаибольший логит в выбранной группе нормализации.
iiКласс, для которого вычисляется вероятность.
jjИндекс класса, перебирающий весь знаменатель.

Ось делит тензор на независимые группы нормализации. Для формы [3,2] и оси 1 получаются три группы по два класса. После удаления оси классов форма групп равна [3], поэтому элементы одномерного массива целевых индексов [1,0,1] соответствуют строкам с индексами ноль, один и два в том же порядке.

Операция log-sum-exp может удалить выбранную ось или сохранить её с размером один. Softmax и log-softmax сохраняют исходную форму. Любой успешно вычисленный тензор результата получает собственный непрерывный буфер в построчном порядке, даже если на вход передан срез или транспонированное представление тензора.

От softmax по словарю к вероятностям Transformer

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

Историческое основание для этого шага — работа Бенжио и соавторов A Neural Probabilistic Language Model. Бенжио и соавторы описывают softmax на выходе: его положительные значения в сумме дают единицу, а входные значения интерпретируются как ненормированные логарифмы вероятностей следующего слова.

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

Transformer использует softmax и для масштабированных оценок «запрос — ключ» внутри механизма внимания, и для предсказания следующего токена. В опубликованном коде GPT-2 для внимания используется устойчивый вариант: перед экспоненцированием из оценок по последней оси вычитают максимум, затем складывают сдвинутые экспоненты и нормируют их, прежде чем взвешивать значения.

Более поздние источники — работа Васвани и соавторов Attention Is All You Need и опубликованный OpenAI файл GPT-2 model.py. Васвани и соавторы определяют внимание на основе масштабированного скалярного произведения: к масштабированным произведениям запросов и ключей применяют softmax, после чего полученными весами взвешивают значения. Для получения вероятностей следующего токена к выходам декодера применяют обучаемое линейное преобразование и softmax. В исходном коде GPT-2 от OpenAI softmax по последней оси вычисляется так: максимум вычитается с сохранением оси единичного размера, затем значения экспоненцируются и делятся на сумму, вычисленную с таким же сохранением оси. В механизме внимания эта функция применяется к масштабированным и замаскированным оценкам до объединения значений. В коде эти операции обозначены как reduce_max и reduce_sum.

В точной арифметике прибавление одной и той же константы ко всем логитам не меняет результат softmax. Вычитание максимума сохраняет это распределение и позволяет избежать сбоев прямого экспоненцирования для строк примера. Log-sum-exp даёт численно устойчивый логарифм нормирующей суммы; log-softmax сохраняет оценки классов в логарифмической шкале, а совмещённое вычисление среднего NLL по индексам позволяет вычислить потерю для целевого класса, даже если соответствующая обычная вероятность округляется до нуля. Интерфейс для произвольной оси, требование конечных входов, схема расположения целей, правила выделения памяти и порядок ошибок — решения о корректности, принятые в реализации курса.

Небольшой пример на Rust начинается с прямого переноса формулы в код. Для [0,1] вычисление завершается успешно; при больших положительных значениях получается деление бесконечности на бесконечность, а при больших по модулю отрицательных — деление нуля на ноль. Так проявляется численная проблема, которую на пути к современным LLM решает устойчивая нормализация. Причина этой проблемы — арифметика с плавающей точкой, а не конкретный язык программирования: реализации на разных языках сталкиваются с теми же проблемами — переполнением и округлением до нуля, хотя их интерфейсы и обработка ошибок могут различаться.

Показать прямую нормализацию экспонент для одной обычной строки и двух строк с большими по модулю конечными значениями rust/demos/ch12-stable-softmax/src/lib.rs#direct-output-softmax
/// Applies the literal exponential normalization used as a bounded baseline.
///
/// This exposes finite-precision overflow and underflow; it is not attributed
/// to the software implementation of any cited language model.
pub fn direct_output_softmax(logits: &[f64]) -> Vec<f64> {
    let exponentials = logits.iter().map(|value| value.exp()).collect::<Vec<_>>();
    let denominator = exponentials.iter().sum::<f64>();
    exponentials
        .into_iter()
        .map(|value| value / denominator)
        .collect()
}

Реализуйте вычисления в логарифмической шкале с явными проверками

Порядок проверок однозначно определяет, о какой ошибке функция сообщит первой. Сначала проверяется допустимость номера оси. Затем softmax, log-softmax и NLL по индексам отклоняют пустую ось классов. Операции, возвращающие тензор, до чтения логитов проверяют схему размещения результата и резервируют для него память. Операция среднего NLL по индексам вместо этого проверяет схему групп, соответствие числа целей числу групп, наличие хотя бы одной цели и границы всех целевых индексов до чтения любого логита. Только после этого функция может сообщить о первом NaN или первой положительной либо отрицательной бесконечности в порядке «сначала группа, затем класс».

Различать ошибки осей, пустых классов, результата, неконечных логитов и целевых индексов rust/crates/llm-from-scratch/src/nn/probability.rs#probability-errors
/// A rejected probability operation, target, output, or converted view operation.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ProbabilityError {
    /// An owned output layout violates the tensor storage invariant.
    Tensor(TensorError),
    /// A tensor-view error was converted into the probability error type.
    View(TensorViewError),
    /// The requested class axis does not exist.
    AxisOutOfBounds { axis: usize, rank: usize },
    /// Softmax, log-softmax, and indexed NLL need at least one class.
    EmptyNormalizationAxis { axis: usize },
    /// The checked output shape is valid, but its value buffer cannot be reserved.
    OutputAllocationFailed { elements: usize },
    /// The first rejected logit in group-major, class-minor order is NaN.
    NaNLogit { group: usize, class: usize },
    /// The first rejected logit in group-major, class-minor order is positive infinity.
    PositiveInfinityLogit { group: usize, class: usize },
    /// The first rejected logit in group-major, class-minor order is negative infinity.
    NegativeInfinityLogit { group: usize, class: usize },
    /// There must be one flat target for every class-axis group.
    TargetCountMismatch { expected: usize, actual: usize },
    /// A mean is undefined when there are no target groups.
    EmptyTargets,
    /// One target does not name a class on the selected axis.
    TargetOutOfBounds {
        group: usize,
        target: usize,
        classes: usize,
    },
}

impl fmt::Display for ProbabilityError {
    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
        match self {
            Self::Tensor(error) => error.fmt(formatter),
            Self::View(error) => error.fmt(formatter),
            Self::AxisOutOfBounds { axis, rank } => {
                write!(
                    formatter,
                    "probability axis {axis} is out of bounds for rank {rank}"
                )
            }
            Self::EmptyNormalizationAxis { axis } => {
                write!(formatter, "probability axis {axis} has no classes")
            }
            Self::OutputAllocationFailed { elements } => write!(
                formatter,
                "cannot allocate probability output for {elements} f64 values"
            ),
            Self::NaNLogit { group, class } => {
                write!(formatter, "logit at group {group}, class {class} is NaN")
            }
            Self::PositiveInfinityLogit { group, class } => write!(
                formatter,
                "logit at group {group}, class {class} is positive infinity"
            ),
            Self::NegativeInfinityLogit { group, class } => write!(
                formatter,
                "logit at group {group}, class {class} is negative infinity"
            ),
            Self::TargetCountMismatch { expected, actual } => write!(
                formatter,
                "indexed mean NLL needs {expected} targets, but received {actual}"
            ),
            Self::EmptyTargets => formatter.write_str("indexed mean NLL needs at least one target"),
            Self::TargetOutOfBounds {
                group,
                target,
                classes,
            } => write!(
                formatter,
                "target {target} at group {group} is out of bounds for {classes} classes"
            ),
        }
    }
}

impl Error for ProbabilityError {
    fn source(&self) -> Option<&(dyn Error + 'static)> {
        match self {
            Self::Tensor(error) => Some(error),
            Self::View(error) => Some(error),
            _ => None,
        }
    }
}

impl From<TensorError> for ProbabilityError {
    fn from(error: TensorError) -> Self {
        Self::Tensor(error)
    }
}

impl From<TensorViewError> for ProbabilityError {
    fn from(error: TensorViewError) -> Self {
        Self::View(error)
    }
}

Перед чтением логитов AxisPlan удаляет выбранную ось классов из формы и шагов исходного тензора. Оставшиеся форма и шаги задают проверяемый курсор. Каждое значение этого курсора — отсчитываемое от нуля смещение элемента класса 0 в одной группе нормализации относительно начала плоского буфера исходного тензора. Смещение измеряется количеством элементов f64, а не байтами. В этой главе такое значение называется базовым смещением группы. Шаг удалённой оси становится шагом по оси классов. Если один раз прибавить этот шаг к базовому смещению, получится смещение класса 11; если прибавить два раза — смещение класса 22 и так далее.

В непрерывном тензоре из примера форма [3,2] имеет шаги [2,1], а ось классов имеет индекс 1. После её удаления остаются форма групп [3] и шаг групп [2], поэтому курсор выдаёт базовые смещения [0,2,4]. Шаг по оси классов равен 1. Следовательно, три группы читают элементы исходного хранилища со смещениями [0,1], [2,3] и [4,5]. Эти группы соответствуют строкам с индексами ноль, один и два — в том же порядке, в котором в одномерном массиве целевых индексов записаны [1,0,1].

Для каждой непустой группы нормализации первый проход находит максимум mm. Второй проход возвращается к тому же базовому смещению, посещает классы в порядке возрастания индексов, отдельно учитывает одно слагаемое exp(0)=1\exp(0)=1 и накапливает остальные сдвинутые экспоненты в tail. Затем он сохраняет S=1+tailS=1+\mathrm{tail} для деления при вычислении вероятностей и lnS=ln(1+tail)\ln S=\ln(1+\mathrm{tail}) для результатов в логарифмической шкале; последний логарифм вычисляется с помощью ln_1p. Если максимальный логит одинаков у нескольких классов, отдельно учитывается только одно единичное слагаемое, а экспоненты остальных классов с тем же максимумом входят в tail. Все классы группы используют эти три величины. ln_1p сохраняет представимую субнормальную поправку в логарифмической шкале, которую обычное ln(1+tail)\ln(1+\mathrm{tail}) могло бы потерять из-за округления до взятия логарифма.

Один раз вычислить повторно используемый набор статистик для каждой проверенной группы нормализации rust/crates/llm-from-scratch/src/nn/probability.rs#checked-probability-groups
#[derive(Debug)]
struct AxisPlan {
    axis: usize,
    classes: usize,
    group_shape: Vec<usize>,
    group_strides: Vec<usize>,
    groups: usize,
    class_stride: usize,
}

#[derive(Clone, Copy, Debug)]
struct RowStats {
    maximum: f64,
    shifted_exponential_sum: f64,
    log_shifted_exponential_sum: f64,
}

#[derive(Clone, Copy, Debug)]
struct FiniteLogits;

#[derive(Clone, Copy, Debug)]
enum LogitFiniteness {
    Check,
    Validated(FiniteLogits),
}

impl AxisPlan {
    fn new(
        input: &TensorView<'_>,
        axis: usize,
        allow_empty_axis: bool,
    ) -> Result<Self, ProbabilityError> {
        if axis >= input.rank() {
            return Err(ProbabilityError::AxisOutOfBounds {
                axis,
                rank: input.rank(),
            });
        }

        let classes = input.shape()[axis];
        if classes == 0 && !allow_empty_axis {
            return Err(ProbabilityError::EmptyNormalizationAxis { axis });
        }

        let mut group_shape = input.shape().to_vec();
        group_shape.remove(axis);
        let mut group_strides = input.strides().to_vec();
        let class_stride = group_strides.remove(axis);
        let (_, groups) = checked_row_major_layout(&group_shape)?;
        Ok(Self {
            axis,
            classes,
            group_shape,
            group_strides,
            groups,
            class_stride,
        })
    }

    fn group_offsets(&self, input: &TensorView<'_>) -> StridedOffsets {
        input
            .projected_offsets(&self.group_shape, &self.group_strides, self.groups)
            .expect("a checked probability plan retains valid group-base offsets")
    }

    fn output_group_offsets(&self, output_strides: &[usize], output_len: usize) -> StridedOffsets {
        let mut group_strides = output_strides.to_vec();
        group_strides.remove(self.axis);
        StridedOffsets::checked(
            &self.group_shape,
            &group_strides,
            0,
            self.groups,
            output_len,
        )
        .expect("a checked probability output retains valid group-base offsets")
    }

    fn target_offset(&self, group_base: usize, target: usize) -> usize {
        let class_offset = target
            .checked_mul(self.class_stride)
            .expect("a checked probability plan cannot overflow a class offset");
        group_base
            .checked_add(class_offset)
            .expect("a checked probability plan cannot overflow a target offset")
    }

    fn for_each_group(
        &self,
        input: &TensorView<'_>,
        finiteness: LogitFiniteness,
        mut visit: impl FnMut(usize, usize, RowStats) -> Result<(), ProbabilityError>,
    ) -> Result<(), ProbabilityError> {
        for (group, group_base) in self.group_offsets(input).enumerate() {
            let stats = row_stats(input, self, finiteness, group, group_base)?;
            visit(group, group_base, stats)?;
        }
        Ok(())
    }
}

fn checked_finite_logit(value: f64, group: usize, class: usize) -> Result<f64, ProbabilityError> {
    if value.is_nan() {
        Err(ProbabilityError::NaNLogit { group, class })
    } else if value == f64::INFINITY {
        Err(ProbabilityError::PositiveInfinityLogit { group, class })
    } else if value == f64::NEG_INFINITY {
        Err(ProbabilityError::NegativeInfinityLogit { group, class })
    } else {
        Ok(value)
    }
}

fn validate_finite_logits(
    input: &TensorView<'_>,
    plan: &AxisPlan,
) -> Result<FiniteLogits, ProbabilityError> {
    for (group, group_base) in plan.group_offsets(input).enumerate() {
        let mut input_offset = group_base;
        for class in 0..plan.classes {
            checked_finite_logit(input.value_at_storage_offset(input_offset), group, class)?;
            if class + 1 < plan.classes {
                input_offset = input_offset
                    .checked_add(plan.class_stride)
                    .expect("a checked probability plan cannot overflow along the class axis");
            }
        }
    }
    Ok(FiniteLogits)
}

fn row_stats(
    input: &TensorView<'_>,
    plan: &AxisPlan,
    finiteness: LogitFiniteness,
    group: usize,
    group_base: usize,
) -> Result<RowStats, ProbabilityError> {
    debug_assert!(plan.classes > 0);
    let mut maximum = f64::NEG_INFINITY;
    let mut input_offset = group_base;
    for class in 0..plan.classes {
        let value = input.value_at_storage_offset(input_offset);
        let value = match finiteness {
            LogitFiniteness::Check => checked_finite_logit(value, group, class)?,
            LogitFiniteness::Validated(_) => value,
        };
        maximum = maximum.max(value);
        if class + 1 < plan.classes {
            input_offset = input_offset
                .checked_add(plan.class_stride)
                .expect("a checked probability plan cannot overflow along the class axis");
        }
    }

    let mut exponential_tail = 0.0;
    let mut skipped_one_maximum = false;
    input_offset = group_base;
    for class in 0..plan.classes {
        let value = input.value_at_storage_offset(input_offset);
        let shifted = value - maximum;
        if shifted == 0.0 && !skipped_one_maximum {
            skipped_one_maximum = true;
        } else {
            exponential_tail += shifted.exp();
        }
        if class + 1 < plan.classes {
            input_offset = input_offset
                .checked_add(plan.class_stride)
                .expect("a checked probability plan cannot overflow along the class axis");
        }
    }
    debug_assert!(skipped_one_maximum);

    Ok(RowStats {
        maximum,
        shifted_exponential_sum: 1.0 + exponential_tail,
        log_shifted_exponential_sum: exponential_tail.ln_1p(),
    })
}

Каждый вызов прямого прохода создаёт один проверенный план оси и групп и ровно один раз для каждой группы вычисляет её статистики. После этого из полученных величин формируется запрошенный результат: log-sum-exp, softmax, log-softmax или среднее NLL по индексам. Доступные только внутри крейта вспомогательные функции log_softmax_forward и indexed_mean_nll_forward могут вместе с основным результатом сформировать вероятности, чтобы последующее вычисление градиента при обратном проходе использовало их без повторной нормализации логитов. Общедоступные функции не возвращают этот дополнительный тензор вероятностей; каждая из них возвращает только результат, указанный в её интерфейсе.

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

Слова «ровно один раз для каждой группы» не означают, что каждый логит читается лишь однажды. Для устойчивого вычисления статистик группы всё равно нужны проход поиска максимума и отдельный проход суммирования сдвинутых экспонент; для записи значений по классам нужен ещё один проход по классам. Это также не объединяет отдельные вызовы общедоступных функций: softmax, а затем log_softmax — это два независимых вызова прямого прохода.

Срез или транспонированное представление передаёт собственные шаги групп, шаг по оси классов и начальное смещение. Поэтому тот же обход читает логические группы такого представления без копирования. Softmax и log-softmax создают отдельный проверяемый курсор базовых смещений для непрерывного результата. Курсоры входа и результата посещают группы в одинаковом построчном порядке, а внутри группы отдельные шаги по оси классов помещают каждое значение в нужную логическую позицию результата. Ни один численный проход не создаёт вектор координат класса и не вызывает TensorView::get для каждого скаляра. Чтение и запись скаляров по-прежнему используют обычную безопасную индексацию с проверкой границ.

Softmax делит сдвинутые экспоненты на SS. Log-softmax вычитает lnS\ln S из сдвинутого логита. Для TT целей совмещённое вычисление среднего NLL по индексам ведёт два накопителя. Основной накопитель total складывает полную неотрицательную потерю целевого класса в каждой группе. Если потери всех групп и текущая сумма остаются конечными, функция один раз делит total на TT и возвращает среднее в натах на одну цель. Благодаря единственному делению в конце сохраняется корректное округление представимого субнормального среднего.

Параллельно запасной накопитель scaled_mean складывает две неотрицательные части потери каждой группы, предварительно разделив каждую часть на TT: разность между максимумом и целевым логитом (mtr)/T(m-\ell_{t_r})/T, где trt_r — целевой класс группы rr, и логарифм нормирующей суммы ln(1+tail)/T\ln(1+\mathrm{tail})/T. Если сама разность mtrm-\ell_{t_r} переполняется, запасной способ вычисляет вместо неё m/Ttr/Tm/T-\ell_{t_r}/T. Функция возвращает scaled_mean только при переполнении полной потери группы или текущего значения total; иначе она делит total на TT и возвращает полученное частное. После проверки границ всех целей выбранный целевой логит читается по смещению, равному базовому смещению группы плюс индекс цели, умноженный на шаг исходного тензора по оси классов. Поэтому все целевые индексы проверяются до чтения любого логита, а ошибки неконечных значений по-прежнему выдаются в порядке «сначала группа, затем класс».

Сформировать запрошенные вероятностные результаты за один проверенный обход групп rust/crates/llm-from-scratch/src/nn/probability.rs#stable-probability-operations
fn output_buffer(elements: usize) -> Result<Vec<f64>, ProbabilityError> {
    let mut values = Vec::new();
    values
        .try_reserve_exact(elements)
        .map_err(|_| ProbabilityError::OutputAllocationFailed { elements })?;
    values.resize(elements, 0.0);
    Ok(values)
}

fn positive_zero(value: f64) -> f64 {
    if value == 0.0 { 0.0 } else { value }
}

/// Requested normalized values emitted by one checked forward traversal.
#[derive(Debug)]
struct NormalizedForward {
    probabilities: Option<Tensor>,
    log_probabilities: Option<Tensor>,
}

/// Log-softmax output and optional probabilities from the same forward traversal.
#[derive(Debug)]
pub(crate) struct LogSoftmaxForward {
    pub(crate) value: Tensor,
    pub(crate) probabilities: Option<Tensor>,
}

/// Indexed mean NLL and optional probabilities emitted by its forward traversal.
#[derive(Debug)]
pub(crate) struct IndexedMeanNllForward {
    pub(crate) loss: f64,
    pub(crate) probabilities: Option<Tensor>,
}

struct NormalizedGroupOutput<'a> {
    output_group_base: usize,
    output_class_stride: usize,
    probabilities: Option<&'a mut [f64]>,
    log_probabilities: Option<&'a mut [f64]>,
}

impl NormalizedGroupOutput<'_> {
    fn emit(
        mut self,
        input: &TensorView<'_>,
        plan: &AxisPlan,
        input_group_base: usize,
        stats: RowStats,
    ) {
        debug_assert!(self.probabilities.is_some() || self.log_probabilities.is_some());
        let mut input_offset = input_group_base;
        let mut output_offset = self.output_group_base;
        for class in 0..plan.classes {
            let shifted = input.value_at_storage_offset(input_offset) - stats.maximum;
            if let Some(values) = self.probabilities.as_mut() {
                values[output_offset] =
                    positive_zero(shifted.exp() / stats.shifted_exponential_sum);
            }
            if let Some(values) = self.log_probabilities.as_mut() {
                values[output_offset] = positive_zero(shifted - stats.log_shifted_exponential_sum);
            }

            if class + 1 < plan.classes {
                input_offset = input_offset
                    .checked_add(plan.class_stride)
                    .expect("a checked probability plan cannot overflow along the class axis");
                output_offset = output_offset
                    .checked_add(self.output_class_stride)
                    .expect("a checked probability output cannot overflow along the class axis");
            }
        }
    }
}

/// Reduces one axis with max-shifted log-sum-exp.
///
/// An empty selected axis returns the log-additive identity, negative infinity,
/// once per remaining-axis group. Other non-finite logits are rejected in
/// group-major, class-minor order.
pub fn log_sum_exp(
    input: &TensorView<'_>,
    axis: usize,
    keep_dim: bool,
) -> Result<Tensor, ProbabilityError> {
    let plan = AxisPlan::new(input, axis, true)?;
    let output_shape = if keep_dim {
        let mut shape = input.shape().to_vec();
        shape[axis] = 1;
        shape
    } else {
        plan.group_shape.clone()
    };
    let (_, output_len) = checked_row_major_layout(&output_shape)?;
    debug_assert_eq!(output_len, plan.groups);
    let mut values = output_buffer(output_len)?;

    if plan.classes == 0 {
        values.fill(f64::NEG_INFINITY);
    } else {
        plan.for_each_group(
            input,
            LogitFiniteness::Check,
            |group, _group_base, stats| {
                values[group] = stats.maximum + stats.log_shifted_exponential_sum;
                Ok(())
            },
        )?;
    }

    Tensor::from_vec(output_shape, values).map_err(Into::into)
}

/// Converts finite logits to normalized probabilities along one explicit axis.
pub fn softmax(input: &TensorView<'_>, axis: usize) -> Result<Tensor, ProbabilityError> {
    let forward = normalized_forward(input, axis, true, false)?;
    Ok(forward
        .probabilities
        .expect("softmax requests a probability output"))
}

/// Converts finite logits to normalized log-probabilities along one explicit axis.
pub fn log_softmax(input: &TensorView<'_>, axis: usize) -> Result<Tensor, ProbabilityError> {
    let forward = log_softmax_forward(input, axis, false)?;
    debug_assert!(forward.probabilities.is_none());
    Ok(forward.value)
}

pub(crate) fn log_softmax_forward(
    input: &TensorView<'_>,
    axis: usize,
    emit_probabilities: bool,
) -> Result<LogSoftmaxForward, ProbabilityError> {
    let forward = normalized_forward(input, axis, emit_probabilities, true)?;
    Ok(LogSoftmaxForward {
        value: forward
            .log_probabilities
            .expect("log-softmax forward requests a log-probability output"),
        probabilities: forward.probabilities,
    })
}

fn normalized_forward(
    input: &TensorView<'_>,
    axis: usize,
    emit_probabilities: bool,
    emit_log_probabilities: bool,
) -> Result<NormalizedForward, ProbabilityError> {
    debug_assert!(emit_probabilities || emit_log_probabilities);
    let plan = AxisPlan::new(input, axis, false)?;
    let (output_strides, output_len) = checked_row_major_layout(input.shape())?;
    let mut log_probability_values = emit_log_probabilities
        .then(|| output_buffer(output_len))
        .transpose()?;
    let finiteness = if emit_probabilities && emit_log_probabilities {
        LogitFiniteness::Validated(validate_finite_logits(input, &plan)?)
    } else {
        LogitFiniteness::Check
    };
    let mut probability_values = emit_probabilities
        .then(|| output_buffer(output_len))
        .transpose()?;
    let output_class_stride = output_strides[axis];

    let mut output_group_offsets = plan.output_group_offsets(&output_strides, output_len);
    plan.for_each_group(input, finiteness, |_group, input_group_base, stats| {
        let output_group_base = output_group_offsets
            .next()
            .expect("a checked probability output has one base per input group");
        NormalizedGroupOutput {
            output_group_base,
            output_class_stride,
            probabilities: probability_values.as_deref_mut(),
            log_probabilities: log_probability_values.as_deref_mut(),
        }
        .emit(input, &plan, input_group_base, stats);
        Ok(())
    })?;
    debug_assert!(output_group_offsets.next().is_none());

    let log_probabilities = log_probability_values
        .map(|values| Tensor::from_vec(input.shape().to_vec(), values))
        .transpose()?;
    let probabilities = probability_values
        .map(|values| Tensor::from_vec(input.shape().to_vec(), values))
        .transpose()?;
    Ok(NormalizedForward {
        probabilities,
        log_probabilities,
    })
}

/// Scores one class index per remaining-axis group with fused stable mean NLL.
///
/// Targets follow the row-major group shape obtained by removing `axis` from
/// the logits. Bounds are checked for every target before a logit is read.
pub fn indexed_mean_nll(
    logits: &TensorView<'_>,
    axis: usize,
    targets: &[usize],
) -> Result<f64, ProbabilityError> {
    let forward = indexed_mean_nll_forward(logits, axis, targets, false)?;
    debug_assert!(forward.probabilities.is_none());
    Ok(forward.loss)
}

pub(crate) fn indexed_mean_nll_forward(
    logits: &TensorView<'_>,
    axis: usize,
    targets: &[usize],
    emit_probabilities: bool,
) -> Result<IndexedMeanNllForward, ProbabilityError> {
    let plan = AxisPlan::new(logits, axis, false)?;
    if targets.len() != plan.groups {
        return Err(ProbabilityError::TargetCountMismatch {
            expected: plan.groups,
            actual: targets.len(),
        });
    }
    if targets.is_empty() {
        return Err(ProbabilityError::EmptyTargets);
    }
    for (group, &target) in targets.iter().enumerate() {
        if target >= plan.classes {
            return Err(ProbabilityError::TargetOutOfBounds {
                group,
                target,
                classes: plan.classes,
            });
        }
    }

    let finiteness = if emit_probabilities {
        LogitFiniteness::Validated(validate_finite_logits(logits, &plan)?)
    } else {
        LogitFiniteness::Check
    };

    let output_layout = emit_probabilities
        .then(|| checked_row_major_layout(logits.shape()))
        .transpose()?;
    let mut probability_values = output_layout
        .as_ref()
        .map(|(_, output_len)| output_buffer(*output_len))
        .transpose()?;
    let mut output_group_offsets = output_layout
        .as_ref()
        .map(|(output_strides, output_len)| plan.output_group_offsets(output_strides, *output_len));
    let output_class_stride = output_layout
        .as_ref()
        .map(|(output_strides, _)| output_strides[axis]);

    let mut total = 0.0;
    let mut scaled_mean = 0.0;
    let mut needs_scaled_fallback = false;
    let target_count = targets.len() as f64;
    plan.for_each_group(logits, finiteness, |group, group_base, stats| {
        let target = targets[group];
        let target_logit = logits.value_at_storage_offset(plan.target_offset(group_base, target));
        let gap = stats.maximum - target_logit;
        let scaled_gap = if gap.is_finite() {
            gap / target_count
        } else {
            stats.maximum / target_count - target_logit / target_count
        };
        scaled_mean += scaled_gap + stats.log_shifted_exponential_sum / target_count;

        let loss = gap + stats.log_shifted_exponential_sum;
        if loss.is_finite() && !needs_scaled_fallback {
            total += loss;
            if !total.is_finite() {
                needs_scaled_fallback = true;
            }
        } else {
            needs_scaled_fallback = true;
        }

        if let Some(values) = probability_values.as_deref_mut() {
            let output_group_base = output_group_offsets
                .as_mut()
                .and_then(Iterator::next)
                .expect("a checked probability output has one base per input group");
            NormalizedGroupOutput {
                output_group_base,
                output_class_stride: output_class_stride
                    .expect("a requested probability output has a class stride"),
                probabilities: Some(values),
                log_probabilities: None,
            }
            .emit(logits, &plan, group_base, stats);
        }
        Ok(())
    })?;
    debug_assert!(
        output_group_offsets
            .as_mut()
            .is_none_or(|offsets| offsets.next().is_none())
    );

    let loss = positive_zero(if needs_scaled_fallback {
        scaled_mean
    } else {
        total / target_count
    });
    let probabilities = probability_values
        .map(|values| Tensor::from_vec(logits.shape().to_vec(), values))
        .transpose()?;
    Ok(IndexedMeanNllForward {
        loss,
        probabilities,
    })
}

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

Выполнить все проверяемые вероятностные операции над общим примером из трёх строк rust/demos/ch12-stable-softmax/src/lib.rs#tiny-stable-softmax-example
/// Normalizes the same relative logits after three different constant shifts.
pub fn tiny_stable_softmax_example() -> Result<TinyStableSoftmaxExample, ProbabilityError> {
    let logits = Tensor::from_vec(LOGIT_SHAPE.to_vec(), LOGIT_VALUES.to_vec())?;
    let probabilities = softmax(&logits.view(), CLASS_AXIS)?;
    let log_probabilities = log_softmax(&logits.view(), CLASS_AXIS)?;
    let log_normalizers = log_sum_exp(&logits.view(), CLASS_AXIS, false)?;
    let mean_nll = indexed_mean_nll(&logits.view(), CLASS_AXIS, &TARGETS)?;

    Ok(TinyStableSoftmaxExample {
        logits,
        probabilities,
        log_probabilities,
        log_normalizers,
        mean_nll,
    })
}

Для пустой выбранной оси в логарифмической шкале существует полезный нейтральный элемент: log-sum-exp возвращает отрицательную бесконечность для каждой группы по оставшимся осям. Нормированного распределения в этом случае нет, поэтому остальные операции завершаются ошибкой. Нулевой размер другой оси даёт корректный пустой тензор без чтения значений. Точный ноль приводится к положительному нулю, поэтому log-softmax и NLL для единственного класса также в точности равны положительному нулю.

Вычитание максимума устраняет предотвратимые сбои, но не ограничения f64. Вероятность крайне маловероятного класса с конечным логитом всё ещё может округлиться до нуля, тогда как её логарифм и потеря, вычисленная непосредственно из логитов, останутся конечными. Если математическая разность между логитами классов не представима в f64, её вычисленное значение всё ещё может стать положительной или отрицательной бесконечностью, а log-sum-exp у верхней границы может округлиться до f64::MAX. В главе 7 уже округлённая до нуля вероятность наблюдаемой цели правильно даёт бесконечное NLL; совмещённое вычисление потери непосредственно из логитов позволяет избежать этого округления, если ответ в логарифмической шкале представим.

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

Подготовить устойчивые результаты, признаки сбоев прямого вычисления, потерю цели, инвариантность и типизированные ошибки rust/demos/ch12-stable-softmax/src/main.rs#learner-stable-softmax-output
    let example = tiny_stable_softmax_example()?;
    let row_sums = example
        .probabilities
        .as_slice()
        .chunks_exact(2)
        .map(|row| row.iter().sum())
        .collect::<Vec<f64>>();
    let target_losses = TARGETS
        .iter()
        .enumerate()
        .map(|(row, &target)| -example.log_probabilities.as_slice()[row * 2 + target])
        .collect::<Vec<_>>();
    let ordinary_direct = direct_output_softmax(&[0.0, 1.0]);
    let overflow_direct = direct_output_softmax(&[1000.0, 1001.0]);
    let underflow_direct = direct_output_softmax(&[-1001.0, -1000.0]);

    let axis_error = softmax(&example.logits.view(), 2).unwrap_err();
    let empty_logits = Tensor::from_vec(vec![2, 0], vec![])?;
    let empty_error = softmax(&empty_logits.view(), 1).unwrap_err();
    let nonfinite_logits = Tensor::from_vec(vec![1, 2], vec![0.0, f64::INFINITY])?;
    let nonfinite_error = softmax(&nonfinite_logits.view(), 1).unwrap_err();
    let target_error = indexed_mean_nll(&example.logits.view(), 1, &[1, 2, 1]).unwrap_err();
./course run cargo run --quiet --locked -p ch12-stable-softmax
stable softmax: shape=[3, 2] values=[0.268941421370, 0.731058578630, 0.268941421370, 0.731058578630, 0.268941421370, 0.731058578630]
log-sum-exp: shape=[3] values=[1.313261687518, 1001.313261687518, -999.686738312482]
targets: [1, 0, 1] losses=[0.313261687518, 1.313261687518, 0.313261687518] mean_nll=0.646595020852
naive overflow [1000, 1001]: undefined=true
naive underflow [-1001, -1000]: undefined=true

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

Сравните прямое вычисление softmax с устойчивым

Для каждой строки таблица сопоставляет результат прямого экспоненцирования с этапами устойчивого вычисления после вычитания максимума. В таблице показаны уже полученные программой на Rust максимум, сдвинутые значения, знаменатель, вероятности, выбранная цель и потеря. Разные символы и типы рамок позволяют без опоры на цвет различить пять состояний: конечный результат прямого вычисления, устойчивый результат, переполнение, округление экспонент до нуля и отклонённый запрос.

Как вычитание максимума делает вычисление устойчивым

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

Форма логитов
[3, 2]
Ось классов
1
Среднее NLL
0.646595020852

Вычтите максимум каждой строки перед экспоненцированием

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

Вычтите максимум каждой строки перед экспоненцированием Строка 0 Строка 1 Строка 2
Исходные логиты [0.000000000000, 1.000000000000][1000.000000000000, 1001.000000000000][-1001.000000000000, -1000.000000000000]
Максимум 1.0000000000001001.000000000000-1000.000000000000
Сдвинутые логиты [-1.000000000000, 0.000000000000][-1.000000000000, 0.000000000000][-1.000000000000, 0.000000000000]
Экспоненты [0.367879441171, 1.000000000000][0.367879441171, 1.000000000000][0.367879441171, 1.000000000000]
Вычисление без сдвига конечный результат : [1.000000000000, 2.718281828459] не определено из-за переполнения не определено из-за округления экспонент до нуля
Вероятности [0.268941421370, 0.731058578630][0.268941421370, 0.731058578630][0.268941421370, 0.731058578630]

Во всех трёх строках записаны одинаковые вероятности. Знаменатель: 1.367879441171

Выберите логарифм вероятности по целевому индексу каждой строки

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

Строка 0

Класс: 1

Логарифмы вероятностей: -0.313261687518

Потеря NLL: 0.313261687518

Строка 1

Класс: 0

Логарифмы вероятностей: -1.313261687518

Потеря NLL: 1.313261687518

Строка 2

Класс: 1

Логарифмы вероятностей: -0.313261687518

Потеря NLL: 0.313261687518

Отклоните недопустимые оси, логиты и цели

Карточки со штриховой рамкой показывают точные значения оси, группы, класса и цели, из-за которых запрос отклонён до выполнения недопустимого вычисления.

axis-out-of-bounds

Ось классов 2; Ранг 2

empty-normalization-axis

Ось классов 1

positive-infinity-logit

Группа 0; Класс 1

target-out-of-bounds

Группа 1; Класс 2; Классы 2

Сделайте прогноз перед запуском Rust

  1. Предскажите сдвинутые значения для [1000,1001].
  2. Изменится ли какая-либо вероятность, если прибавить -1001 к [0,1]?
  3. Объясните, почему прямая нормализация [1000,1001] не определена в f64.
  4. Предскажите обе вероятности для одинаковых логитов [7,7].
  5. Найдите потерю для целевого класса 0 в строке [1000,1001].
  6. Предскажите формы результата log-sum-exp для входа [2,3,4] и оси 1 при выключенном и включённом keep_dim.
  7. У какой операции над пустой осью классов определён нейтральный элемент?
  8. Для формы [3,2], шагов исходного тензора [2,1] и оси классов 1 перечислите три базовых смещения групп и два смещения в исходном хранилище, которые читаются для каждой группы.
  9. Пусть во время обучения одна операция должна вернуть значения log-softmax и сохранить вероятности softmax для последующего вычисления градиента при обратном проходе. Какие величины, общие для всей группы, можно использовать для обоих результатов? Какую работу по каждому классу всё равно потребуется выполнить?
  10. Проверка заблуждения: превращает ли само вычитание максимума логиты в вероятности?
Проверить прогнозы
  1. Максимум равен 1001, поэтому сдвинутые значения в точности равны [-1,0].
  2. Нет. Одна и та же прибавка изменяет максимум на такую же величину, поэтому обе разности после вычитания остаются прежними.
  3. Обе исходные экспоненты переполняются до бесконечности, поэтому каждая вероятность вычисляется как отношение бесконечности к бесконечности.
  4. У одинаковых сдвинутых логитов равны экспоненты, поэтому обе вероятности равны 0.5.
  5. Логарифм вероятности класса ноль равен -1.313261687518, поэтому значение NLL равно 1.313261687518.
  6. После удаления оси с индексом 1 получается [2,4], а после её сохранения — [2,1,4].
  7. Log-sum-exp возвращает нейтральный элемент логарифмического сложения — отрицательную бесконечность. Для распределений softmax и log-softmax, а также для потери по целевому индексу нужен хотя бы один класс.
  8. После удаления оси классов 1 остаётся шаг групп [2], поэтому базовые смещения равны [0,2,4]. Шаг по оси классов равен 1, и три группы читают элементы со смещениями [0,1], [2,3] и [4,5].
  9. Для обоих результатов используются уже вычисленные максимум mm, сумма сдвинутых экспонент SS и логарифм нормирующей суммы lnS\ln S. За один проход по классам из этих величин можно сформировать оба запрошенных результата. Повторно искать максимум или сумму сдвинутых экспонент не нужно, но каждый класс всё равно требуется посетить, чтобы записать значения результатов.
  10. Нет. Вычитание максимума лишь делает вычисление относительных различий между логитами численно устойчивее; само вычитание не превращает эти значения в вероятности. Вероятности получаются после экспоненцирования и деления на сумму всех сдвинутых экспонент.

После прогноза запустите пример:

./course run cargo run --quiet --locked -p ch12-stable-softmax

Подготовьте выборочную сверку градиентов

Теперь тензорное ядро может по любой явно заданной оси преобразовывать конечные логиты из представлений с произвольными шагами в тензоры с собственным хранилищем, содержащие вероятности, логарифмы вероятностей или значения log-sum-exp, а также выполнять совмещённое вычисление среднего NLL по индексам целевых классов. Эти операции будут нормировать оценки по словарю и в механизме внимания, а вычисленное ими на прямом проходе среднее NLL по индексам глава 13 использует для выборочной сверки конечными разностями с отдельным аналитическим путём. Аналитический и численный пути всё ещё используют одни и те же логиты примера и целевые индексы, арифметику IEEE f64 и элементарную функцию exp, хранилище Tensor и соглашения о построчной индексации. Поэтому совпадение служит свидетельством для выбранных точек при локальной гладкости целевой функции, а не доказательством полного градиента или всех общих предпосылок.

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

При вычислении конечных разностей по очереди изменяется один скалярный логит; целевая функция должна быть гладкой на каждом интервале между точками вычисления. Оба пути используют одни и те же логиты примера и целевые индексы, арифметику IEEE f64 и элементарную функцию exp, хранилище Tensor и соглашения о построчной индексации. Поэтому совпадение в выбранных координатах служит свидетельством только для этого примера, этих точек, шага и допуска. Оно не доказывает дифференцируемость, правильность всего градиента или всей вероятностной реализации и не проверяет общие предпосылки двух путей.