36 · Версия материала 5
Сформируйте набор вариантов, затем сделайте один случайный выбор
Разберитесь, как положительная температура, фильтрация top-k с однозначным порядком и восстановленное состояние генератора псевдослучайных чисел превращают логиты декодера в управляемую и воспроизводимую генерацию LLM без кэша.
Начните с четырёх логитов и сделайте каждый выбор явным
Декодер выдаёт логиты последней позиции в порядке ID токенов: . Если сначала упорядочить их по убыванию логита, а равные значения — по возрастанию ID, получится . Токен однозначно занимает первое место. У токенов и равные логиты, поэтому эта реализация ставит меньший ID раньше. Благодаря этому локальному правилу порядок кандидатов при равенстве остаётся детерминированным.
Сначала оставим все четыре токена и будем менять только положительную температуру. При токен получает вероятность ; при — ; при — . Низкая температура усиливает уже существующие различия между логитами, а высокая сглаживает их. Порядок токенов от этого не меняется.
Теперь зададим и . Остаются токены в порядке . Токен остаётся по нужную сторону границы равных логитов и получает ; токен исключается, поэтому в точности. Для токена получаем . Сумма вероятностей оставленных токенов равна .
При генератор SplitMix64 с начальным значением воспроизводимо выдаёт последовательность . Например, второе равномерное случайное число равно . Оно попадает в полуоткрытый интервал токена , поэтому выбирается токен .
Отфильтруйте кандидатов, затем нормализуйте заново
Для конечных логитов, конечной положительной температуры и итоговая вероятность токена равна
Перед вычислением экспоненты реализация вычитает наибольший оставленный масштабированный логит. Это не меняет отношений вероятностей, но не даёт большому общему смещению вызвать переполнение. Вероятность исключённых ID остаётся точно равной , а не просто очень малой.
Температура должна быть строго положительной. Математический предел помогает понять концентрацию около единственного максимума, но буквальное значение привело бы к делению на ноль. Поэтому жадное декодирование оформлено как отдельное правило: оно проверяет те же логиты, выбирает первый токен в однозначном порядке и не запрашивает случайное число. Стохастический режим с выбирает тот же ID, но намеренно запрашивает одно случайное число, поскольку по-прежнему следует правилу случайного выбора.
Не смешивайте логиты, кандидатов и вероятности
- — итоговая вероятность токена после всех трёх операций.
- — конечная положительная температура. Меньшие значения усиливают различия, большие — сглаживают.
- — точное число оставленных ID токенов.
- — размер словаря, поэтому допустимое число кандидатов удовлетворяет условию .
- — множество top- с однозначным порядком. В этой реализации равные логиты упорядочиваются по возрастанию ID токена.
- равен для оставленного ID и для исключённого.
- — логит декодера для токена в последней позиции префикса.
- обозначает кандидата, чья итоговая вероятность описывается.
- перебирает словарь в знаменателе.
- превращает каждый сдвинутый масштабированный логит в положительный вес softmax.
Эти понятия разделяют важные этапы. Логит — неограниченная оценка декодера, место в однозначно заданном порядке определяет принадлежность множеству, а вероятность — нормализованная масса. Температура меняет разности оценок до нормализации. Top-k меняет состав слагаемых в знаменателе. Случайный выбор из категориального распределения выполняется только после обоих решений.
Алгоритм обходит массу оставленных токенов по возрастанию ID, используя полуоткрытые интервалы . Равномерное случайное число выбирает токен тогда и только тогда, когда . Здесь — нижняя граница накопленной вероятности токена , а — верхняя. Последний интервал допускает крошечное расхождение границы из-за суммирования чисел с плавающей точкой; исключённый токен от этого не возвращается.
От поиска при жёстких ограничениях к свободной генерации LLM
Лучевое декодирование, максимизирующее правдоподобие, полезно, когда исходный текст жёстко ограничивает результат. Но у свободного продолжения много допустимых вариантов: лучевой поиск может давать шаблонный или повторяющийся текст, а случайный выбор без ограничения кандидатов способен захватить ненадёжный хвост маловероятных токенов.
Фэн, Льюис и Дофин наглядно показывают это различие для свободной генерации. На каждом шаге они выбирают случайное слово среди десяти наиболее вероятных, настраивают температуру softmax во время генерации и сообщают, что в их задаче продолжения историй такой подход работает лучше лучевого поиска. Выбор без ограничений, напротив, может добавить редкое слово, которое серьёзно ухудшит текст.
Отчёт о GPT-2 показывает это правило уже в крупной языковой модели Transformer. В одном режиме составления краткого изложения используется случайный выбор top-k с , а для свободных продолжений WebText — с . Это пример усечённого стохастического декодирования, но не утверждение об универсальности одного значения .
Хольцман и соавторы разбирают компромисс непосредственно. Они сравнивают на GPT-2 декодирование с максимизацией и со случайным выбором, определяют top-k как выбор среди самых вероятных токенов после повторной нормализации, приводят softmax с температурой и показывают, почему фиксированное даёт несовершенный компромисс как при равномерном, так и при сильно концентрированном распределении.
Итак, системы генерации историй сочетали температуру softmax с top-k, GPT-2 применяла top-k к кратким изложениям и длинным продолжениям, а последующий анализ GPT-2 явно описал компромисс между отсечением и повторной нормализацией и показал, почему одно фиксированное не подходит для любого контекста.
Управляемое стохастическое декодирование превращает распределение авторегрессионной LLM в настраиваемый компромисс между разнообразием и концентрацией. Если восстановить состояние генератора псевдослучайных чисел и сохранить однозначные правила разрешения равенств и обхода интервалов, выбор можно воспроизвести. Более поздние методы заменяют фиксированное число кандидатов, не меняя логиты декодера.
Поэтому top-k — понятный первый пример управляемого стохастического декодера, но не универсальная гарантия качества, не защита от галлюцинаций и не конечная точка исследований декодирования. Разрешение равных логитов по ID токена, восстановление состояния генератора, поведение при ошибочных настройках, EOS и остановка по ёмкости контекста — правила воспроизводимости этой реализации, а не утверждения о процитированных системах.
Исполняемый пример измеряет механизм на четырёх логитах этой главы. Жадный режим выбирает токен , не изменяя состояние генератора. При и фиксированном остаются ID . До повторной нормализации на них приходилось массы полного softmax, а отсечение удалило . Это измерение не определяет, какое правило даёт лучший текст; оно лишь делает видимой границу множества фиксированной мощности.
rust/demos/ch36-temperature-top-k/src/lib.rs#historical-decoding-contrast /// Measures how greedy choice and fixed top-k truncation differ on one logit row.
pub fn historical_decoding_contrast() -> Result<HistoricalDecodingContrast, FixtureError> {
let full = sampling_distribution(
&LOGITS,
SamplingMode::TemperatureTopK {
temperature: 1.0,
top_k: LOGITS.len(),
},
)?;
let truncated = sampling_distribution(
&LOGITS,
SamplingMode::TemperatureTopK {
temperature: 1.0,
top_k: SAMPLE_TOP_K,
},
)?;
let retained_token_ids = truncated.survivors().to_vec();
let retained_full_probability_mass = full
.candidates()
.iter()
.filter(|candidate| retained_token_ids.contains(&candidate.token_id()))
.map(|candidate| candidate.probability())
.sum::<f64>();
let removed_full_probability_mass = full
.candidates()
.iter()
.filter(|candidate| !retained_token_ids.contains(&candidate.token_id()))
.map(|candidate| candidate.probability())
.sum::<f64>();
let mut greedy_rng = SplitMix64::from_seed(SAMPLE_SEED);
let initial_rng_state = greedy_rng.state();
let greedy = sample_next_token(&LOGITS, SamplingMode::Greedy, &mut greedy_rng)?;
require(
retained_token_ids.len() == SAMPLE_TOP_K,
"historical top-k candidate count changed",
)?;
require(
(retained_full_probability_mass + removed_full_probability_mass - 1.0).abs() < 1e-12,
"historical full-softmax mass no longer sums to one",
)?;
Ok(HistoricalDecodingContrast {
greedy_token: greedy.token_id(),
greedy_rng_advanced: greedy_rng.state() != initial_rng_state,
top_k: SAMPLE_TOP_K,
retained_token_ids,
retained_full_probability_mass,
removed_full_probability_mass,
})
} Сначала проверьте входы, однозначно задайте порядок и лишь затем измените состояние генератора
SamplingMode::Greedy и SamplingMode::TemperatureTopK задают два явных режима.
Оба отклоняют пустой набор логитов и неконечные значения. Стохастический режим
дополнительно требует конечную положительную температуру и в диапазоне
. Все проверки завершаются до изменения состояния SplitMix64.
Все три функции начинают с одних и тех же вычислений: они проверяют входные
данные, сортируют логиты с конечными значениями по убыванию, а при равенстве — по
возрастанию ID, оставляют заданное число кандидатов и преобразуют их логиты в
нормализованные вероятности. sample_next_token возвращает выбранный ID токена,
соответствующий полуоткрытый интервал и, для стохастического режима, равномерное
случайное число. Эта функция не формирует записи по каждому токену и отдельный
список прошедших фильтр токенов в порядке ранжирования, которые нужны для
подробного разбора. sample_next_token_with_trace использует те же рассчитанные
вероятности и тот же код выбора по интервалам, но перед обращением к генератору
случайных чисел формирует полное распределение для подробного разбора.
sampling_distribution формирует такое же распределение, не выполняя случайный
выбор. Обычному вызову всё равно нужны временные массивы с ID токенов в порядке
ранжирования и вероятностями. Разница в том, что дополнительные записи для
подробного разбора не создаются.
В распределении для подробного разбора кандидаты расположены по возрастанию ID токена, а отдельный список прошедших фильтр токенов — по убыванию логита; при равных логитах первым идёт меньший ID. После нормализации со сдвигом на максимум исключённым ID соответствует ровно нулевая вероятность. Случайный выбор обходит положительные интервалы по возрастанию ID. Жадный режим не запрашивает случайное число, а любой допустимый стохастический вызов запрашивает ровно одно, в том числе при top-k с .
rust/crates/llm-from-scratch/src/generation/sampling.rs#sampling-policy /// The two intentionally distinct next-token policies taught in Chapter 36.
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum SamplingMode {
/// Select the highest logit, resolving equal logits by lower token ID.
Greedy,
/// Sample after positive-temperature scaling and stable top-k truncation.
TemperatureTopK { temperature: f64, top_k: usize },
}
/// One vocabulary entry after ranking and probability construction.
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct SamplingCandidate {
token_id: u32,
logit: f64,
rank: usize,
retained: bool,
probability: f64,
}
impl SamplingCandidate {
pub const fn token_id(self) -> u32 {
self.token_id
}
pub const fn logit(self) -> f64 {
self.logit
}
/// One-based rank under descending logit and ascending token-ID ties.
pub const fn rank(self) -> usize {
self.rank
}
pub const fn retained(self) -> bool {
self.retained
}
pub const fn probability(self) -> f64 {
self.probability
}
}
/// The complete token-ID-ordered distribution plus rank-ordered survivors.
#[derive(Clone, Debug, PartialEq)]
pub struct SamplingDistribution {
mode: SamplingMode,
candidates: Vec<SamplingCandidate>,
survivors: Vec<u32>,
}
impl SamplingDistribution {
pub const fn mode(&self) -> SamplingMode {
self.mode
}
/// Candidates are always returned in ascending token-ID order.
pub fn candidates(&self) -> &[SamplingCandidate] {
&self.candidates
}
/// Survivors are returned in stable descending-logit rank order.
pub fn survivors(&self) -> &[u32] {
&self.survivors
}
pub fn probability_sum(&self) -> f64 {
compensated_sum(
self.candidates
.iter()
.map(|candidate| candidate.probability),
)
}
pub fn candidate(&self, token_id: u32) -> Option<SamplingCandidate> {
self.candidates
.get(usize::try_from(token_id).ok()?)
.copied()
}
}
/// One selected token and the half-open categorical interval that selected it.
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct SampledToken {
token_id: u32,
unit_draw: Option<f64>,
interval_start: f64,
interval_end: f64,
}
impl SampledToken {
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
}
}
/// A compact selection paired with the complete distribution used to produce it.
#[derive(Clone, Debug, PartialEq)]
pub struct SamplingDecision {
sampled: SampledToken,
distribution: SamplingDistribution,
}
impl SamplingDecision {
pub const fn token_id(&self) -> u32 {
self.sampled.token_id()
}
pub const fn unit_draw(&self) -> Option<f64> {
self.sampled.unit_draw()
}
pub const fn interval_start(&self) -> f64 {
self.sampled.interval_start()
}
pub const fn interval_end(&self) -> f64 {
self.sampled.interval_end()
}
pub const fn distribution(&self) -> &SamplingDistribution {
&self.distribution
}
}
/// A setting or numerical input that cannot define a sampling distribution.
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum SamplingError {
EmptyLogits,
VocabularyTooLarge {
classes: usize,
},
NonFiniteLogit {
token_id: usize,
value: f64,
},
InvalidTemperature {
value: f64,
},
InvalidTopK {
top_k: usize,
vocabulary_size: usize,
},
AllocationFailed {
values: usize,
},
InvalidProbabilitySum {
value: f64,
},
}
impl fmt::Display for SamplingError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::EmptyLogits => formatter.write_str("sampling needs at least one finite logit"),
Self::VocabularyTooLarge { classes } => write!(
formatter,
"sampling vocabulary of {classes} classes does not fit u32 token IDs"
),
Self::NonFiniteLogit { token_id, value } => {
write!(
formatter,
"logit for token {token_id} is not finite: {value}"
)
}
Self::InvalidTemperature { value } => write!(
formatter,
"sampling temperature must be finite and positive, received {value}"
),
Self::InvalidTopK {
top_k,
vocabulary_size,
} => write!(
formatter,
"top-k must be in 1..={vocabulary_size}, received {top_k}"
),
Self::AllocationFailed { values } => {
write!(
formatter,
"cannot allocate sampling evidence for {values} values"
)
}
Self::InvalidProbabilitySum { value } => write!(
formatter,
"sampling probabilities did not normalize to one: {value}"
),
}
}
}
impl Error for SamplingError {}
fn retained_count(mode: SamplingMode, vocabulary_size: usize) -> Result<usize, SamplingError> {
match mode {
SamplingMode::Greedy => Ok(1),
SamplingMode::TemperatureTopK { temperature, top_k } => {
if !temperature.is_finite() || temperature <= 0.0 {
return Err(SamplingError::InvalidTemperature { value: temperature });
}
if top_k == 0 || top_k > vocabulary_size {
return Err(SamplingError::InvalidTopK {
top_k,
vocabulary_size,
});
}
Ok(top_k)
}
}
}
fn compensated_sum(values: impl IntoIterator<Item = f64>) -> f64 {
let mut sum = 0.0;
let mut compensation = 0.0;
for value in values {
let corrected = value - compensation;
let next = sum + corrected;
compensation = (next - sum) - corrected;
sum = next;
}
sum
}
fn stable_ranks(logits: &[f64]) -> Result<Vec<usize>, SamplingError> {
if logits.is_empty() {
return Err(SamplingError::EmptyLogits);
}
if u32::try_from(logits.len() - 1).is_err() {
return Err(SamplingError::VocabularyTooLarge {
classes: logits.len(),
});
}
for (token_id, &value) in logits.iter().enumerate() {
if !value.is_finite() {
return Err(SamplingError::NonFiniteLogit { token_id, value });
}
}
let mut ranked = Vec::new();
ranked
.try_reserve_exact(logits.len())
.map_err(|_| SamplingError::AllocationFailed {
values: logits.len(),
})?;
ranked.extend(0..logits.len());
ranked.sort_unstable_by(|&left, &right| {
logits[right]
.partial_cmp(&logits[left])
.unwrap_or(Ordering::Equal)
.then_with(|| left.cmp(&right))
});
Ok(ranked)
}
fn scaled_gap(logit: f64, maximum: f64, temperature: f64) -> f64 {
if temperature < 1.0 {
(logit - maximum) / temperature
} else {
logit / temperature - maximum / temperature
}
}
struct PreparedSampling {
mode: SamplingMode,
ranked: Vec<usize>,
probabilities: Vec<f64>,
keep: usize,
}
fn prepare_sampling(logits: &[f64], mode: SamplingMode) -> Result<PreparedSampling, SamplingError> {
let ranked = stable_ranks(logits)?;
let keep = retained_count(mode, logits.len())?;
let maximum = logits[ranked[0]];
let mut probabilities = Vec::new();
probabilities
.try_reserve_exact(logits.len())
.map_err(|_| SamplingError::AllocationFailed {
values: logits.len(),
})?;
probabilities.resize(logits.len(), 0.0);
match mode {
SamplingMode::Greedy => probabilities[ranked[0]] = 1.0,
SamplingMode::TemperatureTopK { temperature, .. } => {
for (position, &token_id) in ranked[..keep].iter().enumerate() {
let weight = if position == 0 {
1.0
} else {
scaled_gap(logits[token_id], maximum, temperature).exp()
};
probabilities[token_id] = weight;
}
let weight_sum = compensated_sum(
ranked[..keep]
.iter()
.map(|&token_id| probabilities[token_id]),
);
if !weight_sum.is_finite() || weight_sum <= 0.0 {
return Err(SamplingError::InvalidProbabilitySum { value: weight_sum });
}
for &token_id in &ranked[..keep] {
probabilities[token_id] /= weight_sum;
}
}
}
let probability_sum = compensated_sum(probabilities.iter().copied());
if !probability_sum.is_finite() || (probability_sum - 1.0).abs() > PROBABILITY_TOLERANCE {
return Err(SamplingError::InvalidProbabilitySum {
value: probability_sum,
});
}
Ok(PreparedSampling {
mode,
ranked,
probabilities,
keep,
})
}
fn materialize_distribution(
logits: &[f64],
prepared: &PreparedSampling,
) -> Result<SamplingDistribution, SamplingError> {
let PreparedSampling {
mode,
ranked,
probabilities,
keep,
} = prepared;
let mut rank_by_token = Vec::new();
rank_by_token
.try_reserve_exact(logits.len())
.map_err(|_| SamplingError::AllocationFailed {
values: logits.len(),
})?;
rank_by_token.resize(logits.len(), 0);
for (position, &token_id) in ranked.iter().enumerate() {
rank_by_token[token_id] = position + 1;
}
let mut candidates = Vec::new();
candidates
.try_reserve_exact(logits.len())
.map_err(|_| SamplingError::AllocationFailed {
values: logits.len(),
})?;
for (token_id, ((&logit, &probability), &rank)) in logits
.iter()
.zip(probabilities)
.zip(&rank_by_token)
.enumerate()
{
candidates.push(SamplingCandidate {
token_id: u32::try_from(token_id).expect("validated token ID must fit u32"),
logit,
rank,
retained: rank <= *keep,
probability,
});
}
let mut survivors = Vec::new();
survivors
.try_reserve_exact(*keep)
.map_err(|_| SamplingError::AllocationFailed { values: *keep })?;
for &token_id in &ranked[..*keep] {
survivors.push(u32::try_from(token_id).expect("validated token ID must fit u32"));
}
Ok(SamplingDistribution {
mode: *mode,
candidates,
survivors,
})
}
fn select_prepared(prepared: &PreparedSampling, rng: &mut SplitMix64) -> SampledToken {
if prepared.mode == SamplingMode::Greedy {
return SampledToken {
token_id: u32::try_from(prepared.ranked[0]).expect("validated token ID must fit u32"),
unit_draw: None,
interval_start: 0.0,
interval_end: 1.0,
};
}
let draw = rng.next_unit_f64();
let final_id = prepared
.probabilities
.iter()
.enumerate()
.rev()
.find(|(_, probability)| **probability > 0.0)
.map(|(token_id, _)| token_id)
.expect("a normalized distribution must have a positive survivor");
let mut start = 0.0;
for (token_id, &probability) in prepared
.probabilities
.iter()
.enumerate()
.filter(|(_, probability)| **probability > 0.0)
{
let end = if token_id == final_id {
1.0
} else {
(start + probability).min(1.0)
};
if draw < end || token_id == final_id {
return SampledToken {
token_id: u32::try_from(token_id).expect("validated token ID must fit u32"),
unit_draw: Some(draw),
interval_start: start,
interval_end: end,
};
}
start = end;
}
unreachable!("the final positive survivor covers every unit draw")
}
fn sample_with_observer<T>(
logits: &[f64],
mode: SamplingMode,
rng: &mut SplitMix64,
observe: impl FnOnce(&PreparedSampling) -> Result<T, SamplingError>,
) -> Result<(SampledToken, T), SamplingError> {
let prepared = prepare_sampling(logits, mode)?;
let observation = observe(&prepared)?;
let sampled = select_prepared(&prepared, rng);
Ok((sampled, observation))
}
/// Builds a complete, inspectable distribution without consuming randomness.
pub fn sampling_distribution(
logits: &[f64],
mode: SamplingMode,
) -> Result<SamplingDistribution, SamplingError> {
let prepared = prepare_sampling(logits, mode)?;
materialize_distribution(logits, &prepared)
}
/// Selects one token without retaining the complete inspectable distribution.
pub fn sample_next_token(
logits: &[f64],
mode: SamplingMode,
rng: &mut SplitMix64,
) -> Result<SampledToken, SamplingError> {
sample_with_observer(logits, mode, rng, |_| Ok(())).map(|(sampled, ())| sampled)
}
/// Selects one token and records the complete distribution used for inspection.
pub fn sample_next_token_with_trace(
logits: &[f64],
mode: SamplingMode,
rng: &mut SplitMix64,
) -> Result<SamplingDecision, SamplingError> {
let (sampled, distribution) = sample_with_observer(logits, mode, rng, |prepared| {
materialize_distribution(logits, prepared)
})?;
Ok(SamplingDecision {
sampled,
distribution,
})
} generate_uncached хранит полный префикс. Перед каждым вызовом декодера функция
проверяет, помещается ли префикс в заданный контекст. Затем она выполняет
вычисления без записи графа вычислений, берёт последнюю строку по оси словаря,
выбирает один токен, добавляет его и сохраняет только краткие сведения о шаге:
длину префикса, выбранный ID, случайное число и интервал. Поскольку функция
вызывает вариант без трассировки, для каждого нового токена она не создаёт и не
хранит полные записи о кандидатах. EOS включается в результат до остановки цикла.
Поэтому допустимый префикс длины при ёмкости может выдать ещё один токен,
хотя получившуюся последовательность длины уже нельзя подать в следующем
вызове.
rust/crates/llm-from-scratch/src/generation/sampling.rs#uncached-generation /// Settings for one bounded, uncached autoregressive call.
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct GenerationConfig {
mode: SamplingMode,
eos_token: Option<u32>,
max_new_tokens: usize,
}
impl GenerationConfig {
pub const fn new(mode: SamplingMode, eos_token: Option<u32>, max_new_tokens: usize) -> Self {
Self {
mode,
eos_token,
max_new_tokens,
}
}
pub const fn mode(self) -> SamplingMode {
self.mode
}
pub const fn eos_token(self) -> Option<u32> {
self.eos_token
}
pub const fn max_new_tokens(self) -> usize {
self.max_new_tokens
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum GenerationStop {
Eos,
TokenLimit,
ContextLimit,
}
#[derive(Clone, Debug, PartialEq)]
pub struct GenerationStep {
prefix_length: usize,
token_id: u32,
unit_draw: Option<f64>,
interval_start: f64,
interval_end: f64,
}
impl GenerationStep {
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
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct GenerationResult {
prompt: Vec<u32>,
generated: Vec<u32>,
steps: Vec<GenerationStep>,
stop: GenerationStop,
full_prefix_calls: usize,
}
impl GenerationResult {
pub fn prompt(&self) -> &[u32] {
&self.prompt
}
pub fn generated(&self) -> &[u32] {
&self.generated
}
pub fn steps(&self) -> &[GenerationStep] {
&self.steps
}
pub const fn stop(&self) -> GenerationStop {
self.stop
}
pub const fn full_prefix_calls(&self) -> usize {
self.full_prefix_calls
}
}
#[derive(Debug)]
pub enum GenerationError {
Sampling(SamplingError),
Model(DecoderModelError),
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,
},
}
impl fmt::Display for GenerationError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Sampling(error) => error.fmt(formatter),
Self::Model(error) => error.fmt(formatter),
Self::EmptyPrompt => formatter.write_str("generation needs a nonempty prompt"),
Self::PromptTooLong {
tokens,
max_positions,
} => write!(
formatter,
"prompt has {tokens} tokens, exceeding context capacity {max_positions}"
),
Self::PromptTokenOutOfBounds {
position,
token_id,
vocabulary_size,
} => write!(
formatter,
"prompt token {token_id} at position {position} is out of bounds for vocabulary {vocabulary_size}"
),
Self::EosTokenOutOfBounds {
token_id,
vocabulary_size,
} => write!(
formatter,
"EOS token {token_id} is out of bounds for vocabulary {vocabulary_size}"
),
Self::LogitCountMismatch { expected, actual } => write!(
formatter,
"last-position logits need {expected} values, received {actual}"
),
Self::AllocationFailed { values } => {
write!(
formatter,
"cannot allocate generation evidence for {values} values"
)
}
}
}
}
impl Error for GenerationError {
fn source(&self) -> Option<&(dyn Error + 'static)> {
match self {
Self::Sampling(error) => Some(error),
Self::Model(error) => Some(error),
_ => None,
}
}
}
impl From<SamplingError> for GenerationError {
fn from(error: SamplingError) -> Self {
Self::Sampling(error)
}
}
impl From<DecoderModelError> for GenerationError {
fn from(error: DecoderModelError) -> Self {
Self::Model(error)
}
}
fn generate_with<F>(
vocabulary_size: usize,
max_positions: usize,
prompt: &[u32],
config: GenerationConfig,
rng: &mut SplitMix64,
mut last_logits: F,
) -> Result<GenerationResult, GenerationError>
where
F: FnMut(&[u32]) -> Result<Vec<f64>, GenerationError>,
{
if prompt.is_empty() {
return Err(GenerationError::EmptyPrompt);
}
if prompt.len() > max_positions {
return Err(GenerationError::PromptTooLong {
tokens: prompt.len(),
max_positions,
});
}
for (position, &token_id) in prompt.iter().enumerate() {
if usize::try_from(token_id)
.ok()
.is_none_or(|token| token >= vocabulary_size)
{
return Err(GenerationError::PromptTokenOutOfBounds {
position,
token_id,
vocabulary_size,
});
}
}
if let Some(token_id) = config.eos_token
&& usize::try_from(token_id)
.ok()
.is_none_or(|token| token >= vocabulary_size)
{
return Err(GenerationError::EosTokenOutOfBounds {
token_id,
vocabulary_size,
});
}
retained_count(config.mode, vocabulary_size)?;
let planned_steps = config.max_new_tokens.min(
max_positions
.checked_sub(prompt.len())
.and_then(|remaining| remaining.checked_add(1))
.ok_or(GenerationError::AllocationFailed { values: usize::MAX })?,
);
let prefix_capacity = prompt
.len()
.checked_add(planned_steps)
.ok_or(GenerationError::AllocationFailed { values: usize::MAX })?;
let mut prompt_copy = Vec::new();
prompt_copy
.try_reserve_exact(prompt.len())
.map_err(|_| GenerationError::AllocationFailed {
values: prompt.len(),
})?;
prompt_copy.extend_from_slice(prompt);
let mut prefix = Vec::new();
prefix
.try_reserve_exact(prefix_capacity)
.map_err(|_| GenerationError::AllocationFailed {
values: prefix_capacity,
})?;
prefix.extend_from_slice(prompt);
let mut generated = Vec::new();
generated
.try_reserve_exact(planned_steps)
.map_err(|_| GenerationError::AllocationFailed {
values: planned_steps,
})?;
let mut steps = Vec::new();
steps
.try_reserve_exact(planned_steps)
.map_err(|_| GenerationError::AllocationFailed {
values: planned_steps,
})?;
let mut full_prefix_calls = 0usize;
if config.max_new_tokens == 0 {
return Ok(GenerationResult {
prompt: prompt_copy,
generated,
steps,
stop: GenerationStop::TokenLimit,
full_prefix_calls,
});
}
let stop = loop {
let prefix_length = prefix.len();
let logits = last_logits(&prefix)?;
full_prefix_calls += 1;
if logits.len() != vocabulary_size {
return Err(GenerationError::LogitCountMismatch {
expected: vocabulary_size,
actual: logits.len(),
});
}
let decision = sample_next_token(&logits, config.mode, rng)?;
let token_id = decision.token_id();
let unit_draw = decision.unit_draw();
let interval_start = decision.interval_start();
let interval_end = decision.interval_end();
prefix.push(token_id);
generated.push(token_id);
steps.push(GenerationStep {
prefix_length,
token_id,
unit_draw,
interval_start,
interval_end,
});
if config.eos_token == Some(token_id) {
break GenerationStop::Eos;
}
if generated.len() == config.max_new_tokens {
break GenerationStop::TokenLimit;
}
if prefix.len() > max_positions {
break GenerationStop::ContextLimit;
}
};
Ok(GenerationResult {
prompt: prompt_copy,
generated,
steps,
stop,
full_prefix_calls,
})
}
/// Recomputes the complete decoder prefix for every selected token.
pub fn generate_uncached(
model: &DecoderModel,
prompt: &[u32],
config: GenerationConfig,
rng: &mut SplitMix64,
) -> Result<GenerationResult, GenerationError> {
let model_config = model.config();
let vocabulary_size = model_config.vocabulary_size();
generate_with(
vocabulary_size,
model_config.max_positions(),
prompt,
config,
rng,
|prefix| {
let forward = no_grad(|| model.forward(prefix, &[1, prefix.len()]))?;
let logits = forward.logits().value();
let expected = prefix
.len()
.checked_mul(vocabulary_size)
.ok_or(GenerationError::AllocationFailed { values: usize::MAX })?;
if logits.len() != expected {
return Err(GenerationError::LogitCountMismatch {
expected,
actual: logits.len(),
});
}
let start = expected - vocabulary_size;
let mut final_logits = Vec::new();
final_logits
.try_reserve_exact(vocabulary_size)
.map_err(|_| GenerationError::AllocationFailed {
values: vocabulary_size,
})?;
final_logits.extend_from_slice(&logits.as_slice()[start..]);
Ok(final_logits)
},
)
} Пример загружает точные байты контрольной точки из главы 35. Сначала он
сохраняет отдельно записанное в ней состояние генератора псевдослучайных чисел.
Затем метод into_model получает контрольную точку по значению и перемещает в
декодер буферы модели, которыми она владеет; после этого объект контрольной
точки больше не нужен. Эта передача владения не меняет ни значения параметров,
ни правило случайного выбора. Промпт порождает последовательность
при длинах префикса : декодер дважды обрабатывает полный префикс, после
чего генерация останавливается из-за исчерпания ёмкости контекста. Если снова
начать с сохранённого состояния генератора, совпадут и сгенерированная
последовательность, и конечное состояние генератора. Если назначить токен
токеном EOS, результат включает и останавливается после одного вызова
декодера.
rust/demos/ch36-temperature-top-k/src/lib.rs#learner-evidence /// Loads the Chapter 35 checkpoint and records sampling plus full-prefix stops.
pub fn learner_evidence() -> Result<LearnerEvidence, FixtureError> {
let temperatures = temperature_evidence()?;
let boundary = sampling_distribution(
&LOGITS,
SamplingMode::TemperatureTopK {
temperature: 1.0,
top_k: BOUNDARY_TOP_K,
},
)?;
let mut greedy_rng = SplitMix64::from_seed(SAMPLE_SEED);
let greedy_state = greedy_rng.state();
let greedy = sample_next_token_with_trace(&LOGITS, SamplingMode::Greedy, &mut greedy_rng)?;
require(
greedy_rng.state() == greedy_state,
"greedy sampling unexpectedly consumed RNG state",
)?;
let seeded_decisions = seeded_decisions()?;
let prior = checkpoint_evidence()?;
let loaded_checkpoint_bytes = prior.encoded.bytes().len();
let checkpoint = Checkpoint::from_bytes(prior.encoded.bytes())?;
let loaded_rng_state = checkpoint.rng_state();
let model = checkpoint.into_model()?;
let model_config = model.config();
let loaded_vocabulary_size = model_config.vocabulary_size();
let loaded_context = model_config.max_positions();
let generation_config = GenerationConfig::new(
SamplingMode::TemperatureTopK {
temperature: 1.0,
top_k: 3,
},
None,
4,
);
let mut loaded_rng = SplitMix64::from_state(loaded_rng_state);
let loaded = generate_uncached(&model, &LOADED_PROMPT, generation_config, &mut loaded_rng)?;
let mut replay_rng = SplitMix64::from_state(loaded_rng_state);
let replay = generate_uncached(&model, &LOADED_PROMPT, generation_config, &mut replay_rng)?;
let loaded_replay_identical = loaded == replay && loaded_rng.state() == replay_rng.state();
let first_token = *loaded.generated().first().ok_or(FixtureError::Invariant(
"loaded generation emitted no token",
))?;
let mut eos_rng = SplitMix64::from_state(loaded_rng_state);
let eos = generate_uncached(
&model,
&LOADED_PROMPT,
GenerationConfig::new(generation_config.mode(), Some(first_token), 4),
&mut eos_rng,
)?;
let errors = error_evidence();
let history = historical_decoding_contrast()?;
require(
boundary.survivors() == [3, 1],
"stable tied top-k boundary changed",
)?;
require(
seeded_decisions
.iter()
.map(SamplingDecision::token_id)
.eq([3, 2, 2, 2, 3, 3, 3, 3]),
"seeded sampling sequence changed",
)?;
require(
loaded
.steps()
.iter()
.map(|step| step.prefix_length())
.eq([1, 2])
&& loaded.full_prefix_calls() == 2
&& loaded.stop() == GenerationStop::ContextLimit,
"loaded uncached context evidence changed",
)?;
require(
loaded_replay_identical,
"loaded checkpoint and RNG no longer replay generation",
)?;
require(
eos.generated() == [first_token]
&& eos.stop() == GenerationStop::Eos
&& eos.full_prefix_calls() == 1,
"EOS stopping evidence changed",
)?;
require(
errors.zero_temperature_rejected
&& errors.zero_top_k_rejected
&& errors.nonfinite_logit_rejected
&& errors.rng_unchanged,
"invalid settings changed RNG state or escaped validation",
)?;
Ok(LearnerEvidence {
temperatures,
boundary,
greedy,
seeded_decisions,
loaded_checkpoint_bytes,
loaded_rng_state,
loaded_vocabulary_size,
loaded_context,
generation_max_new_tokens: generation_config.max_new_tokens(),
loaded,
loaded_replay_identical,
eos,
errors,
history,
})
} Выполните cargo run --quiet --locked -p ch36-temperature-top-k. Основные
результаты программа выводит напрямую:
top_k=k:2 survivors:[3,1] tied_boundary:keep:1 remove:2 sum:1.000000000000
sample=seed:36 top_k:3 sequence:[3,2,2,2,3,3,3,3] draws:8 greedy_token:3 greedy_draw:none
checkpoint=loaded_bytes:6330 rng_state:0x9e3779b97f4a7c38 vocabulary:5 context:2 eos:none max_new_tokens:4 prompt:[0] generated:[4,4] prefixes:[1,2] stop:context-limit full_prefix_calls:2 replay_identical:true
eos=vocabulary:5 context:2 eos_token:4 max_new_tokens:4 generated:[4] stop:eos full_prefix_calls:1
errors=temperature_zero:true top_k_zero:true nonfinite_logit:true rng_unchanged:true
history=greedy_token:3 greedy_rng_advanced:false top_k:3 survivors:[3,1,2] retained_full_mass:0.927670511871 removed_full_mass:0.072329488129
rust/demos/ch36-temperature-top-k/src/main.rs fn main() -> Result<(), Box<dyn std::error::Error>> {
print!("{}", ch36_temperature_top_k::learner_report()?);
Ok(())
} Прочитайте распределение слева направо
Сначала схема совмещает три представления одних и тех же четырёх логитов. Затем она отдельно показывает границу при и явно переключается на правило с , и начальным значением , после чего проводит каждое случайное число по соответствующему интервалу. Обозначенная граница примеров меняет искусственный словарь размером на словарь загруженного декодера размером . Рядом показаны запуск без EOS, который останавливается по ёмкости контекста, и запуск с токеном в роли EOS. Двойная и пунктирная рамки вместе с текстом повторяют состояния «оставлен» и «исключён». В таблице сохранены точный логит каждого токена, его место в порядке, вероятность и принадлежность top-k, чтобы полосы и рамки можно было сопоставить с числовыми данными.
Температура меняет форму, top-k отсекает, случайное число выбирает
Точная трасса из программы на Rust сравнивает три температуры, показывает границу равных логитов с однозначным порядком, восемь полуоткрытых интервалов и причины остановки генерации из контрольной точки.
- оставлен — двойная рамка
- исключён — пунктирная рамка
- выбран случайным числом
Сравните температуры для одних логитов
Искусственный пример с четырьмя токенами
Все четыре токена остаются; меняются только отношения вероятностей. На каждой полосе указано точное значение.
Острее
Исходный масштаб
Ровнее
Оставьте ровно два места в заданном порядке
При равных логитах границу определяет возрастание ID: токен 1 остаётся, а вероятность токена 2 становится точно нулевой.
| Токен | Логит | Место в порядке | Статус top-k | Итоговая вероятность |
|---|---|---|---|---|
| исключён — пунктирная рамка | ||||
| оставлен — двойная рамка | ||||
| исключён — пунктирная рамка | ||||
| оставлен — двойная рамка |
Повторите выбор из того же состояния генератора
Здесь вместо двух кандидатов используется правило с тремя кандидатами, температурой 1 и начальным значением 36. Интервалы оставленных токенов обходятся по возрастанию ID.
seed=36 survivors=[3,1,2] sum=1.000000000000 - Выбор 1 Полуоткрытый интервал выбирает токен
- Выбор 2 Полуоткрытый интервал выбирает токен
- Выбор 3 Полуоткрытый интервал выбирает токен
- Выбор 4 Полуоткрытый интервал выбирает токен
- Выбор 5 Полуоткрытый интервал выбирает токен
- Выбор 6 Полуоткрытый интервал выбирает токен
- Выбор 7 Полуоткрытый интервал выбирает токен
- Выбор 8 Полуоткрытый интервал выбирает токен
Примените правило к генерации без кэша
Здесь пример с четырьмя искусственными токенами сменяется загруженным декодером с пятью токенами. Для каждого запуска показаны точные правила EOS и ёмкости контекста.
Загруженный декодер с пятью токенами context=2
Явный жадный режим
— наибольший логит; при равенстве — меньший ID
draw=none Жадный режим не меняет состояние генератора.
Загруженная контрольная точка
eos=none max_new_tokens=4 prefixes=[1,2] calls=2 остановка при исчерпании ёмкости контекста; восстановленное состояние точно повторяет выбор
Граница EOS
eos=4 max_new_tokens=4 generated=[4] calls=1 EOS остаётся в выданной последовательности
Ошибки до изменения состояния
temperature_zero=true top_k_zero=true nonfinite_logit=true rng_unchanged=true Некорректные настройки возвращают ошибку до изменения состояния генератора.
Сначала предскажите, затем откройте трассу
- Какой ID выберет жадный режим для логитов ?
- Какой из ID с равными логитами останется при ?
- Изменит ли повышение порядок токенов?
- Почему буквальное значение отклоняется, хотя предел помогает понять поведение распределения?
- Запрашивает ли стохастический режим с случайное число?
- Какому токену принадлежит число в примере с ?
- Включается ли EOS в выданную последовательность?
- Почему префикс длины при ёмкости может выдать ещё один токен до остановки по контексту?
Проверьте восемь предсказаний
- Побеждает токен , потому что его логит — единственный максимум.
- Остаётся токен : при равных логитах ID упорядочиваются по возрастанию, поэтому вероятность токена становится точно нулевой.
- Нет. Положительная температура меняет отношения вероятностей, но сохраняет порядок логитов.
- Предел описывает концентрацию, а буквальный ноль потребовал бы деления на ноль; это детерминированное правило реализует отдельный жадный режим.
- Да. Он выбирает тот же ID, что и жадный режим, но остаётся стохастическим и запрашивает ровно одно случайное число.
- Токену принадлежит интервал , содержащий это число.
- Да. Токен добавляется до того, как сообщается об остановке на EOS.
- Допустимый префикс длины предсказывает следующий токен; ёмкость была бы превышена только при следующем вызове декодера.
Сохраните эту последовательность без кэша при переходе к поэтапной генерации
К этому этапу программа умеет загрузить выбранную контрольную точку, преобразовать строку логитов последней позиции в управляемое распределение следующего токена, повторить выбор из восстановленного состояния генератора, остановиться на EOS или при исчерпании ёмкости контекста и получить эталонную последовательность без кэша, которую глава 37 сохранит при поэтапных вычислениях.
Глава 37 будет сохранять в кэше ранее вычисленные векторы ключей и значений одного слоя внимания. Результат для новейшей позиции должен совпасть с вычислением полного префикса из этой главы, прежде чем глава 38 распространит кэширование на весь декодер.