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/__init__.py ADDED
@@ -0,0 +1,6 @@
1
+ """simdref package."""
2
+
3
+ __all__ = ["__version__"]
4
+
5
+ __version__ = "0.0.0"
6
+
simdref/__main__.py ADDED
@@ -0,0 +1,6 @@
1
+ from simdref.cli import main
2
+
3
+
4
+ if __name__ == "__main__":
5
+ raise SystemExit(main())
6
+
simdref/annotate.py ADDED
@@ -0,0 +1,448 @@
1
+ """Annotate assembly (``.s``) files with per-instruction summaries and perf.
2
+
3
+ Given GAS/AT&T-syntax assembly, emit an annotated ``.sa`` file where each
4
+ recognised instruction line carries a trailing ``# ...`` comment describing
5
+ what the instruction does plus latency / CPI figures pulled from the
6
+ simdref catalog.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import json
12
+ import re
13
+ import sqlite3
14
+ import statistics
15
+ from dataclasses import dataclass, field
16
+ from enum import Enum
17
+ from typing import Any, Iterable, Iterator
18
+
19
+ from simdref import perf
20
+ from simdref.models import InstructionRecord
21
+ from simdref.storage import load_instructions_by_mnemonic_from_db
22
+
23
+
24
+ class LineKind(str, Enum):
25
+ BLANK = "blank"
26
+ LABEL = "label"
27
+ DIRECTIVE = "directive"
28
+ COMMENT = "comment"
29
+ INSTRUCTION = "instruction"
30
+
31
+
32
+ @dataclass(slots=True)
33
+ class AsmLine:
34
+ kind: LineKind
35
+ raw: str
36
+ indent: str = ""
37
+ mnemonic: str = ""
38
+ operands: str = ""
39
+ trailing_comment: str = ""
40
+
41
+
42
+ _INSTR_RE = re.compile(
43
+ r"^(?P<indent>[ \t]*)"
44
+ r"(?P<mnemonic>[A-Za-z][A-Za-z0-9_.]*)"
45
+ r"(?:[ \t]+(?P<operands>[^#\n]*?))?"
46
+ r"(?:[ \t]*(?P<comment>#.*))?$"
47
+ )
48
+ _LABEL_RE = re.compile(r"^[ \t]*[A-Za-z_.$][\w.$]*:")
49
+
50
+
51
+ def parse_asm_line(line: str) -> AsmLine:
52
+ stripped = line.rstrip("\n")
53
+ if not stripped.strip():
54
+ return AsmLine(LineKind.BLANK, stripped)
55
+ bare = stripped.lstrip()
56
+ if bare.startswith("#") or bare.startswith("//"):
57
+ return AsmLine(LineKind.COMMENT, stripped)
58
+ if bare.startswith("."):
59
+ return AsmLine(LineKind.DIRECTIVE, stripped)
60
+ if _LABEL_RE.match(stripped):
61
+ return AsmLine(LineKind.LABEL, stripped)
62
+ m = _INSTR_RE.match(stripped)
63
+ if not m:
64
+ return AsmLine(LineKind.COMMENT, stripped)
65
+ return AsmLine(
66
+ kind=LineKind.INSTRUCTION,
67
+ raw=stripped,
68
+ indent=m.group("indent") or "",
69
+ mnemonic=m.group("mnemonic") or "",
70
+ operands=(m.group("operands") or "").strip(),
71
+ trailing_comment=(m.group("comment") or "").strip(),
72
+ )
73
+
74
+
75
+ # ---------------------------------------------------------------------------
76
+ # Catalog lookup
77
+ # ---------------------------------------------------------------------------
78
+
79
+
80
+ # Common AT&T size suffixes that may be absent from the catalog's Intel form.
81
+ _ATT_SUFFIXES = ("b", "w", "l", "q", "s", "d", "t")
82
+
83
+
84
+ def _lookup_variants(mnemonic: str) -> list[str]:
85
+ """Candidate mnemonics to try against the catalog, in preference order."""
86
+ base = mnemonic.lower()
87
+ out = [base]
88
+ # Strip a single trailing size suffix (AT&T style) if present.
89
+ if len(base) > 2 and base[-1] in _ATT_SUFFIXES:
90
+ trimmed = base[:-1]
91
+ if trimmed not in out:
92
+ out.append(trimmed)
93
+ return out
94
+
95
+
96
+ def lookup(mnemonic: str, conn: sqlite3.Connection) -> list[InstructionRecord]:
97
+ for cand in _lookup_variants(mnemonic):
98
+ records = load_instructions_by_mnemonic_from_db(conn, cand)
99
+ if records:
100
+ return records
101
+ return []
102
+
103
+
104
+ _SIZE_SPECIFIER_RE = re.compile(
105
+ r"\b(byte|word|dword|qword|xmmword|ymmword|zmmword)\s+ptr\b",
106
+ re.IGNORECASE,
107
+ )
108
+ _SIZE_TO_BITS = {
109
+ "byte": "8", "word": "16", "dword": "32", "qword": "64",
110
+ "xmmword": "128", "ymmword": "256", "zmmword": "512",
111
+ }
112
+ # Register classes ordered so longer/wider names match first.
113
+ _REG_PATTERNS: tuple[tuple[re.Pattern[str], str], ...] = (
114
+ (re.compile(r"\b(?:r[abcd]x|r[sd]i|rbp|rsp|r(?:8|9|1[0-5]))\b"), "R64"),
115
+ (re.compile(r"\b(?:e[abcd]x|e[sd]i|ebp|esp|r(?:8|9|1[0-5])d)\b"), "R32"),
116
+ (re.compile(r"\b(?:[abcd]x|[sd]i|bp|sp|r(?:8|9|1[0-5])w)\b"), "R16"),
117
+ (re.compile(r"\b(?:[abcd][lh]|[sd]il|bpl|spl|r(?:8|9|1[0-5])b)\b"), "R8"),
118
+ (re.compile(r"\bxmm\d+\b"), "XMM"),
119
+ (re.compile(r"\bymm\d+\b"), "YMM"),
120
+ (re.compile(r"\bzmm\d+\b"), "ZMM"),
121
+ (re.compile(r"\bk[0-7]\b"), "K"),
122
+ )
123
+
124
+
125
+ def _operand_width_tokens(operands: str) -> list[str]:
126
+ if not operands:
127
+ return []
128
+ s = operands.lower()
129
+ tokens: list[str] = []
130
+ for m in _SIZE_SPECIFIER_RE.finditer(s):
131
+ bits = _SIZE_TO_BITS[m.group(1).lower()]
132
+ tokens.append(f"M{bits}")
133
+ for regex, tok in _REG_PATTERNS:
134
+ if regex.search(s):
135
+ tokens.append(tok)
136
+ return tokens
137
+
138
+
139
+ def _operand_match_score(record: InstructionRecord, tokens: list[str]) -> int:
140
+ if not tokens:
141
+ return 0
142
+ key = str(getattr(record, "key", "") or "").upper()
143
+ score = 0
144
+ for tok in tokens:
145
+ if re.search(rf"\b{tok}\b", key):
146
+ score += 10
147
+ elif tok.startswith("R") and ("M" + tok[1:]) in key:
148
+ score += 2
149
+ return score
150
+
151
+
152
+ def pick_record(
153
+ records: list[InstructionRecord],
154
+ *,
155
+ arch: str | None = None,
156
+ operands: str = "",
157
+ ) -> InstructionRecord | None:
158
+ if not records:
159
+ return None
160
+ candidates = records
161
+ if arch is not None:
162
+ pinned = [r for r in records if arch in (r.arch_details or {})]
163
+ if pinned:
164
+ candidates = pinned
165
+
166
+ op_tokens = _operand_width_tokens(operands)
167
+
168
+ def score(rec: InstructionRecord) -> tuple[int, int, int]:
169
+ measured = sum(
170
+ 1 for d in (rec.arch_details or {}).values()
171
+ if (d.get("source_kind") or "measured") == "measured"
172
+ )
173
+ op_score = _operand_match_score(rec, op_tokens)
174
+ # Operand match dominates; measurement coverage tiebreaks.
175
+ return (op_score, measured, len(rec.arch_details or {}))
176
+
177
+ return max(candidates, key=score)
178
+
179
+
180
+ # ---------------------------------------------------------------------------
181
+ # Aggregation
182
+ # ---------------------------------------------------------------------------
183
+
184
+
185
+ @dataclass(slots=True)
186
+ class PerfSummary:
187
+ latency: float | None
188
+ cpi: float | None
189
+ n_archs: int
190
+ source_kind: str # "measured", "modeled", "mixed"
191
+ archs_used: list[str] = field(default_factory=list)
192
+
193
+
194
+ def _per_arch_value(details: dict[str, Any], value_fn) -> float | None:
195
+ for v in value_fn(details):
196
+ try:
197
+ return float(v)
198
+ except (TypeError, ValueError):
199
+ continue
200
+ return None
201
+
202
+
203
+ def _latency_for(details: dict[str, Any]) -> list[str]:
204
+ return perf.latency_cycle_values(details.get("latencies") or [])
205
+
206
+
207
+ def _cpi_for(details: dict[str, Any]) -> list[str]:
208
+ return perf._cpi_values(details)
209
+
210
+
211
+ def aggregate_perf(
212
+ record: InstructionRecord,
213
+ *,
214
+ mode: str = "avg",
215
+ include_modeled: bool = False,
216
+ ) -> PerfSummary:
217
+ """Aggregate latency and CPI across ``record``'s measured microarches."""
218
+ arch_details = record.arch_details or {}
219
+
220
+ def collect(kinds: tuple[str, ...]) -> tuple[list[float], list[float], list[str]]:
221
+ lats: list[float] = []
222
+ cpis: list[float] = []
223
+ archs: list[str] = []
224
+ for core, details in arch_details.items():
225
+ kind = details.get("source_kind") or "measured"
226
+ if kind not in kinds:
227
+ continue
228
+ lat = _per_arch_value(details, _latency_for)
229
+ cpi = _per_arch_value(details, _cpi_for)
230
+ if lat is None and cpi is None:
231
+ continue
232
+ archs.append(core)
233
+ if lat is not None:
234
+ lats.append(lat)
235
+ if cpi is not None:
236
+ cpis.append(cpi)
237
+ return lats, cpis, archs
238
+
239
+ lats, cpis, archs = collect(("measured",))
240
+ source_kind = "measured"
241
+ if not archs and include_modeled:
242
+ lats, cpis, archs = collect(("modeled",))
243
+ source_kind = "modeled"
244
+ if not archs:
245
+ # Last-ditch: anything at all.
246
+ lats, cpis, archs = collect(("measured", "modeled"))
247
+ source_kind = "mixed"
248
+
249
+ def reduce(values: list[float]) -> float | None:
250
+ if not values:
251
+ return None
252
+ if mode == "avg":
253
+ return statistics.fmean(values)
254
+ if mode == "median":
255
+ return statistics.median(values)
256
+ if mode == "best":
257
+ return min(values)
258
+ if mode == "worst":
259
+ return max(values)
260
+ raise ValueError(f"unknown --agg mode: {mode}")
261
+
262
+ return PerfSummary(
263
+ latency=reduce(lats),
264
+ cpi=reduce(cpis),
265
+ n_archs=len(archs),
266
+ source_kind=source_kind,
267
+ archs_used=archs,
268
+ )
269
+
270
+
271
+ # ---------------------------------------------------------------------------
272
+ # Formatting
273
+ # ---------------------------------------------------------------------------
274
+
275
+
276
+ def _fmt_num(x: float | None) -> str:
277
+ if x is None:
278
+ return "-"
279
+ if float(x).is_integer():
280
+ return f"{x:.1f}"
281
+ return f"{x:.2f}"
282
+
283
+
284
+ def _arch_perf(
285
+ record: InstructionRecord, arch: str
286
+ ) -> tuple[float | None, float | None, str]:
287
+ details = (record.arch_details or {}).get(arch) or {}
288
+ kind = details.get("source_kind") or "measured"
289
+ lat = _per_arch_value(details, _latency_for)
290
+ cpi = _per_arch_value(details, _cpi_for)
291
+ return lat, cpi, kind
292
+
293
+
294
+ _SUMMARY_BITS_RE = re.compile(r"(\d+)\s*-?\s*bit", re.IGNORECASE)
295
+ _FORM_WIDTH_RE = re.compile(r"[MRI](\d+)")
296
+
297
+
298
+ def _summary_matches_form(summary: str, record: InstructionRecord) -> bool:
299
+ """Drop obviously mislabeled SDM summaries (e.g. the generic MOV
300
+ blurb "Move 32-bit integer operands." attached to MOV (M64, R64)).
301
+
302
+ We only flag a summary as wrong when it declares an explicit bit-width
303
+ that appears nowhere among the form's operand widths."""
304
+ m = _SUMMARY_BITS_RE.search(summary or "")
305
+ if not m:
306
+ return True
307
+ declared = m.group(1)
308
+ widths: list[str] = []
309
+ for op in record.operand_details or []:
310
+ w = op.get("width")
311
+ if w:
312
+ widths.append(str(w))
313
+ if not widths:
314
+ key = str(getattr(record, "key", "") or "")
315
+ widths = _FORM_WIDTH_RE.findall(key)
316
+ if not widths:
317
+ return True
318
+ return declared in widths
319
+
320
+
321
+ def format_annotation(
322
+ record: InstructionRecord,
323
+ *,
324
+ performance: bool,
325
+ docs: bool,
326
+ arch: str | None,
327
+ agg: str,
328
+ include_modeled: bool,
329
+ ) -> str:
330
+ """Compose the comment fragment (without the leading ``# `` marker)."""
331
+ parts: list[str] = []
332
+ if docs and record.summary and _summary_matches_form(record.summary, record):
333
+ parts.append(record.summary.strip())
334
+ if performance:
335
+ if arch is not None:
336
+ lat, cpi, kind = _arch_perf(record, arch)
337
+ tag = f"[{arch}, {kind}]"
338
+ else:
339
+ summary = aggregate_perf(
340
+ record, mode=agg, include_modeled=include_modeled
341
+ )
342
+ lat, cpi = summary.latency, summary.cpi
343
+ if summary.n_archs == 0:
344
+ tag = "[no data]"
345
+ elif arch is None and agg == "avg":
346
+ tag = f"[avg of {summary.n_archs} archs, {summary.source_kind}]"
347
+ else:
348
+ tag = f"[{agg} of {summary.n_archs} archs, {summary.source_kind}]"
349
+ perf_frag = f"lat={_fmt_num(lat)}c cpi={_fmt_num(cpi)} {tag}"
350
+ parts.append(perf_frag)
351
+ return " | ".join(parts)
352
+
353
+
354
+ # ---------------------------------------------------------------------------
355
+ # Options & streaming
356
+ # ---------------------------------------------------------------------------
357
+
358
+
359
+ @dataclass(slots=True)
360
+ class AnnotateOptions:
361
+ performance: bool = True
362
+ docs: bool = True
363
+ arch: str | None = None
364
+ agg: str = "avg"
365
+ include_modeled: bool = False
366
+ block: bool = False # otherwise inline
367
+ unknown: str = "mark" # "keep" | "drop" | "mark"
368
+ fmt: str = "sa" # "sa" | "md" | "json"
369
+
370
+
371
+ def _annotate_instruction(
372
+ parsed: AsmLine,
373
+ opts: AnnotateOptions,
374
+ conn: sqlite3.Connection,
375
+ ) -> tuple[str, dict[str, Any] | None]:
376
+ """Return the rendered output line and an optional JSON record."""
377
+ records = lookup(parsed.mnemonic, conn)
378
+ record = pick_record(records, arch=opts.arch, operands=parsed.operands)
379
+
380
+ if record is None:
381
+ if opts.unknown == "drop":
382
+ return parsed.raw, None
383
+ if opts.unknown == "mark":
384
+ marker = "# ??"
385
+ if parsed.trailing_comment:
386
+ return parsed.raw, None
387
+ return f"{parsed.raw} {marker}", {
388
+ "mnemonic": parsed.mnemonic,
389
+ "known": False,
390
+ }
391
+ return parsed.raw, None
392
+
393
+ if not (opts.performance or opts.docs):
394
+ return parsed.raw, None
395
+
396
+ annotation = format_annotation(
397
+ record,
398
+ performance=opts.performance,
399
+ docs=opts.docs,
400
+ arch=opts.arch,
401
+ agg=opts.agg,
402
+ include_modeled=opts.include_modeled,
403
+ )
404
+ if not annotation:
405
+ return parsed.raw, None
406
+
407
+ json_record = {
408
+ "mnemonic": parsed.mnemonic,
409
+ "known": True,
410
+ "summary": record.summary,
411
+ "annotation": annotation,
412
+ }
413
+
414
+ if opts.block:
415
+ block_line = f"{parsed.indent}# {annotation}"
416
+ return f"{block_line}\n{parsed.raw}", json_record
417
+
418
+ # Inline: append after the raw line, respecting any pre-existing comment.
419
+ if parsed.trailing_comment:
420
+ return parsed.raw, json_record
421
+ return f"{parsed.raw} # {annotation}", json_record
422
+
423
+
424
+ def annotate_stream(
425
+ lines: Iterable[str],
426
+ *,
427
+ opts: AnnotateOptions,
428
+ conn: sqlite3.Connection,
429
+ ) -> Iterator[str]:
430
+ """Yield annotated lines for each input line (newline-terminated)."""
431
+ json_records: list[dict[str, Any]] = []
432
+ collecting_json = opts.fmt == "json"
433
+
434
+ for line in lines:
435
+ parsed = parse_asm_line(line)
436
+ if parsed.kind != LineKind.INSTRUCTION:
437
+ if not collecting_json:
438
+ yield parsed.raw + "\n"
439
+ continue
440
+ out_line, record = _annotate_instruction(parsed, opts, conn)
441
+ if collecting_json:
442
+ if record is not None:
443
+ json_records.append(record)
444
+ continue
445
+ yield out_line + "\n"
446
+
447
+ if collecting_json:
448
+ yield json.dumps(json_records, indent=2) + "\n"