Make tokenizer byte offsets total
This commit is contained in:
@@ -563,20 +563,33 @@ fn gpt2_codepoint_to_byte(character: char) -> Option<u8> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn next_char(text: &str, position: usize) -> usize {
|
fn next_char(text: &str, position: usize) -> usize {
|
||||||
if position >= text.len() {
|
match text.get(position..).and_then(|rest| rest.chars().next()) {
|
||||||
return text.len();
|
Some(ch) => position + ch.len_utf8(),
|
||||||
|
None => position.saturating_add(1).min(text.len()),
|
||||||
}
|
}
|
||||||
position + text[position..].chars().next().unwrap().len_utf8()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn cjk_at(text: &str, position: usize) -> bool {
|
fn cjk_at(text: &str, position: usize) -> bool {
|
||||||
let character = text[position..].chars().next().unwrap() as u32;
|
let Some(&byte) = text.as_bytes().get(position) else {
|
||||||
|
return false;
|
||||||
|
};
|
||||||
|
if byte < 128 {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
let Some(character) = text
|
||||||
|
.get(position..)
|
||||||
|
.and_then(|rest| rest.chars().next())
|
||||||
|
.map(|ch| ch as u32)
|
||||||
|
else {
|
||||||
|
return false;
|
||||||
|
};
|
||||||
(0x4e00..=0x9fa5).contains(&character)
|
(0x4e00..=0x9fa5).contains(&character)
|
||||||
|| (0x3040..=0x309f).contains(&character)
|
|| (0x3040..=0x309f).contains(&character)
|
||||||
|| (0x30a0..=0x30ff).contains(&character)
|
|| (0x30a0..=0x30ff).contains(&character)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn letter_like(text: &str, position: usize) -> bool {
|
fn letter_like(text: &str, position: usize) -> bool {
|
||||||
|
// Non-ASCII bytes count as letters; CJK and kana are isolated first by cjk_at.
|
||||||
let byte = text.as_bytes()[position];
|
let byte = text.as_bytes()[position];
|
||||||
byte >= 128 || byte.is_ascii_alphabetic()
|
byte >= 128 || byte.is_ascii_alphabetic()
|
||||||
}
|
}
|
||||||
@@ -607,17 +620,16 @@ struct CharInfo {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn char_info(text: &str, position: usize) -> CharInfo {
|
fn char_info(text: &str, position: usize) -> CharInfo {
|
||||||
if position >= text.len() {
|
let Some(ch) = text.get(position..).and_then(|rest| rest.chars().next()) else {
|
||||||
return CharInfo {
|
return CharInfo {
|
||||||
ch: '\0',
|
ch: '\0',
|
||||||
next: text.len(),
|
next: position.saturating_add(1).min(text.len()),
|
||||||
letter: false,
|
letter: false,
|
||||||
number: false,
|
number: false,
|
||||||
whitespace: false,
|
whitespace: false,
|
||||||
punctuation: false,
|
punctuation: false,
|
||||||
};
|
};
|
||||||
}
|
};
|
||||||
let ch = text[position..].chars().next().unwrap();
|
|
||||||
let codepoint = ch as u32;
|
let codepoint = ch as u32;
|
||||||
let whitespace = unicode_whitespace(codepoint);
|
let whitespace = unicode_whitespace(codepoint);
|
||||||
let number = unicode_number(codepoint);
|
let number = unicode_number(codepoint);
|
||||||
@@ -729,3 +741,28 @@ fn unicode_punctuation(cp: u32) -> bool {
|
|||||||
.iter()
|
.iter()
|
||||||
.any(|range| range.contains(&cp))
|
.any(|range| range.contains(&cp))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn character_helpers_are_total_at_every_byte_offset() {
|
||||||
|
let text = "Aé中😊";
|
||||||
|
for position in 0..=text.len() + 1 {
|
||||||
|
let next = next_char(text, position);
|
||||||
|
assert!(next <= text.len());
|
||||||
|
if position < text.len() {
|
||||||
|
assert!(next > position);
|
||||||
|
}
|
||||||
|
assert!(char_info(text, position).next <= text.len());
|
||||||
|
let _ = cjk_at(text, position);
|
||||||
|
}
|
||||||
|
|
||||||
|
assert_eq!(next_char(text, 1), 3);
|
||||||
|
assert!(cjk_at(text, 3));
|
||||||
|
assert!(!cjk_at(text, 6));
|
||||||
|
assert_eq!(next_char("", 0), 0);
|
||||||
|
assert!(!cjk_at("", 0));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user