dirag 0.1.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.
- dirag/__init__.py +3 -0
- dirag/__main__.py +5 -0
- dirag/answer.py +129 -0
- dirag/cli.py +126 -0
- dirag/config.py +85 -0
- dirag/embed.py +130 -0
- dirag/index.py +405 -0
- dirag/jobs.py +129 -0
- dirag/llm.py +121 -0
- dirag/rerank.py +69 -0
- dirag/search.py +177 -0
- dirag/server.py +414 -0
- dirag/state.py +45 -0
- dirag/static/app.css +206 -0
- dirag/static/app.js +696 -0
- dirag/static/favicon.svg +1 -0
- dirag/static/icons.js +27 -0
- dirag/static/index.html +85 -0
- dirag/static/vendor/LICENSE.pdfjs +177 -0
- dirag/static/vendor/pdf.min.mjs +21 -0
- dirag/static/vendor/pdf.worker.min.mjs +21 -0
- dirag/toc.py +141 -0
- dirag-0.1.0.dist-info/METADATA +196 -0
- dirag-0.1.0.dist-info/RECORD +27 -0
- dirag-0.1.0.dist-info/WHEEL +4 -0
- dirag-0.1.0.dist-info/entry_points.txt +2 -0
- dirag-0.1.0.dist-info/licenses/LICENSE +21 -0
dirag/__init__.py
ADDED
dirag/__main__.py
ADDED
dirag/answer.py
ADDED
|
@@ -0,0 +1,129 @@
|
|
|
1
|
+
"""The answer: quotes chosen by a language model, verified against the passages and located on the page.
|
|
2
|
+
|
|
3
|
+
The model sees the top passages of a search and returns passage numbers with
|
|
4
|
+
the words it quotes from each, plus a response of one to three sentences citing
|
|
5
|
+
them by number. Nothing
|
|
6
|
+
it writes is shown as a quote:
|
|
7
|
+
|
|
8
|
+
- Each quote is matched against its passage on letters and digits only, so
|
|
9
|
+
case, punctuation, spacing and line-break hyphens do not matter. A quote
|
|
10
|
+
with "..." is matched part by part, in order. A quote that does not match
|
|
11
|
+
is dropped.
|
|
12
|
+
- The text shown for a quote is the passage's own text for the matched span.
|
|
13
|
+
- The quote is located on its page through the page's words, giving one
|
|
14
|
+
rectangle per line for the reader and the page image to mark.
|
|
15
|
+
- Citations in the response to passages with no kept quote are removed, the
|
|
16
|
+
rest renumbered to the quotes. The response is kept only when at least one
|
|
17
|
+
citation remains; otherwise only the quotes are shown. With no kept quote the answer
|
|
18
|
+
is empty.
|
|
19
|
+
"""
|
|
20
|
+
|
|
21
|
+
import json
|
|
22
|
+
import re
|
|
23
|
+
|
|
24
|
+
import pymupdf
|
|
25
|
+
|
|
26
|
+
from . import llm
|
|
27
|
+
|
|
28
|
+
DEPTH = 10
|
|
29
|
+
MAX_QUOTES = 5
|
|
30
|
+
MIN_ALNUM = 15
|
|
31
|
+
|
|
32
|
+
SYSTEM = ("You answer a question using only the numbered book passages given. Reply with JSON only:\n"
|
|
33
|
+
'{"quotes": [{"n": <passage number>, "text": "<words copied exactly from that passage>"}], '
|
|
34
|
+
'"summary": "<one to three sentences answering the question, each citing passages like [2]>"}\n'
|
|
35
|
+
"Copy every quote word for word from its passage; use ... only to skip words inside a quote. "
|
|
36
|
+
"Cite only passages you quoted. "
|
|
37
|
+
f"Each quote is one to three sentences. Order quotes from most to least useful; at most {MAX_QUOTES}. "
|
|
38
|
+
'If no passage answers the question, reply {"quotes": [], "summary": ""}.')
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def _alnum(text):
|
|
42
|
+
"""Letters and digits of `text`, lowercased, with the index in `text` of each one."""
|
|
43
|
+
chars, where = [], []
|
|
44
|
+
for i, c in enumerate(text):
|
|
45
|
+
if c.isalnum():
|
|
46
|
+
chars.append(c.lower())
|
|
47
|
+
where.append(i)
|
|
48
|
+
return "".join(chars), where
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def locate(quote, text):
|
|
52
|
+
"""(start, end) of `quote` within `text`, matched on letters and digits, or None."""
|
|
53
|
+
haystack, where = _alnum(text)
|
|
54
|
+
parts = [p for p in (_alnum(part)[0] for part in re.split(r"\.\.\.|\u2026", quote)) if p]
|
|
55
|
+
if not parts or sum(len(p) for p in parts) < MIN_ALNUM:
|
|
56
|
+
return None
|
|
57
|
+
start, at = None, 0
|
|
58
|
+
for part in parts:
|
|
59
|
+
found = haystack.find(part, at)
|
|
60
|
+
if found == -1:
|
|
61
|
+
return None
|
|
62
|
+
start = found if start is None else start
|
|
63
|
+
at = found + len(part)
|
|
64
|
+
return where[start], where[at - 1] + 1
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def page_rects(path, page_number, span_text):
|
|
68
|
+
"""Line rectangles of `span_text` on a page, in PDF points, found through the page's words."""
|
|
69
|
+
doc = pymupdf.open(path)
|
|
70
|
+
try:
|
|
71
|
+
words = doc.load_page(page_number - 1).get_text("words")
|
|
72
|
+
finally:
|
|
73
|
+
doc.close()
|
|
74
|
+
joined, owner = [], []
|
|
75
|
+
for index, word in enumerate(words):
|
|
76
|
+
letters = _alnum(word[4])[0]
|
|
77
|
+
joined.append(letters)
|
|
78
|
+
owner.extend([index] * len(letters))
|
|
79
|
+
needle = _alnum(span_text)[0]
|
|
80
|
+
found = "".join(joined).find(needle)
|
|
81
|
+
if found == -1 or not needle:
|
|
82
|
+
return []
|
|
83
|
+
lines = {}
|
|
84
|
+
for index in sorted(set(owner[found:found + len(needle)])):
|
|
85
|
+
x0, y0, x1, y1, _, block, line, _ = words[index]
|
|
86
|
+
box = lines.setdefault((block, line), [x0, y0, x1, y1])
|
|
87
|
+
box[:] = [min(box[0], x0), min(box[1], y0), max(box[2], x1), max(box[3], y1)]
|
|
88
|
+
return [[round(v, 2) for v in box] for box in lines.values()]
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def compose(question, rows, file_of):
|
|
92
|
+
"""{"summary", "quotes"} for `rows` (search rows, best first). `file_of(rel)` gives a book's PDF path.
|
|
93
|
+
|
|
94
|
+
Raises llm.LLMError when the model call fails.
|
|
95
|
+
"""
|
|
96
|
+
top = rows[:DEPTH]
|
|
97
|
+
listing = "\n\n".join(f"[{n}] {row['book']} > {row['section'] or ''}, p. {row['page']}\n{row['text']}"
|
|
98
|
+
for n, row in enumerate(top, 1))
|
|
99
|
+
reply = llm.chat_json(SYSTEM, f"Question: {question}\n\nPassages:\n\n{listing}")
|
|
100
|
+
quotes, numbers = [], {}
|
|
101
|
+
for item in reply.get("quotes", []) if isinstance(reply, dict) else []:
|
|
102
|
+
try:
|
|
103
|
+
n, said = int(item["n"]), str(item["text"])
|
|
104
|
+
except (KeyError, TypeError, ValueError):
|
|
105
|
+
continue
|
|
106
|
+
if not 1 <= n <= len(top) or len(quotes) >= MAX_QUOTES:
|
|
107
|
+
continue
|
|
108
|
+
row = top[n - 1]
|
|
109
|
+
span = locate(said, row["text"])
|
|
110
|
+
if span is None:
|
|
111
|
+
continue
|
|
112
|
+
text = row["text"][span[0]:span[1]]
|
|
113
|
+
if any(q["chunk_id"] == row["id"] and q["text"] == text for q in quotes):
|
|
114
|
+
continue
|
|
115
|
+
path = file_of(row["path"])
|
|
116
|
+
quotes.append({"n": len(quotes) + 1, "chunk_id": row["id"], "book_id": row["book_id"], "book": row["book"],
|
|
117
|
+
"year": row["year"], "section": row["section"], "page": row["page"], "text": text,
|
|
118
|
+
"rects": (page_rects(path, row["page"], text) if path else []) or _bbox(row)})
|
|
119
|
+
numbers.setdefault(n, quotes[-1]["n"])
|
|
120
|
+
summary = str(reply.get("summary") or "").strip() if quotes else ""
|
|
121
|
+
summary = re.sub(r"\s*\[(\d+)\]", lambda m: f" [{numbers[int(m.group(1))]}]" if int(m.group(1)) in numbers else "",
|
|
122
|
+
summary).strip()
|
|
123
|
+
if not re.search(r"\[\d+\]", summary):
|
|
124
|
+
summary = ""
|
|
125
|
+
return {"summary": summary, "quotes": quotes}
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
def _bbox(row):
|
|
129
|
+
return [json.loads(row["bbox_json"])] if row["bbox_json"] else []
|
dirag/cli.py
ADDED
|
@@ -0,0 +1,126 @@
|
|
|
1
|
+
"""Command line: `dirag [serve]`, `dirag index`, `dirag find`, `dirag where`, `dirag toc`."""
|
|
2
|
+
|
|
3
|
+
import argparse
|
|
4
|
+
import shutil
|
|
5
|
+
import sys
|
|
6
|
+
import textwrap
|
|
7
|
+
|
|
8
|
+
from . import __version__, config
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def cmd_serve(args):
|
|
12
|
+
from .server import serve
|
|
13
|
+
if args.library:
|
|
14
|
+
config.set_library(args.library)
|
|
15
|
+
return serve(args.host, args.port, browser=not args.no_browser)
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def cmd_index(args):
|
|
19
|
+
from .index import run
|
|
20
|
+
root = config.require_library(args.library)
|
|
21
|
+
action = "reindex" if args.reindex else "rechunk" if args.rechunk else "update"
|
|
22
|
+
return run(root, config.index_path(root), action=action, limit=args.limit)
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def _index():
|
|
26
|
+
return config.index_path(config.require_library())
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def cmd_find(args):
|
|
30
|
+
from . import search
|
|
31
|
+
conn = search.connect(_index())
|
|
32
|
+
try:
|
|
33
|
+
rows = search.passages(conn, args.query, args.mode, args.rerank, limit=args.k, per_book=args.per_book)
|
|
34
|
+
finally:
|
|
35
|
+
conn.close()
|
|
36
|
+
if not rows:
|
|
37
|
+
print("No matches.")
|
|
38
|
+
return 0
|
|
39
|
+
width = min(shutil.get_terminal_size((100, 24)).columns, 100)
|
|
40
|
+
for number, row in enumerate(rows, 1):
|
|
41
|
+
print(f"{number:>3} {row['book'][:66]} p.{row['page']}{' (dropped by llm)' if row.get('dropped') else ''}")
|
|
42
|
+
if row["section"]:
|
|
43
|
+
print(f" {row['section'][:90]}")
|
|
44
|
+
print(textwrap.fill(row["text"], width=width, initial_indent=" ", subsequent_indent=" "))
|
|
45
|
+
print("")
|
|
46
|
+
return 0
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def cmd_where(args):
|
|
50
|
+
from . import search
|
|
51
|
+
conn = search.connect(_index())
|
|
52
|
+
try:
|
|
53
|
+
ranked = search.chapters(conn, search.candidates(conn, args.query, args.mode), limit=args.k)
|
|
54
|
+
finally:
|
|
55
|
+
conn.close()
|
|
56
|
+
if not ranked:
|
|
57
|
+
print("No matches.")
|
|
58
|
+
return 0
|
|
59
|
+
print(f"{'hits':>5} {'pages':>11} book / chapter")
|
|
60
|
+
for entry in ranked:
|
|
61
|
+
span = f"{min(entry['pages'])}-{max(entry['pages'])}"
|
|
62
|
+
stamp = f" ({entry['year']})" if entry["year"] else ""
|
|
63
|
+
print(f"{entry['hits']:>5} {span:>11} {entry['book'][:60]}{stamp}")
|
|
64
|
+
print(f"{'':>5} {'':>11} {entry['chapter'][:70]}")
|
|
65
|
+
return 0
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def cmd_toc(args):
|
|
69
|
+
import pymupdf
|
|
70
|
+
from .toc import chapters_of, sections_of
|
|
71
|
+
doc = pymupdf.open(args.pdf)
|
|
72
|
+
try:
|
|
73
|
+
sections = sections_of(doc)
|
|
74
|
+
if not sections:
|
|
75
|
+
print("No usable table of contents.")
|
|
76
|
+
return 0
|
|
77
|
+
for chapter in chapters_of(sections):
|
|
78
|
+
print(f"p.{chapter['start_page']:>5}-{chapter['end_page']:>5} {chapter['sections']:>4} sec {chapter['title'][:80]}")
|
|
79
|
+
finally:
|
|
80
|
+
doc.close()
|
|
81
|
+
return 0
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def main(argv=None):
|
|
85
|
+
parser = argparse.ArgumentParser(prog="dirag", description="Search a folder of PDF books and read the page.")
|
|
86
|
+
parser.add_argument("--version", action="version", version=f"dirag {__version__}")
|
|
87
|
+
sub = parser.add_subparsers(dest="command")
|
|
88
|
+
|
|
89
|
+
serve = sub.add_parser("serve", help="run the web app (the default)")
|
|
90
|
+
serve.add_argument("--library", help="the folder of PDFs; remembered once given")
|
|
91
|
+
serve.add_argument("--host", default="127.0.0.1")
|
|
92
|
+
serve.add_argument("--port", type=int, default=8008)
|
|
93
|
+
serve.add_argument("--no-browser", action="store_true", help="do not open the browser")
|
|
94
|
+
serve.set_defaults(func=cmd_serve)
|
|
95
|
+
|
|
96
|
+
index = sub.add_parser("index", help="index new and changed PDFs, drop deleted ones")
|
|
97
|
+
index.add_argument("--library", help="the folder of PDFs; remembered once given")
|
|
98
|
+
index.add_argument("--limit", type=int, default=0, help="only the first N PDFs")
|
|
99
|
+
mode = index.add_mutually_exclusive_group()
|
|
100
|
+
mode.add_argument("--reindex", action="store_true", help="parse and embed every PDF again")
|
|
101
|
+
mode.add_argument("--rechunk", action="store_true", help="rebuild passages and vectors from cached pages")
|
|
102
|
+
index.set_defaults(func=cmd_index)
|
|
103
|
+
|
|
104
|
+
find = sub.add_parser("find", help="passages for a query")
|
|
105
|
+
find.add_argument("query")
|
|
106
|
+
find.add_argument("-k", type=int, default=10)
|
|
107
|
+
find.add_argument("--per-book", type=int, default=2)
|
|
108
|
+
find.add_argument("--mode", choices=("lexical", "semantic", "hybrid"), default="hybrid")
|
|
109
|
+
find.add_argument("--rerank", choices=("off", "neural", "llm"), default="off")
|
|
110
|
+
find.set_defaults(func=cmd_find)
|
|
111
|
+
|
|
112
|
+
where = sub.add_parser("where", help="chapters for a query")
|
|
113
|
+
where.add_argument("query")
|
|
114
|
+
where.add_argument("-k", type=int, default=10)
|
|
115
|
+
where.add_argument("--mode", choices=("lexical", "semantic", "hybrid"), default="hybrid")
|
|
116
|
+
where.set_defaults(func=cmd_where)
|
|
117
|
+
|
|
118
|
+
toc = sub.add_parser("toc", help="the chapter map of one PDF")
|
|
119
|
+
toc.add_argument("pdf")
|
|
120
|
+
toc.set_defaults(func=cmd_toc)
|
|
121
|
+
|
|
122
|
+
argv = list(sys.argv[1:] if argv is None else argv)
|
|
123
|
+
if not argv or (argv[0] not in sub.choices and argv[0] not in ("-h", "--help", "--version")):
|
|
124
|
+
argv = ["serve", *argv]
|
|
125
|
+
args = parser.parse_args(argv)
|
|
126
|
+
return args.func(args)
|
dirag/config.py
ADDED
|
@@ -0,0 +1,85 @@
|
|
|
1
|
+
"""Where dirag keeps its files, and the one library folder it serves.
|
|
2
|
+
|
|
3
|
+
Two directories, both overridable:
|
|
4
|
+
|
|
5
|
+
DIRAG_HOME derived data: one index per library folder (indexes/), the
|
|
6
|
+
embedding model cache (models/) and the indexing job record
|
|
7
|
+
(job.json). Default ~/.local/share/dirag. Rebuildable.
|
|
8
|
+
DIRAG_STATE user state: config.json, positions.json, bookmarks.json,
|
|
9
|
+
cards.json. Default DIRAG_HOME. Small, textual, not derived.
|
|
10
|
+
|
|
11
|
+
The library is a single folder of PDFs. It is resolved in this order: the
|
|
12
|
+
--library option, the DIRAG_LIBRARY variable, then config.json. Passing
|
|
13
|
+
--library, or choosing a folder in the app, saves it to config.json. With
|
|
14
|
+
DIRAG_LIBRARY set the folder is fixed and the app cannot change it.
|
|
15
|
+
|
|
16
|
+
The app's folder picker lists folders under DIRAG_BROWSE_ROOT only (default:
|
|
17
|
+
the home directory).
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
import hashlib
|
|
21
|
+
import json
|
|
22
|
+
import os
|
|
23
|
+
import sys
|
|
24
|
+
from pathlib import Path
|
|
25
|
+
|
|
26
|
+
HOME = Path(os.getenv("DIRAG_HOME") or Path.home() / ".local" / "share" / "dirag").expanduser()
|
|
27
|
+
STATE = Path(os.getenv("DIRAG_STATE") or HOME).expanduser()
|
|
28
|
+
MODELS = HOME / "models"
|
|
29
|
+
JOB = HOME / "job.json"
|
|
30
|
+
CONFIG = STATE / "config.json"
|
|
31
|
+
FIXED = bool(os.getenv("DIRAG_LIBRARY"))
|
|
32
|
+
BROWSE_ROOT = Path(os.getenv("DIRAG_BROWSE_ROOT") or Path.home()).expanduser().resolve()
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def read_json(path, empty):
|
|
36
|
+
"""The file's content, or `empty` when it is missing, unreadable or the wrong type."""
|
|
37
|
+
try:
|
|
38
|
+
data = json.loads(Path(path).read_text(encoding="utf-8"))
|
|
39
|
+
return data if isinstance(data, type(empty)) else empty
|
|
40
|
+
except FileNotFoundError:
|
|
41
|
+
return empty
|
|
42
|
+
except (OSError, ValueError) as exc:
|
|
43
|
+
print(f"{Path(path).name} unusable ({exc}); ignoring it", file=sys.stderr)
|
|
44
|
+
return empty
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def write_json(path, data):
|
|
48
|
+
"""Write through a temp file and rename, so a reader never sees half a file."""
|
|
49
|
+
path = Path(path)
|
|
50
|
+
path.parent.mkdir(parents=True, exist_ok=True)
|
|
51
|
+
tmp = path.with_suffix(path.suffix + ".tmp")
|
|
52
|
+
tmp.write_text(json.dumps(data, indent=1, sort_keys=True) + "\n", encoding="utf-8")
|
|
53
|
+
tmp.replace(path)
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def index_path(root):
|
|
57
|
+
"""The index file for a library folder: its name plus a hash of its absolute path."""
|
|
58
|
+
digest = hashlib.sha1(str(root).encode()).hexdigest()[:10]
|
|
59
|
+
return HOME / "indexes" / f"{root.name or 'root'}-{digest}.sqlite3"
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def set_library(path):
|
|
63
|
+
"""Save `path` as the library folder and return it as an absolute Path."""
|
|
64
|
+
root = Path(path).expanduser().resolve()
|
|
65
|
+
if not root.is_dir():
|
|
66
|
+
raise SystemExit(f"Not a folder: {root}")
|
|
67
|
+
settings = read_json(CONFIG, {})
|
|
68
|
+
settings["library"] = str(root)
|
|
69
|
+
write_json(CONFIG, settings)
|
|
70
|
+
return root
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def library(override=None):
|
|
74
|
+
"""The library folder as an absolute Path, or None when none is set."""
|
|
75
|
+
if override:
|
|
76
|
+
return set_library(override)
|
|
77
|
+
value = os.getenv("DIRAG_LIBRARY") or read_json(CONFIG, {}).get("library")
|
|
78
|
+
return Path(value).expanduser().resolve() if value else None
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def require_library(override=None):
|
|
82
|
+
root = library(override)
|
|
83
|
+
if root is None:
|
|
84
|
+
raise SystemExit("No library folder set. Pass --library PATH once; it is remembered.")
|
|
85
|
+
return root
|
dirag/embed.py
ADDED
|
@@ -0,0 +1,130 @@
|
|
|
1
|
+
"""The local models (the embedder here, the cross-encoder in rerank.py) and where they run.
|
|
2
|
+
|
|
3
|
+
They run on an NVIDIA GPU when one is present and working, else on the CPU.
|
|
4
|
+
Present means an NVIDIA driver is loaded and onnxruntime has its CUDA provider;
|
|
5
|
+
working means a first inference succeeds. A GPU that fails that check is
|
|
6
|
+
reported once and the CPU is used instead.
|
|
7
|
+
|
|
8
|
+
DIRAG_DEVICE=cpu never use the GPU
|
|
9
|
+
DIRAG_EMBED_THREADS CPU threads for onnxruntime; default all cores
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
import os
|
|
13
|
+
import shutil
|
|
14
|
+
import subprocess
|
|
15
|
+
import sys
|
|
16
|
+
import threading
|
|
17
|
+
import warnings
|
|
18
|
+
from pathlib import Path
|
|
19
|
+
|
|
20
|
+
from . import config
|
|
21
|
+
|
|
22
|
+
MODEL = "BAAI/bge-small-en-v1.5"
|
|
23
|
+
DIM = 384
|
|
24
|
+
THREADS = int(os.getenv("DIRAG_EMBED_THREADS")) if os.getenv("DIRAG_EMBED_THREADS") else None
|
|
25
|
+
|
|
26
|
+
_model = None
|
|
27
|
+
_lock = threading.Lock()
|
|
28
|
+
_gpu = None # None until decided, then True or False
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def gpu_name():
|
|
32
|
+
"""The NVIDIA GPU's name, or None when there is no usable one. Does not load onnxruntime.
|
|
33
|
+
|
|
34
|
+
nvidia-smi naming a GPU is the test; without nvidia-smi, the GPU device file is.
|
|
35
|
+
Driver files alone (as in a container without the GPU passed through) do not count.
|
|
36
|
+
"""
|
|
37
|
+
if os.getenv("DIRAG_DEVICE", "").lower() == "cpu":
|
|
38
|
+
return None
|
|
39
|
+
if shutil.which("nvidia-smi"):
|
|
40
|
+
try:
|
|
41
|
+
out = subprocess.run(["nvidia-smi", "--query-gpu=name", "--format=csv,noheader"],
|
|
42
|
+
capture_output=True, text=True, timeout=5)
|
|
43
|
+
names = out.stdout.strip().splitlines()
|
|
44
|
+
return names[0] if out.returncode == 0 and names else None
|
|
45
|
+
except (OSError, subprocess.SubprocessError):
|
|
46
|
+
return None
|
|
47
|
+
return "NVIDIA GPU" if Path("/dev/nvidia0").exists() else None
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def use_gpu():
|
|
51
|
+
"""Whether the local models should try the GPU."""
|
|
52
|
+
global _gpu
|
|
53
|
+
if _gpu is None:
|
|
54
|
+
_gpu = False
|
|
55
|
+
if gpu_name():
|
|
56
|
+
import onnxruntime
|
|
57
|
+
if "CUDAExecutionProvider" in onnxruntime.get_available_providers():
|
|
58
|
+
# CUDA and cuDNN installed as pip packages load from here.
|
|
59
|
+
try:
|
|
60
|
+
onnxruntime.preload_dlls()
|
|
61
|
+
except Exception:
|
|
62
|
+
pass
|
|
63
|
+
_gpu = True
|
|
64
|
+
return _gpu
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def build(make, probe):
|
|
68
|
+
"""make(**options) on the GPU when it works, else on the CPU. probe(model) runs one inference.
|
|
69
|
+
|
|
70
|
+
onnxruntime falls back to the CPU on its own when the CUDA provider cannot
|
|
71
|
+
start, with only a warning; that warning counts as the GPU not working.
|
|
72
|
+
"""
|
|
73
|
+
global _gpu
|
|
74
|
+
import onnxruntime
|
|
75
|
+
# dirag reports the device itself; onnxruntime's own provider errors are noise here.
|
|
76
|
+
onnxruntime.set_default_logger_severity(4)
|
|
77
|
+
if use_gpu():
|
|
78
|
+
try:
|
|
79
|
+
with warnings.catch_warnings(record=True) as caught:
|
|
80
|
+
warnings.simplefilter("always")
|
|
81
|
+
model = make(cuda=True, device_ids=[0])
|
|
82
|
+
probe(model)
|
|
83
|
+
failed = [str(w.message) for w in caught if "CUDAExecutionProvider" in str(w.message)]
|
|
84
|
+
if not failed:
|
|
85
|
+
return model
|
|
86
|
+
reason = failed[0]
|
|
87
|
+
except Exception as exc:
|
|
88
|
+
reason = str(exc)
|
|
89
|
+
print(f"dirag: the GPU did not work ({reason.splitlines()[0][:160]}); using the CPU", file=sys.stderr, flush=True)
|
|
90
|
+
_gpu = False
|
|
91
|
+
with warnings.catch_warnings():
|
|
92
|
+
warnings.simplefilter("ignore", RuntimeWarning)
|
|
93
|
+
model = make(threads=THREADS, providers=["CPUExecutionProvider"])
|
|
94
|
+
probe(model)
|
|
95
|
+
return model
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
def _load():
|
|
99
|
+
global _model
|
|
100
|
+
with _lock:
|
|
101
|
+
if _model is None:
|
|
102
|
+
from fastembed import TextEmbedding
|
|
103
|
+
_model = build(lambda **o: TextEmbedding(model_name=MODEL, cache_dir=str(config.MODELS), **o),
|
|
104
|
+
lambda m: list(m.embed(["probe"])))
|
|
105
|
+
return _model
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def batch_size():
|
|
109
|
+
"""Passages per embedding call: large on the GPU, small on the CPU so a stop is honoured within seconds."""
|
|
110
|
+
_load()
|
|
111
|
+
return 256 if _gpu else 32
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
def passages(texts):
|
|
115
|
+
"""Vectors for a batch of passages, as lists of floats."""
|
|
116
|
+
texts = list(texts)
|
|
117
|
+
return [vector.tolist() for vector in _load().embed(texts)] if texts else []
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
def query(text):
|
|
121
|
+
"""Vector for one query. bge models prefix queries with an instruction; query_embed adds it."""
|
|
122
|
+
return list(_load().query_embed([text]))[0].tolist()
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
def preflight():
|
|
126
|
+
"""Load the embedder before a long run, so a broken setup fails before the first book."""
|
|
127
|
+
try:
|
|
128
|
+
_load()
|
|
129
|
+
except Exception as exc:
|
|
130
|
+
raise SystemExit(f"Embedding is not working:\n {exc}")
|