shethjenil's picture
Upload 5 files
a313012 verified
Raw
History Blame Contribute Delete
3.29 kB
from dataclasses import dataclass
from heapq import heappop, heappush, nlargest
ROOT_NODE_OFFSET = 2
@dataclass(slots=True)
class Candidate:
text: str = ""
is_complete: bool = False
score: int = 0
address: int = 0
class OnDeviceHeadModel:
def __init__(self, path):
self.fp = open(path, "rb")
self.address_size = self._read_int(1)
self.score_size = self._read_int(1)
if self.address_size not in (3, 4):
raise ValueError(f"invalid address_size={self.address_size}")
if self.score_size not in (2, 3, 4):
raise ValueError(f"invalid score_size={self.score_size}")
def _read_int(self, size):
b = self.fp.read(size)
if len(b) != size:
raise EOFError
return int.from_bytes(b, "little")
def _node_info(self, address):
self.fp.seek(address)
block = self._read_int(self.score_size)
leaf = None
if block & 1:
leaf = Candidate(is_complete=True, score=self._read_int(self.score_size))
return block >> 1, leaf
def _child(self, prefix):
b = self.fp.read(1)
if not b:
return None
size = b[0]
if size == 0:
return None
text = prefix + self.fp.read(size).decode("utf-8", errors="ignore")
first = self.fp.read(1)
if not first:
return None
first = first[0]
if first & 1 == 0:
score = int.from_bytes(bytes([first]) + self.fp.read(self.score_size - 1), "little") >> 1
return Candidate(text=text, is_complete=True, score=score)
address = int.from_bytes(bytes([first]) + self.fp.read(self.address_size - 1), "little") >> 1
pos = self.fp.tell()
score, _ = self._node_info(address)
self.fp.seek(pos)
return Candidate(text=text, score=score, address=address)
def _children(self, node):
if node.is_complete:
return
_, leaf = self._node_info(node.address)
if leaf:
leaf.text = node.text
yield leaf
while child := self._child(node.text):
yield child
def _find_start(self, prefix):
node = Candidate(address=ROOT_NODE_OFFSET)
while len(node.text) < len(prefix):
for child in self._children(node):
if child.text.startswith(prefix) or prefix.startswith(child.text):
if child.is_complete and len(child.text) < len(prefix):
continue
node = child
break
else:
return None
return node
def suggest(self, prefix, limit=10):
prefix = prefix.strip().lower()
if not prefix:
return []
start = self._find_start(prefix)
if not start:
return []
heap = [(-start.score, start)]
leaves = []
while heap:
_, node = heappop(heap)
for child in self._children(node):
if child.is_complete:
leaves.append(child)
else:
heappush(heap, (-child.score, child))
return [(c.text, c.score) for c in nlargest(limit, leaves, key=lambda c: c.score)]