Skip to content

Commit e53862d

Browse files
committed
feat: better interface
1 parent 6a9db71 commit e53862d

5 files changed

Lines changed: 85 additions & 14 deletions

File tree

python/mtc_token_healing/mtc_token_healing.pyi

Lines changed: 19 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
from typing import Optional, Sequence, Tuple
1+
from typing import Optional, Sequence, Tuple, overload
22

33
TokenId = int
44
SortedTokenId = int
@@ -19,15 +19,30 @@ class VocabPrefixAutomaton:
1919
def parse_tokens(
2020
self, token_ids: Sequence[TokenId]
2121
) -> Sequence[Tuple[bytes, SortedTokenRange]]: ...
22+
def parse_tokens_str_suffix(
23+
self, token_ids: Sequence[TokenId]
24+
) -> Sequence[Tuple[str, SortedTokenRange]]: ...
25+
@overload
26+
def get_original_token_ids(self, sorted_token_id: SortedTokenId) -> TokenId: ...
27+
@overload
28+
def get_original_token_ids(
29+
self, sorted_token_ids: Sequence[SortedTokenId]
30+
) -> Sequence[TokenId]: ...
31+
@overload
32+
def get_sorted_token_ids(self, token_id: TokenId) -> SortedTokenId: ...
33+
@overload
34+
def get_sorted_token_ids(
35+
self, token_ids: Sequence[TokenId]
36+
) -> Sequence[SortedTokenId]: ...
2237

2338
class TokenSeqTrieNode:
24-
token: TokenId
39+
token: int
2540
pred_range: Optional[SortedTokenRange]
2641
parent: int
2742
subtree_lower: int
2843
subtree_upper: int
2944

3045
def dfs_token_seq_trie(
31-
token_ids_seq: Sequence[Sequence[TokenId]],
46+
token_ids_seq: Sequence[Sequence[int]],
3247
pred_rank_ranges: Sequence[SortedTokenRange],
33-
) -> Sequence[TokenSeqTrieNode]: ...
48+
) -> Tuple[Sequence[TokenSeqTrieNode], int]: ...

python/src/lib.rs

Lines changed: 17 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,16 +1,16 @@
11
use std::borrow::Cow;
22

33
use ::mtc_token_healing::{
4-
SortedTokenRange, TokenId, TokenSeqInput, TokenSeqTrieNode, VocabPrefixAutomaton,
5-
dfs_token_seq_trie,
4+
dfs_token_seq_trie, SortedTokenRange, TokenId, TokenSeqInput, TokenSeqTrieNode,
5+
VocabPrefixAutomaton,
66
};
77
use pyo3::prelude::*;
88

99
#[pyfunction(name = "dfs_token_seq_trie")]
1010
fn dfs_token_seq_trie_py(
1111
token_ids: Vec<Vec<TokenId>>,
1212
pred_rank_ranges: Vec<SortedTokenRange>,
13-
) -> Vec<TokenSeqTrieNode> {
13+
) -> (Vec<TokenSeqTrieNode>, usize) {
1414
let inputs = token_ids
1515
.into_iter()
1616
.zip(pred_rank_ranges)
@@ -19,7 +19,20 @@ fn dfs_token_seq_trie_py(
1919
pred_range: r,
2020
})
2121
.collect();
22-
dfs_token_seq_trie(inputs)
22+
let nodes = dfs_token_seq_trie(inputs);
23+
let parent_chain_len = {
24+
let mut res = 0;
25+
while res < nodes.len() {
26+
let node = &nodes[res];
27+
if node.parent == res.saturating_sub(1) && node.pred_range.is_none() {
28+
res += 1;
29+
continue;
30+
}
31+
break;
32+
}
33+
res
34+
};
35+
(nodes, parent_chain_len)
2336
}
2437

2538
#[pymodule]

python/tests/test_dfs_token_seq_trie.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -34,7 +34,7 @@ def test_dfs_token_seq_trie():
3434

3535
for i in range(len(nodes)):
3636
masks = [
37-
j < i and nodes[j].subtree_upper >= nodes[i].subtree_upper
37+
j <= i and nodes[j].subtree_upper >= nodes[i].subtree_upper
3838
for j in range(len(nodes))
3939
]
4040
print("".join(map(str, map(int, masks))))

