#[derive(Default)]
pub(crate) struct RepetitionDetector {
last_char: Option<char>,
run_length: usize,
alt_pair: Option<(char, char)>,
alt_count: usize,
alt_buf: [char; 2],
alt_buf_len: usize,
alt_run_chars: usize,
}
const MAX_CHAR_RUN: usize = 80;
const MAX_ALT_CYCLES: usize = 30;
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum RepetitionKind {
CharRun { ch: char, count: usize },
AlternatingPattern { pattern: String, cycles: usize },
}
#[derive(Debug, Clone)]
pub(crate) struct RepetitionWarning {
pub kind: RepetitionKind,
pub message: String,
}
impl RepetitionDetector {
pub fn feed_text(&mut self, text: &str) -> Option<RepetitionWarning> {
for ch in text.chars() {
self.feed_char(ch);
}
self.check_warnings()
}
pub fn feed_thinking(&mut self, text: &str) -> Option<RepetitionWarning> {
for ch in text.chars() {
self.feed_char(ch);
}
self.check_warnings()
}
fn feed_char(&mut self, ch: char) {
if Some(ch) == self.last_char {
self.run_length += 1;
} else {
self.last_char = Some(ch);
self.run_length = 1;
}
if self.alt_buf_len < 2 {
self.alt_buf[self.alt_buf_len] = ch;
self.alt_buf_len += 1;
self.alt_run_chars += 1;
return;
}
let prev = self.alt_buf[1];
self.alt_buf[0] = prev;
self.alt_buf[1] = ch;
self.alt_run_chars += 1;
let pair = (self.alt_buf[0], self.alt_buf[1]);
if pair.0 == pair.1 {
self.alt_pair = None;
self.alt_count = 0;
return;
}
if Some(pair) == self.alt_pair {
self.alt_count += 1;
} else if let Some(ref current) = self.alt_pair {
if pair.0 == current.1 && pair.1 == current.0 {
} else {
self.alt_pair = Some(pair);
self.alt_count = 1;
self.alt_run_chars = 2;
}
} else {
self.alt_pair = Some(pair);
self.alt_count = 1;
}
}
fn check_warnings(&mut self) -> Option<RepetitionWarning> {
if self.run_length >= MAX_CHAR_RUN {
let ch = self.last_char.unwrap_or('?');
return Some(RepetitionWarning {
kind: RepetitionKind::CharRun {
ch,
count: self.run_length,
},
message: format!(
"Repetitive character \"{ch}\" detected ({}+ times in a row). Model output may be degenerate.",
self.run_length
),
});
}
if self.alt_count >= MAX_ALT_CYCLES {
if let Some((a, b)) = self.alt_pair {
return Some(RepetitionWarning {
kind: RepetitionKind::AlternatingPattern {
pattern: format!("{a}{b}"),
cycles: self.alt_count,
},
message: format!(
"Alternating pattern \"{a}{b}\" repeated {}+ times. Model output may be degenerate.",
self.alt_count
),
});
}
}
None
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn detects_character_run() {
let mut d = RepetitionDetector::default();
let mut warning = None;
for _ in 0..79 {
warning = d.feed_text("a");
}
assert!(warning.is_none());
warning = d.feed_text("a");
assert!(warning.is_some());
assert!(matches!(
warning.unwrap().kind,
RepetitionKind::CharRun { .. }
));
}
#[test]
fn resets_char_run_on_different_char() {
let mut d = RepetitionDetector::default();
for _ in 0..50 {
assert!(d.feed_text("a").is_none());
}
assert!(d.run_length == 50);
assert!(d.feed_text("b").is_none());
assert!(d.run_length == 1);
assert!(d.last_char == Some('b'));
}
#[test]
fn detects_alternating_pattern() {
let mut d = RepetitionDetector::default();
for _ in 0..30 {
d.feed_text("-");
d.feed_text("_");
}
let w = d.feed_text("-");
assert!(w.is_some());
assert!(matches!(
w.unwrap().kind,
RepetitionKind::AlternatingPattern { .. }
));
}
#[test]
fn feeds_multibyte_text() {
let mut d = RepetitionDetector::default();
let s = "ç".repeat(80);
let w = d.feed_text(&s);
assert!(w.is_some());
assert!(matches!(
w.unwrap().kind,
RepetitionKind::CharRun { ch: 'ç', .. }
));
}
}