use std::collections::HashMap;
use std::fmt;
use std::sync::Arc;
use crate::penalty_window::PenaltyWindow;
const MAX_CHAR_LEN: usize = 40;
const MAX_SEQ_LEN: usize = 20;
const FLOAT_MAX_LOG: f32 = 88.722_84;
pub const DEFAULT_SEQUENCE_BREAKERS: [&str; 4] = ["\n", ":", "\"", "*"];
pub trait DryVocab {
fn n_tokens(&self) -> usize;
fn detokenize(&self, token: usize) -> String;
fn tokenize(&self, text: &str) -> Vec<usize>;
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct DryBreakers {
raw: Vec<String>,
heads: HashMap<usize, Vec<Vec<usize>>>,
}
impl DryBreakers {
pub fn none() -> Self {
DryBreakers::default()
}
pub fn raw(&self) -> &[String] {
&self.raw
}
pub fn is_empty(&self) -> bool {
self.heads.is_empty()
}
pub fn from_vocab(vocab: &dyn DryVocab, breakers: &[String]) -> Self {
let mut out = DryBreakers {
raw: breakers.to_vec(),
heads: HashMap::new(),
};
for breaker in breakers {
if breaker.is_empty() {
continue;
}
let mut bytes = breaker.as_bytes();
if bytes.len() > MAX_CHAR_LEN {
bytes = &bytes[..MAX_CHAR_LEN];
}
out.add_overlapping(vocab, bytes);
}
out
}
pub fn from_token_sequences(sequences: &[Vec<usize>]) -> Self {
let mut heads: HashMap<usize, Vec<Vec<usize>>> = HashMap::new();
for sequence in sequences {
let Some((&head, tail)) = sequence.split_first() else {
continue;
};
let mut tail = tail.to_vec();
tail.truncate(MAX_SEQ_LEN);
let entry = heads.entry(head).or_default();
if !entry.contains(&tail) {
entry.push(tail);
}
}
DryBreakers {
raw: Vec::new(),
heads,
}
}
fn tails(&self, token: usize) -> Option<&[Vec<usize>]> {
self.heads.get(&token).map(|v| v.as_slice())
}
fn is_single_token_breaker(&self, token: usize) -> bool {
self.heads
.get(&token)
.is_some_and(|tails| tails.iter().any(|tail| tail.is_empty()))
}
fn push(&mut self, head: usize, tail: Vec<usize>) {
let entry = self.heads.entry(head).or_default();
if !entry.contains(&tail) {
entry.push(tail);
}
}
fn add_overlapping(&mut self, vocab: &dyn DryVocab, breaker: &[u8]) {
for token in 0..vocab.n_tokens() {
let word = vocab.detokenize(token);
let word = word.as_bytes();
if contains_subslice(word, breaker) {
self.push(token, Vec::new());
continue;
}
let mut from = 0usize;
while from < word.len() {
let Some(offset) = word[from..].iter().position(|&c| c == breaker[0]) else {
break;
};
let pos = from + offset;
from = pos + 1;
let mut i = 1usize;
let mut matched = true;
while i < breaker.len() && i + pos < word.len() {
if word[pos + i] != breaker[i] {
matched = false;
break;
}
i += 1;
}
if !matched {
continue;
}
let Ok(rest) = std::str::from_utf8(&breaker[i..]) else {
continue;
};
let mut tail = vocab.tokenize(rest);
tail.truncate(MAX_SEQ_LEN);
self.push(token, tail);
}
}
}
}
fn contains_subslice(haystack: &[u8], needle: &[u8]) -> bool {
if needle.is_empty() || needle.len() > haystack.len() {
return needle.is_empty();
}
haystack.windows(needle.len()).any(|w| w == needle)
}
#[derive(Debug, Clone)]
pub struct DryParams {
multiplier: f32,
base: f32,
allowed_length: i32,
penalty_last_n: i32,
total_context_size: usize,
breakers: Arc<DryBreakers>,
}
impl DryParams {
pub fn off() -> Self {
DryParams {
multiplier: 0.0,
base: 1.75,
allowed_length: 2,
penalty_last_n: -1,
total_context_size: 0,
breakers: Arc::new(DryBreakers::none()),
}
}
pub fn new(
multiplier: f32,
base: f32,
allowed_length: i32,
penalty_last_n: i32,
total_context_size: usize,
breakers: DryBreakers,
) -> Self {
DryParams {
multiplier,
base,
allowed_length,
penalty_last_n,
total_context_size,
breakers: Arc::new(breakers),
}
}
pub fn is_enabled(&self) -> bool {
self.multiplier != 0.0 && self.base >= 1.0 && self.penalty_last_n != 0
}
pub fn multiplier(&self) -> f32 {
self.multiplier
}
pub fn base(&self) -> f32 {
self.base
}
pub fn allowed_length(&self) -> i32 {
self.allowed_length
}
pub fn penalty_last_n(&self) -> i32 {
self.penalty_last_n
}
pub fn total_context_size(&self) -> usize {
self.total_context_size
}
pub fn breakers(&self) -> &DryBreakers {
&self.breakers
}
fn effective_last_n(&self) -> usize {
if self.penalty_last_n == -1 {
self.total_context_size
} else {
self.penalty_last_n.max(0) as usize
}
}
pub fn penalties(&self, history: PenaltyWindow<'_>) -> HashMap<usize, f32> {
let empty = HashMap::new();
if !self.is_enabled() {
return empty;
}
let effective = self.effective_last_n();
let recent: Vec<usize> = history
.recent(effective.min(self.total_context_size))
.collect();
let n = recent.len();
if n as i32 <= self.allowed_length {
return empty;
}
let rat = |i: usize| recent[n - 1 - i];
let mut rep_limit = n as i32;
for i in 0..n {
let Some(tails) = self.breakers.tails(rat(i)) else {
continue;
};
let mut longest_match: i32 = -1;
for tail in tails {
let seq_len = tail.len() as i32;
if seq_len <= longest_match || seq_len > i as i32 {
continue;
}
let matched = (0..tail.len()).all(|offset| tail[offset] == rat(i - offset - 1));
if matched {
longest_match = seq_len;
}
}
if longest_match >= 0 {
rep_limit = i as i32 - longest_match;
break;
}
}
if rep_limit < self.allowed_length {
return empty;
}
let mut repeat_count = vec![0i32; n];
let last = n - 1;
let mut lt = 0usize;
let mut rt = 0usize;
for k in 1..n {
if k > rt {
let mut matched = 0usize;
while matched + k < n && rat(matched) == rat(matched + k) {
matched += 1;
}
repeat_count[last - k] = (matched as i32).min(rep_limit);
if matched > 0 {
lt = k;
rt = k + matched - 1;
}
} else {
let p = k - lt;
let right_part_len = (rt - k + 1) as i32;
if repeat_count[last - p] < right_part_len {
repeat_count[last - k] = repeat_count[last - p].min(rep_limit);
} else {
let mut i = rt + 1;
while i < n && rat(i) == rat(i - k) {
i += 1;
}
repeat_count[last - k] = ((i - k) as i32).min(rep_limit);
lt = k;
rt = i - 1;
}
}
}
let mut max_token_repeat: HashMap<usize, i32> = HashMap::new();
for i in 0..n - 1 {
let repeat_len = repeat_count[i];
if repeat_len < self.allowed_length {
continue;
}
let token = recent[i + 1];
let slot = max_token_repeat.entry(token).or_insert(repeat_len);
if *slot < repeat_len {
*slot = repeat_len;
}
}
let mut max_exponent = 0i32;
if self.base > 1.000_001 {
max_exponent = (FLOAT_MAX_LOG / self.base.ln()) as i32;
}
let mut out = HashMap::with_capacity(max_token_repeat.len());
for (&token, &max_repeat) in &max_token_repeat {
if self.breakers.is_single_token_breaker(token) {
continue;
}
let mut repeat_exp = max_repeat - self.allowed_length;
if max_exponent > 0 && repeat_exp > max_exponent {
repeat_exp = max_exponent;
}
let penalty = (self.multiplier as f64) * (self.base as f64).powi(repeat_exp);
out.insert(token, penalty as f32);
}
out
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct DryRequest {
pub multiplier: f32,
pub base: f32,
pub allowed_length: i32,
pub penalty_last_n: i32,
pub sequence_breakers: Vec<String>,
}
impl Default for DryRequest {
fn default() -> Self {
DryRequest {
multiplier: 0.0,
base: 1.75,
allowed_length: 2,
penalty_last_n: -1,
sequence_breakers: DEFAULT_SEQUENCE_BREAKERS
.iter()
.map(|s| s.to_string())
.collect(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DryVocabMissing;
impl fmt::Display for DryVocabMissing {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(
"the DRY sampler needs the model's vocabulary to tokenise its sequence \
breakers, and this checkpoint has none loaded. Pass \
`--dry-sequence-breaker none` to run DRY with no breakers, or \
`--dry-multiplier 0` to switch DRY off",
)
}
}
impl std::error::Error for DryVocabMissing {}
impl DryRequest {
pub fn is_enabled(&self) -> bool {
self.multiplier != 0.0 && self.base >= 1.0 && self.penalty_last_n != 0
}
pub fn resolve(
&self,
vocab: Option<&dyn DryVocab>,
total_context_size: usize,
) -> Result<DryParams, DryVocabMissing> {
if !self.is_enabled() {
return Ok(DryParams::off());
}
let breakers = match (self.sequence_breakers.is_empty(), vocab) {
(true, _) => DryBreakers::none(),
(false, Some(vocab)) => DryBreakers::from_vocab(vocab, &self.sequence_breakers),
(false, None) => return Err(DryVocabMissing),
};
Ok(DryParams::new(
self.multiplier,
self.base,
self.allowed_length,
self.penalty_last_n,
total_context_size,
breakers,
))
}
}
#[cfg(test)]
mod tests {
use super::*;
fn dry(
multiplier: f32,
base: f32,
allowed: i32,
last_n: i32,
breakers: &[Vec<usize>],
) -> DryParams {
DryParams::new(
multiplier,
base,
allowed,
last_n,
1024,
DryBreakers::from_token_sequences(breakers),
)
}
fn penalties_of(params: &DryParams, history: &[usize]) -> Vec<(usize, f32)> {
let mut out: Vec<(usize, f32)> = params
.penalties(PenaltyWindow::new(&[], history))
.into_iter()
.collect();
out.sort_unstable_by_key(|&(token, _)| token);
out
}
#[test]
fn a_repeated_suffix_penalises_only_the_token_that_would_extend_it() {
let params = dry(1.0, 1.1, 2, 5, &[]);
assert_eq!(penalties_of(¶ms, &[0, 1, 2, 0, 1]), vec![(2, 1.0)]);
let doubled = dry(2.0, 1.1, 2, 5, &[]);
assert_eq!(penalties_of(&doubled, &[0, 1, 2, 0, 1]), vec![(2, 2.0)]);
}
#[test]
fn a_window_within_the_allowed_length_is_never_penalised() {
let params = dry(1.0, 1.1, 2, 4, &[]);
assert!(penalties_of(¶ms, &[0, 1]).is_empty());
}
#[test]
fn a_repetition_shorter_than_the_allowed_length_is_free() {
let params = dry(1.0, 1.1, 4, 7, &[]);
assert!(penalties_of(¶ms, &[0, 1, 2, 3, 4, 0, 1]).is_empty());
}
#[test]
fn the_penalty_is_multiplier_times_base_to_the_length_over_the_allowance() {
let params = dry(0.8, 1.75, 2, -1, &[]);
let got = penalties_of(¶ms, &[0, 1, 2, 3, 0, 1, 2, 3, 0, 1, 2]);
assert_eq!(got.len(), 1, "only token 3 extends the repetition: {got:?}");
assert_eq!(got[0].0, 3);
assert!(
(got[0].1 - 13.130_469).abs() < 1e-4,
"0.8 * 1.75^5 = 13.130469, got {}",
got[0].1
);
}
#[test]
fn a_single_token_sequence_breaker_is_never_itself_penalised() {
let with_breaker = dry(1.0, 1.1, 2, 6, &[vec![3]]);
assert!(penalties_of(&with_breaker, &[0, 1, 3, 4, 0, 1]).is_empty());
let without = dry(1.0, 1.1, 2, 6, &[]);
assert_eq!(penalties_of(&without, &[0, 1, 3, 4, 0, 1]), vec![(3, 1.0)]);
}
#[test]
fn a_zero_multiplier_a_base_below_one_or_a_zero_window_all_disable_it() {
let history = [0usize, 1, 2, 0, 1];
for params in [
DryParams::off(),
dry(0.0, 1.75, 2, -1, &[]),
dry(1.0, 0.5, 2, -1, &[]),
dry(1.0, 1.75, 2, 0, &[]),
] {
assert!(!params.is_enabled(), "{params:?} must be disabled");
assert!(
params
.penalties(PenaltyWindow::new(&[], &history))
.is_empty(),
"{params:?} penalised something while disabled"
);
}
assert!(dry(1.0, 1.75, 2, -1, &[]).is_enabled());
}
#[test]
fn the_scan_only_sees_the_last_n_tokens() {
let history = [0usize, 1, 2, 0, 1];
assert_eq!(
penalties_of(&dry(1.0, 1.1, 2, 5, &[]), &history),
vec![(2, 1.0)]
);
assert!(penalties_of(&dry(1.0, 1.1, 2, 3, &[]), &history).is_empty());
assert_eq!(
penalties_of(&dry(1.0, 1.1, 2, -1, &[]), &history),
vec![(2, 1.0)]
);
}
#[test]
fn the_scan_reads_across_the_prompt_and_generation_seam() {
let params = dry(1.0, 1.1, 2, 5, &[]);
let whole = params.penalties(PenaltyWindow::new(&[], &[0, 1, 2, 0, 1]));
let split = params.penalties(PenaltyWindow::new(&[0, 1, 2], &[0, 1]));
let all_prompt = params.penalties(PenaltyWindow::new(&[0, 1, 2, 0, 1], &[]));
assert_eq!(whole, split);
assert_eq!(whole, all_prompt);
assert_eq!(whole.len(), 1);
}
#[test]
fn a_very_long_repetition_clamps_the_exponent_rather_than_overflowing() {
let params = dry(1.0, 1.75, 2, -1, &[]);
let mut history: Vec<usize> = vec![0usize; 400];
history.push(1);
history.extend(std::iter::repeat_n(0usize, 400));
let penalties = params.penalties(PenaltyWindow::new(&[], &history));
for (token, penalty) in penalties {
assert!(
penalty.is_finite(),
"token {token} got a non-finite penalty {penalty}"
);
}
}
struct WordVocab {
words: Vec<&'static str>,
}
impl DryVocab for WordVocab {
fn n_tokens(&self) -> usize {
self.words.len()
}
fn detokenize(&self, token: usize) -> String {
self.words[token].to_string()
}
fn tokenize(&self, text: &str) -> Vec<usize> {
let mut out = Vec::new();
let mut rest = text;
'outer: while !rest.is_empty() {
let mut candidates: Vec<usize> = (0..self.words.len()).collect();
candidates.sort_by_key(|&i| std::cmp::Reverse(self.words[i].len()));
for i in candidates {
let w = self.words[i];
if !w.is_empty() && rest.starts_with(w) {
out.push(i);
rest = &rest[w.len()..];
continue 'outer;
}
}
break;
}
out
}
}
#[test]
fn a_breaker_is_found_inside_a_token_and_across_a_token_boundary() {
let vocab = WordVocab {
words: vec!["foo", "bar", ":", "ab", "c", "abc", "xab"],
};
let colon = DryBreakers::from_vocab(&vocab, &[":".to_string()]);
assert_eq!(colon.tails(2), Some([Vec::new()].as_slice()));
assert_eq!(colon.tails(0), None, "`foo` does not contain a colon");
assert!(colon.is_single_token_breaker(2));
let abc = DryBreakers::from_vocab(&vocab, &["abc".to_string()]);
assert_eq!(abc.tails(5), Some([Vec::new()].as_slice()));
assert_eq!(abc.tails(3), Some([vec![4usize]].as_slice()));
assert_eq!(abc.tails(6), Some([vec![4usize]].as_slice()));
assert!(!abc.is_single_token_breaker(3), "`ab` needs `c` to follow");
assert!(abc.is_single_token_breaker(5));
assert_eq!(abc.tails(0), None);
}
#[test]
fn a_breaker_with_a_tail_only_stops_the_scan_when_the_tail_follows() {
let vocab = WordVocab {
words: vec!["foo", "bar", ":", "ab", "c", "abc", "xab"],
};
let breakers = DryBreakers::from_vocab(&vocab, &["abc".to_string()]);
assert_eq!(breakers.tails(3), Some([vec![4usize]].as_slice()));
let params = DryParams::new(1.0, 1.1, 2, -1, 1024, breakers);
let broken = [0usize, 1, 2, 3, 4, 0, 1, 2, 3, 4];
let unbroken = [0usize, 1, 2, 3, 6, 0, 1, 2, 3, 6];
assert!(
params
.penalties(PenaltyWindow::new(&[], &broken))
.is_empty(),
"`ab` followed by `c` is the breaker, one token from the end"
);
assert!(
!params
.penalties(PenaltyWindow::new(&[], &unbroken))
.is_empty(),
"`ab` followed by `xab` is not the breaker, so the scan runs"
);
let no_breakers = DryParams::new(1.0, 1.1, 2, -1, 1024, DryBreakers::none());
assert!(
!no_breakers
.penalties(PenaltyWindow::new(&[], &broken))
.is_empty(),
"the same window without breakers repeats, so the first \
assertion is about the breaker and not about the window"
);
}
#[test]
fn an_empty_breaker_string_is_skipped() {
let vocab = WordVocab {
words: vec!["foo", "bar"],
};
let breakers = DryBreakers::from_vocab(&vocab, &[String::new()]);
assert!(breakers.is_empty());
assert_eq!(breakers.raw(), &[String::new()]);
}
#[test]
fn the_default_breakers_are_llama_cpps_four() {
assert_eq!(DEFAULT_SEQUENCE_BREAKERS, ["\n", ":", "\"", "*"]);
}
}