Spaces:
Sleeping
Sleeping
| from dataclasses import dataclass | |
| from heapq import heappop, heappush, nlargest | |
| ROOT_NODE_OFFSET = 2 | |
| 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)] | |