simdref 0.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
simdref/search.py ADDED
@@ -0,0 +1,288 @@
1
+ """Fuzzy search and ranking for intrinsics and instructions.
2
+
3
+ Scores candidates through exact/prefix/substring matching, normalised token
4
+ overlap, rapidfuzz similarity, SIMD width bonuses, and intent-based biasing
5
+ (intrinsic vs instruction preference). See ARCHITECTURE.md for details.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from dataclasses import dataclass
11
+ import re
12
+
13
+ try:
14
+ from rapidfuzz import fuzz
15
+ except Exception: # pragma: no cover - fallback when dependency is unavailable
16
+ from difflib import SequenceMatcher
17
+
18
+ class _FallbackFuzz:
19
+ @staticmethod
20
+ def ratio(a: str, b: str) -> float:
21
+ return SequenceMatcher(a=a, b=b).ratio() * 100.0
22
+
23
+ @staticmethod
24
+ def partial_ratio(a: str, b: str) -> float:
25
+ return SequenceMatcher(a=a, b=b).ratio() * 100.0
26
+
27
+ @staticmethod
28
+ def token_set_ratio(a: str, b: str) -> float:
29
+ return SequenceMatcher(a=" ".join(sorted(set(a.split()))), b=" ".join(sorted(set(b.split())))).ratio() * 100.0
30
+
31
+ fuzz = _FallbackFuzz()
32
+
33
+ from simdref.models import Catalog
34
+
35
+
36
+ @dataclass(slots=True)
37
+ class SearchResult:
38
+ kind: str
39
+ key: str
40
+ title: str
41
+ subtitle: str
42
+ score: float
43
+
44
+
45
+ TOKEN_RE = re.compile(r"[A-Za-z0-9]+")
46
+ WIDTH_TOKEN_RE = re.compile(r"^(mm|mm\d+|xmm|ymm|zmm)$")
47
+
48
+
49
+ def _normalize_tokens(value: str) -> list[str]:
50
+ text = value.replace("_", " ").replace(",", " ").replace("{", " ").replace("}", " ")
51
+ return [token.casefold() for token in TOKEN_RE.findall(text)]
52
+
53
+
54
+ def _normalize_text(value: str) -> str:
55
+ return " ".join(_normalize_tokens(value))
56
+
57
+
58
+ def _normalized_instruction_query(value: str) -> str:
59
+ return _normalize_text(value)
60
+
61
+
62
+ def _classify_query(query: str) -> str:
63
+ lowered = query.casefold().strip()
64
+ normalized = _normalize_text(query)
65
+ if lowered.startswith("__riscv_"):
66
+ return "intrinsic"
67
+ if "." in lowered and not lowered.startswith("_"):
68
+ return "instruction"
69
+ if lowered == "v" or lowered.startswith("zv") or lowered.startswith("zve") or lowered.startswith("rv"):
70
+ return "instruction"
71
+ if lowered.startswith("_mm") or lowered.startswith("__m") or normalized.startswith("mm ") or normalized == "mm":
72
+ return "intrinsic"
73
+ if "_" in lowered and "mm" in lowered:
74
+ return "intrinsic"
75
+ if normalized:
76
+ tokens = normalized.split()
77
+ if tokens and all(token.isalpha() or token.isalnum() for token in tokens):
78
+ first = tokens[0]
79
+ if first in {"add", "sub", "mul", "div", "mov", "cmp", "and", "or", "xor"} or first.startswith("v"):
80
+ return "instruction"
81
+ return "neutral"
82
+
83
+
84
+ def _token_prefix_score(query_tokens: list[str], candidate_tokens: list[str]) -> float:
85
+ if not query_tokens or not candidate_tokens:
86
+ return 0.0
87
+ matched = 0
88
+ for q in query_tokens:
89
+ if any(token.startswith(q) for token in candidate_tokens):
90
+ matched += 1
91
+ return 100.0 * matched / len(query_tokens)
92
+
93
+
94
+ def _token_overlap_count(query_tokens: list[str], candidate_tokens: list[str]) -> int:
95
+ if not query_tokens or not candidate_tokens:
96
+ return 0
97
+ count = 0
98
+ for q in query_tokens:
99
+ if any(token == q or token.startswith(q) or q.startswith(token) for token in candidate_tokens):
100
+ count += 1
101
+ return count
102
+
103
+
104
+ def _width_family_bonus(query: str, candidate: str) -> float:
105
+ query_tokens = _normalize_tokens(query)
106
+ candidate_tokens = _normalize_tokens(candidate)
107
+ query_widths = {token for token in query_tokens if WIDTH_TOKEN_RE.match(token)}
108
+ candidate_widths = {token for token in candidate_tokens if WIDTH_TOKEN_RE.match(token)}
109
+ if not query_widths or not candidate_widths:
110
+ return 0.0
111
+ if query_widths & candidate_widths:
112
+ return 22.0
113
+ return -22.0
114
+
115
+
116
+ def _meaningful_query_tokens(tokens: list[str]) -> list[str]:
117
+ return [token for token in tokens if not WIDTH_TOKEN_RE.match(token)]
118
+
119
+
120
+ def _has_structural_overlap(query: str, candidate: str) -> bool:
121
+ q = query.casefold().strip()
122
+ c = candidate.casefold().strip()
123
+ if not q or not c:
124
+ return False
125
+ if q == c or c.startswith(q) or q in c:
126
+ return True
127
+ query_tokens = _normalize_tokens(query)
128
+ candidate_tokens = _normalize_tokens(candidate)
129
+ meaningful_query_tokens = _meaningful_query_tokens(query_tokens)
130
+ if meaningful_query_tokens:
131
+ return _token_overlap_count(meaningful_query_tokens, candidate_tokens) > 0
132
+ return _token_overlap_count(query_tokens, candidate_tokens) > 0
133
+
134
+
135
+ def _base_score(query: str, candidate: str) -> float:
136
+ q = query.casefold().strip()
137
+ c = candidate.casefold().strip()
138
+ if not q or not c:
139
+ return 0.0
140
+ if q == c:
141
+ return 220.0
142
+ if c.startswith(q):
143
+ return 175.0
144
+ if q in c:
145
+ return 135.0
146
+ qnorm = _normalize_text(query)
147
+ cnorm = _normalize_text(candidate)
148
+ if qnorm and qnorm == cnorm:
149
+ return 190.0
150
+ if qnorm and cnorm.startswith(qnorm):
151
+ return 165.0
152
+ query_tokens = qnorm.split()
153
+ candidate_tokens = cnorm.split()
154
+ token_prefix = _token_prefix_score(query_tokens, candidate_tokens)
155
+ token_overlap = _token_overlap_count(query_tokens, candidate_tokens)
156
+ if token_overlap == 0:
157
+ return 0.0
158
+ prefix_score = token_prefix + 40.0
159
+ if prefix_score >= 155.0:
160
+ return prefix_score
161
+ token_set = fuzz.token_set_ratio(qnorm, cnorm) if qnorm and cnorm else 0.0
162
+ partial = fuzz.partial_ratio(qnorm, cnorm) if qnorm and cnorm else 0.0
163
+ ratio = fuzz.ratio(qnorm, cnorm) if qnorm and cnorm else 0.0
164
+ return max(prefix_score, token_set + 20.0, partial + 10.0, ratio)
165
+
166
+
167
+ def _intent_bias(query_kind: str, result_kind: str) -> float:
168
+ if query_kind == "intrinsic":
169
+ return 45.0 if result_kind == "intrinsic" else -25.0
170
+ if query_kind == "instruction":
171
+ return 35.0 if result_kind == "instruction" else -10.0
172
+ return 0.0
173
+
174
+
175
+ def _isa_match_bias(query: str, isa_values: list[str], result_kind: str) -> float:
176
+ normalized_query = _normalize_text(query)
177
+ if not normalized_query:
178
+ return 0.0
179
+ isa_tokens = {_normalize_text(value) for value in isa_values if _normalize_text(value)}
180
+ if normalized_query not in isa_tokens:
181
+ return 0.0
182
+ return 30.0 if result_kind == "instruction" else -15.0
183
+
184
+
185
+ def _is_pure_isa_query(query: str, isa_values: list[str]) -> bool:
186
+ normalized_query = _normalize_text(query)
187
+ if not normalized_query:
188
+ return False
189
+ return normalized_query in {_normalize_text(value) for value in isa_values if _normalize_text(value)}
190
+
191
+
192
+ def _looks_like_isa_query(query: str) -> bool:
193
+ lowered = query.casefold().strip()
194
+ return lowered == "v" or lowered.startswith("zv") or lowered.startswith("zve") or lowered.startswith("rv")
195
+
196
+
197
+ def search_records(intrinsics: list, instructions: list, query: str, limit: int = 20) -> list[SearchResult]:
198
+ results: list[SearchResult] = []
199
+ query_kind = _classify_query(query)
200
+ for item in intrinsics:
201
+ if query_kind == "instruction" and _is_pure_isa_query(query, item.isa) and not _has_structural_overlap(query, item.name):
202
+ continue
203
+ if query_kind == "intrinsic" and not _has_structural_overlap(query, item.name):
204
+ continue
205
+ score = max(_base_score(query, item.name), _base_score(query, item.search_blob)) + _intent_bias(query_kind, "intrinsic")
206
+ score += _isa_match_bias(query, item.isa, "intrinsic")
207
+ score += _width_family_bonus(query, item.name)
208
+ if score >= 35:
209
+ results.append(
210
+ SearchResult(
211
+ kind="intrinsic",
212
+ key=item.name,
213
+ title=item.name,
214
+ subtitle=item.description,
215
+ score=score,
216
+ )
217
+ )
218
+ for item in instructions:
219
+ if _looks_like_isa_query(query) and not _is_pure_isa_query(query, item.isa):
220
+ continue
221
+ if query_kind == "instruction" and not _looks_like_isa_query(query) and not _has_structural_overlap(query, item.key):
222
+ continue
223
+ score = max(_base_score(query, item.mnemonic), _base_score(query, item.key), _base_score(query, item.search_blob))
224
+ score += _intent_bias(query_kind, "instruction")
225
+ score += _isa_match_bias(query, item.isa, "instruction")
226
+ if score >= 35:
227
+ results.append(
228
+ SearchResult(
229
+ kind="instruction",
230
+ key=item.db_key,
231
+ title=item.key,
232
+ subtitle=item.summary,
233
+ score=score,
234
+ )
235
+ )
236
+
237
+ def sort_key(item: SearchResult):
238
+ preferred = 0
239
+ if query_kind == "intrinsic":
240
+ preferred = 0 if item.kind == "intrinsic" else 1
241
+ elif query_kind == "instruction":
242
+ preferred = 0 if item.kind == "instruction" else 1
243
+ return (-item.score, preferred, len(item.title), item.title)
244
+
245
+ results.sort(key=sort_key)
246
+ return results[:limit]
247
+
248
+
249
+ def search_catalog(catalog: Catalog, query: str, limit: int = 20) -> list[SearchResult]:
250
+ return search_records(catalog.intrinsics, catalog.instructions, query, limit=limit)
251
+
252
+
253
+ def find_intrinsic(catalog: Catalog, name: str):
254
+ target = name.casefold()
255
+ for item in catalog.intrinsics:
256
+ if item.name.casefold() == target:
257
+ return item
258
+ return None
259
+
260
+
261
+ def find_instruction(catalog: Catalog, query: str):
262
+ target = query.casefold()
263
+ normalized_target = _normalized_instruction_query(query)
264
+ for item in catalog.instructions:
265
+ if item.key.casefold() == target or item.mnemonic.casefold() == target:
266
+ return item
267
+ if normalized_target and (
268
+ _normalized_instruction_query(item.key) == normalized_target
269
+ or _normalized_instruction_query(item.mnemonic) == normalized_target
270
+ ):
271
+ return item
272
+ return None
273
+
274
+
275
+ def find_instructions(catalog: Catalog, query: str) -> list:
276
+ target = query.casefold()
277
+ normalized_target = _normalized_instruction_query(query)
278
+ matches = []
279
+ for item in catalog.instructions:
280
+ if item.key.casefold() == target or item.mnemonic.casefold() == target:
281
+ matches.append(item)
282
+ continue
283
+ if normalized_target and (
284
+ _normalized_instruction_query(item.key) == normalized_target
285
+ or _normalized_instruction_query(item.mnemonic) == normalized_target
286
+ ):
287
+ matches.append(item)
288
+ return matches