src/automaton.rs

Lines changed: 47 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -120,12 +120,32 @@ impl VocabPrefixAutomaton {
120120

121121
#[cfg(feature = "pyo3")]
122122
pub mod pyo3 {
123-
use pyo3::{Bound, Python, pymethods, types::PyBytes};
123+
use pyo3::{Bound, FromPyObject, IntoPyObject, Python, pymethods, types::PyBytes};
124124

125125
use crate::{SortedTokenId, SortedTokenRange, TokenId};
126126

127127
use super::VocabPrefixAutomaton;
128128

129+
#[derive(Debug, FromPyObject, IntoPyObject)]
130+
enum TokenIdSeq {
131+
TokenId(TokenId),
132+
Seq(Vec<TokenId>),
133+
}
134+
135+
impl TokenIdSeq {
136+
fn map<F: FnMut(TokenId) -> TokenId>(self, mut f: F) -> Self {
137+
match self {
138+
Self::TokenId(id) => Self::TokenId(f(id)),
139+
Self::Seq(mut items) => {
140+
items.iter_mut().for_each(|item| {
141+
*item = f(*item);
142+
});
143+
Self::Seq(items)
144+
}
145+
}
146+
}
147+
}
148+
129149
#[pymethods]
130150
impl VocabPrefixAutomaton {
131151
#[new]
@@ -151,10 +171,11 @@ pub mod pyo3 {
151171
#[pyo3(name = "parse_bytes")]
152172
fn parse_bytes_py(
153173
&self,
174+
py: Python<'_>,
154175
bytes: &[u8],
155176
start_from: usize,
156177
) -> Vec<(usize, SortedTokenRange)> {
157-
self.parse_bytes(bytes, start_from)
178+
py.allow_threads(|| self.parse_bytes(bytes, start_from))
158179
}
159180

160181
#[pyo3(name = "parse_tokens")]
@@ -163,10 +184,32 @@ pub mod pyo3 {
163184
py: Python<'py>,
164185
tokens: Vec<usize>,
165186
) -> Vec<(Bound<'py, PyBytes>, SortedTokenRange)> {
166-
self.parse_rev_token_id_seq(tokens.into_iter().rev())
167-
.into_iter()
187+
let res = py.allow_threads(|| self.parse_rev_token_id_seq(tokens.into_iter().rev()));
188+
res.into_iter()
168189
.map(|(b, c)| (PyBytes::new(py, &b), c))
169190
.collect()
170191
}
192+
193+
#[pyo3(name = "parse_tokens_str_suffix")]
194+
fn parse_tokens_str_suffix_py(
195+
&self,
196+
py: Python<'_>,
197+
tokens: Vec<usize>,
198+
) -> Vec<(String, SortedTokenRange)> {
199+
py.allow_threads(|| {
200+
self.parse_rev_token_id_seq(tokens.into_iter().rev())
201+
.into_iter()
202+
.filter_map(|(b, c)| String::from_utf8(b.into()).ok().map(|s| (s, c)))
203+
.collect()
204+
})
205+
}
206+
207+
fn get_original_token_ids(&self, py: Python<'_>, seq: TokenIdSeq) -> TokenIdSeq {
208+
py.allow_threads(|| seq.map(|id| self.order.get(id as usize).copied().unwrap_or(id)))
209+
}
210+
211+
fn get_sorted_token_ids(&self, py: Python<'_>, seq: TokenIdSeq) -> TokenIdSeq {
212+
py.allow_threads(|| seq.map(|id| self.rank.get(id as usize).copied().unwrap_or(id)))
213+
}
171214
}
172215
}

src/token.rs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,7 @@ mod _pyo3 {
2929
impl SortedTokenRange {
3030
pub(crate) fn repr_py(&self) -> String {
3131
let Self { lower, upper } = self;
32-
format!("SortedTokenRange(lower={}, upper={})", lower, upper)
32+
format!("SortedTokenRange(lower={lower}, upper={upper})")
3333
}
3434
}
3535

0 commit comments

Comments
 (0)