keywordmoves 0.4.2__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.
- keywordmoves/__init__.py +26 -0
- keywordmoves/__main__.py +4 -0
- keywordmoves/builtin/__init__.py +2 -0
- keywordmoves/builtin/google_trends.py +252 -0
- keywordmoves/builtin/huggingface_llm.py +128 -0
- keywordmoves/builtin/keybert_keywords.py +468 -0
- keywordmoves/builtin/native_export.py +171 -0
- keywordmoves/builtin/nltk_keywords.py +530 -0
- keywordmoves/builtin/observed_evidence.py +102 -0
- keywordmoves/builtin/openai_llm.py +197 -0
- keywordmoves/builtin/spacy_keywords.py +438 -0
- keywordmoves/builtin/text_library.py +302 -0
- keywordmoves/cli.py +139 -0
- keywordmoves/errors.py +15 -0
- keywordmoves/models.py +77 -0
- keywordmoves/monitoring/__init__.py +1 -0
- keywordmoves/monitoring/analysis.py +276 -0
- keywordmoves/monitoring/captures.py +56 -0
- keywordmoves/monitoring/cli.py +199 -0
- keywordmoves/monitoring/demo.py +79 -0
- keywordmoves/monitoring/imports.py +163 -0
- keywordmoves/monitoring/runner.py +149 -0
- keywordmoves/monitoring/store.py +240 -0
- keywordmoves/monitoring/validation.py +246 -0
- keywordmoves/online/__init__.py +45 -0
- keywordmoves/online/bing_search.py +229 -0
- keywordmoves/online/bing_search_analysis.py +542 -0
- keywordmoves/online/bing_search_imports.py +186 -0
- keywordmoves/online/bing_search_providers.py +384 -0
- keywordmoves/online/bing_webmaster.py +250 -0
- keywordmoves/online/commercial.py +252 -0
- keywordmoves/online/common.py +269 -0
- keywordmoves/online/discovery.py +167 -0
- keywordmoves/online/google.py +160 -0
- keywordmoves/online/google_search.py +280 -0
- keywordmoves/online/google_search_analysis.py +409 -0
- keywordmoves/online/google_search_console.py +314 -0
- keywordmoves/online/google_search_imports.py +250 -0
- keywordmoves/online/google_search_providers.py +414 -0
- keywordmoves/online/instagram.py +555 -0
- keywordmoves/online/instagram_analysis.py +415 -0
- keywordmoves/online/instagram_imports.py +145 -0
- keywordmoves/online/reddit.py +355 -0
- keywordmoves/online/reddit_analysis.py +422 -0
- keywordmoves/online/reddit_imports.py +245 -0
- keywordmoves/online/reddit_providers.py +220 -0
- keywordmoves/online/tiktok.py +523 -0
- keywordmoves/online/tiktok_analysis.py +560 -0
- keywordmoves/online/tiktok_imports.py +218 -0
- keywordmoves/online/websites.py +299 -0
- keywordmoves/online/youtube.py +494 -0
- keywordmoves/online/youtube_analysis.py +460 -0
- keywordmoves/online/youtube_imports.py +207 -0
- keywordmoves/online/youtube_providers.py +306 -0
- keywordmoves/protocols.py +24 -0
- keywordmoves/registry.py +87 -0
- keywordmoves-0.4.2.data/data/share/keywordmoves/branding/README.md +11 -0
- keywordmoves-0.4.2.data/data/share/keywordmoves/branding/logo-monochrome.svg +6 -0
- keywordmoves-0.4.2.data/data/share/keywordmoves/branding/logo.png +0 -0
- keywordmoves-0.4.2.data/data/share/keywordmoves/branding/logo.svg +6 -0
- keywordmoves-0.4.2.dist-info/METADATA +454 -0
- keywordmoves-0.4.2.dist-info/RECORD +66 -0
- keywordmoves-0.4.2.dist-info/WHEEL +5 -0
- keywordmoves-0.4.2.dist-info/entry_points.txt +39 -0
- keywordmoves-0.4.2.dist-info/licenses/LICENSE +22 -0
- keywordmoves-0.4.2.dist-info/top_level.txt +1 -0
keywordmoves/__init__.py
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
1
|
+
"""KeywordMoves public package."""
|
|
2
|
+
|
|
3
|
+
from .models import (
|
|
4
|
+
KeywordCandidate,
|
|
5
|
+
KeywordEvidence,
|
|
6
|
+
LLMRequest,
|
|
7
|
+
LLMResult,
|
|
8
|
+
PluginDescriptor,
|
|
9
|
+
PluginRequest,
|
|
10
|
+
PluginResult,
|
|
11
|
+
)
|
|
12
|
+
from .registry import LLMRegistry, PluginRegistry
|
|
13
|
+
|
|
14
|
+
__all__ = [
|
|
15
|
+
"KeywordCandidate",
|
|
16
|
+
"KeywordEvidence",
|
|
17
|
+
"LLMRegistry",
|
|
18
|
+
"LLMRequest",
|
|
19
|
+
"LLMResult",
|
|
20
|
+
"PluginDescriptor",
|
|
21
|
+
"PluginRegistry",
|
|
22
|
+
"PluginRequest",
|
|
23
|
+
"PluginResult",
|
|
24
|
+
]
|
|
25
|
+
|
|
26
|
+
__version__ = "0.4.2"
|
keywordmoves/__main__.py
ADDED
|
@@ -0,0 +1,252 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import csv
|
|
4
|
+
import hashlib
|
|
5
|
+
import io
|
|
6
|
+
import math
|
|
7
|
+
import re
|
|
8
|
+
from collections import defaultdict
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
from statistics import mean
|
|
11
|
+
|
|
12
|
+
from ..errors import ConfigurationError, InputError
|
|
13
|
+
from ..models import (
|
|
14
|
+
ExecutionContext,
|
|
15
|
+
KeywordCandidate,
|
|
16
|
+
KeywordEvidence,
|
|
17
|
+
PluginDescriptor,
|
|
18
|
+
PluginRequest,
|
|
19
|
+
PluginResult,
|
|
20
|
+
)
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def _rows(path: Path) -> tuple[list[list[str]], str]:
|
|
24
|
+
if not path.is_file():
|
|
25
|
+
raise InputError(f"Google Trends export does not exist: {path}.")
|
|
26
|
+
try:
|
|
27
|
+
if path.stat().st_size > 5_000_000:
|
|
28
|
+
raise InputError("Google Trends CSV exceeds the 5 MB import bound.")
|
|
29
|
+
raw = path.read_bytes()
|
|
30
|
+
if len(raw) > 5_000_000:
|
|
31
|
+
raise InputError("Google Trends CSV exceeds the 5 MB import bound.")
|
|
32
|
+
reader = csv.reader(io.StringIO(raw.decode("utf-8-sig"), newline=""), strict=True)
|
|
33
|
+
return [[cell.strip() for cell in row] for row in reader], hashlib.sha256(raw).hexdigest()
|
|
34
|
+
except (UnicodeDecodeError, csv.Error, OSError) as exc:
|
|
35
|
+
raise InputError(f"Google Trends export must be readable, valid UTF-8 CSV: {path}.") from exc
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def _number(value: str) -> float | None:
|
|
39
|
+
cleaned = value.strip()
|
|
40
|
+
if cleaned not in {"", "-", "<1"} and not re.fullmatch(
|
|
41
|
+
r"[+-]?(?:\d+|\d{1,3}(?:,\d{3})+)(?:\.\d+)?%?", cleaned):
|
|
42
|
+
return None
|
|
43
|
+
cleaned = cleaned.replace(",", "")
|
|
44
|
+
if cleaned in {"", "-", "<1"}:
|
|
45
|
+
return None
|
|
46
|
+
cleaned = cleaned.rstrip("%")
|
|
47
|
+
try:
|
|
48
|
+
value = float(cleaned)
|
|
49
|
+
return value if math.isfinite(value) else None
|
|
50
|
+
except ValueError:
|
|
51
|
+
return None
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _header_index(rows: list[list[str]], first_cells: set[str]) -> int:
|
|
55
|
+
for index, row in enumerate(rows):
|
|
56
|
+
if row and row[0].strip().casefold() in first_cells:
|
|
57
|
+
return index
|
|
58
|
+
raise InputError("The CSV does not contain a recognised Google Trends export header.")
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def _clean_series_name(value: str) -> str:
|
|
62
|
+
return re.sub(r"\s*:\s*\([^)]*\)\s*$", "", value).strip()
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
class GoogleTrendsPlugin:
|
|
66
|
+
descriptor = PluginDescriptor(
|
|
67
|
+
name="google-trends",
|
|
68
|
+
summary="Normalise Google Trends CSV exports into keyword evidence.",
|
|
69
|
+
capabilities=("discover", "analyse", "relative-popularity"),
|
|
70
|
+
operations=("import-interest", "import-related"),
|
|
71
|
+
)
|
|
72
|
+
|
|
73
|
+
def run(self, request: PluginRequest, context: ExecutionContext) -> PluginResult:
|
|
74
|
+
del context
|
|
75
|
+
if request.operation not in self.descriptor.operations:
|
|
76
|
+
raise ConfigurationError(
|
|
77
|
+
f"google-trends operation must be one of: {', '.join(self.descriptor.operations)}."
|
|
78
|
+
)
|
|
79
|
+
if len(request.inputs) != 1:
|
|
80
|
+
raise ConfigurationError("Google Trends import expects exactly one CSV input.")
|
|
81
|
+
path = Path(request.inputs[0])
|
|
82
|
+
rows, raw_sha = _rows(path)
|
|
83
|
+
options = {**request.options, "source_sha256": raw_sha}
|
|
84
|
+
geography = request.options.get("geography")
|
|
85
|
+
observed_at = request.options.get("observed_at")
|
|
86
|
+
if request.operation == "import-interest":
|
|
87
|
+
return self._interest(rows, path, geography, observed_at, options)
|
|
88
|
+
return self._related(rows, path, geography, observed_at, options)
|
|
89
|
+
|
|
90
|
+
def _interest(
|
|
91
|
+
self, rows: list[list[str]], path: Path, geography: str | None,
|
|
92
|
+
observed_at: str | None, options: dict,
|
|
93
|
+
) -> PluginResult:
|
|
94
|
+
index = _header_index(rows, {"week", "day", "month", "date", "time", "region", "subregion"})
|
|
95
|
+
header = rows[index]
|
|
96
|
+
if len(header) < 2:
|
|
97
|
+
raise InputError("The Google Trends interest export has no keyword series.")
|
|
98
|
+
series: dict[str, list[dict]] = defaultdict(list)
|
|
99
|
+
partial = False
|
|
100
|
+
for row in rows[index + 1:]:
|
|
101
|
+
if not row or not any(row):
|
|
102
|
+
continue
|
|
103
|
+
if len(row) != len(header) or not row[0]:
|
|
104
|
+
raise InputError("Malformed or truncated Trends interest row.")
|
|
105
|
+
for column, raw in enumerate(row[1:], start=1):
|
|
106
|
+
if header[column].casefold() == "ispartial":
|
|
107
|
+
if raw.casefold() not in {"true", "false"}:
|
|
108
|
+
raise InputError("IsPartial must be true or false.")
|
|
109
|
+
partial = partial or raw.casefold() == "true"
|
|
110
|
+
continue
|
|
111
|
+
censored = raw == "<1"
|
|
112
|
+
missing = raw in {"", "-"}
|
|
113
|
+
value = _number(raw)
|
|
114
|
+
if not (censored or missing) and (
|
|
115
|
+
value is None or not 0 <= value <= 100 or "%" in raw):
|
|
116
|
+
raise InputError("Interest values must be finite 0..100 indices or explicit missing values.")
|
|
117
|
+
series[_clean_series_name(header[column])].append({
|
|
118
|
+
"period": row[0], "raw_value": raw, "value": value, "censored": censored})
|
|
119
|
+
candidates = []
|
|
120
|
+
for phrase, points in series.items():
|
|
121
|
+
if not phrase or not points:
|
|
122
|
+
continue
|
|
123
|
+
values = [point["value"] for point in points if point["value"] is not None]
|
|
124
|
+
complete = len(values) == len(points) and not partial
|
|
125
|
+
average = round(mean(values), 4) if complete else None
|
|
126
|
+
peak = max(values) if complete else None
|
|
127
|
+
evidence = (
|
|
128
|
+
KeywordEvidence("Google Trends CSV export", "relative_interest_mean", average,
|
|
129
|
+
"index_0_100", str(observed_at) if observed_at else None,
|
|
130
|
+
str(geography) if geography else None,
|
|
131
|
+
"Relative sampled index; a low-volume zero is not proof of no searches."),
|
|
132
|
+
KeywordEvidence("Google Trends CSV export", "relative_interest_peak", peak,
|
|
133
|
+
"index_0_100", str(observed_at) if observed_at else None,
|
|
134
|
+
str(geography) if geography else None),
|
|
135
|
+
)
|
|
136
|
+
candidates.append(KeywordCandidate(
|
|
137
|
+
phrase=phrase.casefold(), relationship="trend-series",
|
|
138
|
+
score=round(average / 100.0, 4) if average is not None else None,
|
|
139
|
+
evidence=evidence, metadata={
|
|
140
|
+
"observations": len(points), "numeric_observations": len(values),
|
|
141
|
+
"series": points, "censored": any(point["censored"] for point in points),
|
|
142
|
+
"completeness": "complete" if complete else "partial",
|
|
143
|
+
"availability": "observed" if values else "no-data",
|
|
144
|
+
"platform": "YouTube" if options.get("search_property") == "youtube" else "Google",
|
|
145
|
+
"evidence_kind": "relative_interest",
|
|
146
|
+
**{key: options[key] for key in ("window", "scope", "category", "search_property",
|
|
147
|
+
"normalization_id") if key in options},
|
|
148
|
+
}))
|
|
149
|
+
if not candidates:
|
|
150
|
+
raise InputError("No interest series were found; unavailable data are not zero demand.")
|
|
151
|
+
return PluginResult(
|
|
152
|
+
self.descriptor.name, "import-interest", tuple(candidates),
|
|
153
|
+
("Google Trends values are relative indices, not search-volume counts.",
|
|
154
|
+
"Censored, missing or partial series do not receive exact summary values."),
|
|
155
|
+
{"source_file": str(path.resolve()), "source_sha256": options["source_sha256"],
|
|
156
|
+
"observed_at": observed_at, "geography": geography,
|
|
157
|
+
**{key: options[key] for key in ("window", "scope", "category", "search_property",
|
|
158
|
+
"normalization_id") if key in options}},
|
|
159
|
+
)
|
|
160
|
+
|
|
161
|
+
def _related(
|
|
162
|
+
self, rows: list[list[str]], path: Path, geography: str | None,
|
|
163
|
+
observed_at: str | None, options: dict,
|
|
164
|
+
) -> PluginResult:
|
|
165
|
+
section = options.get("section")
|
|
166
|
+
if section is not None and section not in {"top", "rising"}:
|
|
167
|
+
raise ConfigurationError("section must be top or rising.")
|
|
168
|
+
candidates: list[KeywordCandidate] = []
|
|
169
|
+
seen: dict[tuple[str, str], str] = {}
|
|
170
|
+
headers, duplicate_rows = [], 0
|
|
171
|
+
for number, row in enumerate(rows, start=1):
|
|
172
|
+
if not row or not any(row):
|
|
173
|
+
continue
|
|
174
|
+
first = row[0].strip()
|
|
175
|
+
if first.casefold() in {"top", "rising"} and not any(row[1:]):
|
|
176
|
+
section = first.casefold()
|
|
177
|
+
continue
|
|
178
|
+
if first.casefold() in {"related queries", "related topics"} and not any(row[1:]):
|
|
179
|
+
headers.append(row)
|
|
180
|
+
continue
|
|
181
|
+
if first.casefold() in {"query", "queries", "topic", "topics"} and (
|
|
182
|
+
len(row) < 2 or row[1].casefold() in {"value", "score", "interest"}):
|
|
183
|
+
continue
|
|
184
|
+
if section is None:
|
|
185
|
+
if len(row) >= 2 and (_number(row[1]) is not None
|
|
186
|
+
or row[1].casefold() == "breakout"):
|
|
187
|
+
raise InputError("Related-query values need a Top/Rising section or section option.")
|
|
188
|
+
headers.append(row)
|
|
189
|
+
continue
|
|
190
|
+
if len(row) != 2 or not first:
|
|
191
|
+
raise InputError(f"Malformed related-query data at CSV row {number}.")
|
|
192
|
+
raw = row[1].strip()
|
|
193
|
+
breakout = raw.casefold() == "breakout"
|
|
194
|
+
unavailable = raw in {"", "-"}
|
|
195
|
+
censored = raw == "<1"
|
|
196
|
+
value = _number(raw)
|
|
197
|
+
if section == "top":
|
|
198
|
+
if breakout or "%" in raw or (value is not None and not 0 <= value <= 100):
|
|
199
|
+
raise InputError("Top query values must be 0..100 relative indices.")
|
|
200
|
+
unit, metric = "index_0_100", "related_top"
|
|
201
|
+
elif breakout:
|
|
202
|
+
unit, metric = "growth_label", "related_rising_breakout"
|
|
203
|
+
value = "Breakout"
|
|
204
|
+
else:
|
|
205
|
+
unit, metric = "percent_growth", "related_rising"
|
|
206
|
+
if censored:
|
|
207
|
+
raise InputError("A censored Top index cannot be used as Rising growth.")
|
|
208
|
+
if not (unavailable or censored or breakout) and value is None:
|
|
209
|
+
raise InputError(f"Invalid or non-finite Trends value at CSV row {number}.")
|
|
210
|
+
key = (section, first.casefold())
|
|
211
|
+
signature = str(value) + ":" + raw
|
|
212
|
+
if key in seen:
|
|
213
|
+
if seen[key] != signature:
|
|
214
|
+
raise InputError("Conflicting duplicate query values in the same Trends section.")
|
|
215
|
+
duplicate_rows += 1
|
|
216
|
+
continue
|
|
217
|
+
seen[key] = signature
|
|
218
|
+
evidence = KeywordEvidence(
|
|
219
|
+
source="Google Trends CSV export", metric=metric, value=value, unit=unit,
|
|
220
|
+
observed_at=str(observed_at) if observed_at else None,
|
|
221
|
+
geography=str(geography) if geography else None,
|
|
222
|
+
notes=("Breakout means growth greater than 5000%; no absolute volume."
|
|
223
|
+
if breakout else "Source-relative Top index or Rising growth, not search volume."),
|
|
224
|
+
)
|
|
225
|
+
context = {key: options[key] for key in (
|
|
226
|
+
"seed_keyword", "window", "scope", "category", "search_property", "normalization_id")
|
|
227
|
+
if key in options}
|
|
228
|
+
candidates.append(KeywordCandidate(
|
|
229
|
+
phrase=first.casefold(), relationship=f"trends-{section}", score=None,
|
|
230
|
+
evidence=(evidence,), metadata={
|
|
231
|
+
**context, "section": section, "raw_value": raw, "csv_row": number,
|
|
232
|
+
"platform": "YouTube" if options.get("search_property") == "youtube" else "Google",
|
|
233
|
+
"evidence_kind": "relative_interest",
|
|
234
|
+
"availability": "no-data" if unavailable else "observed",
|
|
235
|
+
"censored": censored, "approximate": False, "completeness": "complete",
|
|
236
|
+
"breakout_lower_bound_percent": 5000 if breakout else None,
|
|
237
|
+
},
|
|
238
|
+
))
|
|
239
|
+
if not candidates:
|
|
240
|
+
raise InputError("No related queries were found; this is unavailable evidence, not zero demand.")
|
|
241
|
+
sha = options["source_sha256"]
|
|
242
|
+
return PluginResult(
|
|
243
|
+
plugin=self.descriptor.name, operation="import-related", keywords=tuple(candidates),
|
|
244
|
+
notes=("Top, Rising and Breakout retain separate units; no cross-platform demand score.",
|
|
245
|
+
"Different Trends exports may use different normalisation and cannot be compared blindly."),
|
|
246
|
+
metadata={"source_file": str(path.resolve()), "source_sha256": sha,
|
|
247
|
+
"export_headers": headers, "duplicate_rows": duplicate_rows,
|
|
248
|
+
"geography": geography, "observed_at": observed_at,
|
|
249
|
+
**{key: options[key] for key in (
|
|
250
|
+
"seed_keyword", "window", "scope", "category", "search_property",
|
|
251
|
+
"normalization_id") if key in options}},
|
|
252
|
+
)
|
|
@@ -0,0 +1,128 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
from ..errors import ConfigurationError
|
|
6
|
+
from ..models import LLMRequest, LLMResult, PluginDescriptor
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def _as_bool(value: Any, default: bool = False) -> bool:
|
|
10
|
+
if value is None:
|
|
11
|
+
return default
|
|
12
|
+
if isinstance(value, bool):
|
|
13
|
+
return value
|
|
14
|
+
return str(value).strip().lower() in {"1", "true", "yes", "on"}
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class HuggingFaceTransformersLLM:
|
|
18
|
+
"""Local text2text generation through Hugging Face Transformers and PyTorch."""
|
|
19
|
+
|
|
20
|
+
descriptor = PluginDescriptor(
|
|
21
|
+
name="huggingface-transformers",
|
|
22
|
+
summary="Run a local Hugging Face text2text model through PyTorch.",
|
|
23
|
+
capabilities=("generate", "local-inference"),
|
|
24
|
+
operations=("generate",),
|
|
25
|
+
)
|
|
26
|
+
default_model = "Qwen/Qwen2.5-0.5B-Instruct"
|
|
27
|
+
default_revision = "2b01de6d1108f9b2b5e46a726aa678a359b6c03b"
|
|
28
|
+
|
|
29
|
+
def __init__(self) -> None:
|
|
30
|
+
self._loaded: dict[
|
|
31
|
+
tuple[str, str | None, str | None], tuple[Any, Any, Any, str, bool]
|
|
32
|
+
] = {}
|
|
33
|
+
|
|
34
|
+
def _load(self, request: LLMRequest) -> tuple[Any, Any, Any, str, bool]:
|
|
35
|
+
try:
|
|
36
|
+
import torch
|
|
37
|
+
from transformers import (
|
|
38
|
+
AutoConfig,
|
|
39
|
+
AutoModelForCausalLM,
|
|
40
|
+
AutoModelForSeq2SeqLM,
|
|
41
|
+
AutoTokenizer,
|
|
42
|
+
)
|
|
43
|
+
except ImportError as exc:
|
|
44
|
+
raise ConfigurationError(
|
|
45
|
+
"The Hugging Face LLM plugin needs the optional dependencies. "
|
|
46
|
+
"Install KeywordMoves with: pip install 'keywordmoves[huggingface]'."
|
|
47
|
+
) from exc
|
|
48
|
+
|
|
49
|
+
model_name = str(request.options.get("model", self.default_model))
|
|
50
|
+
revision = request.options.get("revision")
|
|
51
|
+
if revision is None and model_name == self.default_model:
|
|
52
|
+
revision = self.default_revision
|
|
53
|
+
cache_dir = request.options.get("cache_dir")
|
|
54
|
+
key = (model_name, str(revision) if revision else None, str(cache_dir) if cache_dir else None)
|
|
55
|
+
if key in self._loaded:
|
|
56
|
+
return self._loaded[key]
|
|
57
|
+
|
|
58
|
+
local_only = _as_bool(request.options.get("local_files_only"), False)
|
|
59
|
+
common = {
|
|
60
|
+
"revision": revision,
|
|
61
|
+
"cache_dir": cache_dir,
|
|
62
|
+
"local_files_only": local_only,
|
|
63
|
+
"trust_remote_code": False,
|
|
64
|
+
}
|
|
65
|
+
common = {name: value for name, value in common.items() if value is not None}
|
|
66
|
+
config = AutoConfig.from_pretrained(model_name, **common)
|
|
67
|
+
tokenizer = AutoTokenizer.from_pretrained(model_name, **common)
|
|
68
|
+
model_class = AutoModelForSeq2SeqLM if config.is_encoder_decoder else AutoModelForCausalLM
|
|
69
|
+
model = model_class.from_pretrained(model_name, use_safetensors=True, **common)
|
|
70
|
+
|
|
71
|
+
requested_device = str(request.options.get("device", "auto"))
|
|
72
|
+
if requested_device == "auto":
|
|
73
|
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
74
|
+
elif requested_device in {"cpu", "cuda"}:
|
|
75
|
+
device = requested_device
|
|
76
|
+
else:
|
|
77
|
+
raise ConfigurationError("Hugging Face device must be 'auto', 'cpu', or 'cuda'.")
|
|
78
|
+
if device == "cuda" and not torch.cuda.is_available():
|
|
79
|
+
raise ConfigurationError("CUDA was selected, but PyTorch reports that CUDA is unavailable.")
|
|
80
|
+
model.to(device)
|
|
81
|
+
model.eval()
|
|
82
|
+
loaded = (torch, tokenizer, model, device, bool(config.is_encoder_decoder))
|
|
83
|
+
self._loaded[key] = loaded
|
|
84
|
+
return loaded
|
|
85
|
+
|
|
86
|
+
def generate(self, request: LLMRequest) -> LLMResult:
|
|
87
|
+
torch, tokenizer, model, device, is_encoder_decoder = self._load(request)
|
|
88
|
+
model_name = str(request.options.get("model", self.default_model))
|
|
89
|
+
max_input_tokens = int(request.options.get("max_input_tokens", 512))
|
|
90
|
+
prompt = request.prompt
|
|
91
|
+
if not is_encoder_decoder and getattr(tokenizer, "chat_template", None):
|
|
92
|
+
prompt = tokenizer.apply_chat_template(
|
|
93
|
+
[{"role": "user", "content": request.prompt}],
|
|
94
|
+
tokenize=False,
|
|
95
|
+
add_generation_prompt=True,
|
|
96
|
+
)
|
|
97
|
+
encoded = tokenizer(
|
|
98
|
+
prompt,
|
|
99
|
+
return_tensors="pt",
|
|
100
|
+
truncation=True,
|
|
101
|
+
max_length=max_input_tokens,
|
|
102
|
+
)
|
|
103
|
+
encoded = {key: value.to(device) for key, value in encoded.items()}
|
|
104
|
+
with torch.inference_mode():
|
|
105
|
+
output = model.generate(
|
|
106
|
+
**encoded,
|
|
107
|
+
max_new_tokens=request.max_new_tokens,
|
|
108
|
+
do_sample=False,
|
|
109
|
+
num_beams=int(request.options.get("num_beams", 1)),
|
|
110
|
+
)
|
|
111
|
+
generated = output[0]
|
|
112
|
+
if not is_encoder_decoder:
|
|
113
|
+
generated = generated[encoded["input_ids"].shape[-1] :]
|
|
114
|
+
text = tokenizer.decode(generated, skip_special_tokens=True).strip()
|
|
115
|
+
revision = request.options.get("revision")
|
|
116
|
+
if revision is None and model_name == self.default_model:
|
|
117
|
+
revision = self.default_revision
|
|
118
|
+
return LLMResult(
|
|
119
|
+
plugin=self.descriptor.name,
|
|
120
|
+
model=model_name,
|
|
121
|
+
text=text,
|
|
122
|
+
metadata={
|
|
123
|
+
"device": device,
|
|
124
|
+
"task": request.task,
|
|
125
|
+
"revision": revision,
|
|
126
|
+
"architecture": "seq2seq" if is_encoder_decoder else "causal",
|
|
127
|
+
},
|
|
128
|
+
)
|