init
This commit is contained in:
@@ -0,0 +1,289 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Build static vector-search data for the ASMR browser."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import re
|
||||
import struct
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable
|
||||
|
||||
|
||||
DEFAULT_MODEL = "intfloat/multilingual-e5-small"
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Generate web/data assets from asmr_works.jsonl."
|
||||
)
|
||||
parser.add_argument("--input", default="asmr_works.jsonl", help="Input JSONL file.")
|
||||
parser.add_argument("--output-dir", default="web/data", help="Output data directory.")
|
||||
parser.add_argument(
|
||||
"--method",
|
||||
choices=("hashing", "sentence-transformers"),
|
||||
default="hashing",
|
||||
help="Embedding backend. Hashing works without dependencies; sentence-transformers is semantic.",
|
||||
)
|
||||
parser.add_argument("--model", default=DEFAULT_MODEL, help="SentenceTransformer model name.")
|
||||
parser.add_argument("--dimensions", type=int, default=384, help="Hashing vector dimensions.")
|
||||
parser.add_argument("--batch-size", type=int, default=64, help="SentenceTransformer batch size.")
|
||||
parser.add_argument("--limit", type=int, default=None, help="Optional cap for development.")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def safe_list(value: Any) -> list[Any]:
|
||||
return value if isinstance(value, list) else []
|
||||
|
||||
|
||||
def clean_text(value: Any) -> str:
|
||||
if not isinstance(value, str):
|
||||
return ""
|
||||
return re.sub(r"\s+", " ", value).strip()
|
||||
|
||||
|
||||
def unique_strings(values: Iterable[Any]) -> list[str]:
|
||||
seen: set[str] = set()
|
||||
result: list[str] = []
|
||||
for value in values:
|
||||
text = clean_text(value)
|
||||
if text and text not in seen:
|
||||
seen.add(text)
|
||||
result.append(text)
|
||||
return result
|
||||
|
||||
|
||||
def tag_name(tag: Any) -> str:
|
||||
if not isinstance(tag, dict):
|
||||
return ""
|
||||
i18n = tag.get("i18n")
|
||||
if isinstance(i18n, dict):
|
||||
ja = i18n.get("ja-jp")
|
||||
if isinstance(ja, dict):
|
||||
name = clean_text(ja.get("name"))
|
||||
if name:
|
||||
return name
|
||||
return clean_text(tag.get("name"))
|
||||
|
||||
|
||||
def actor_name(actor: Any) -> str:
|
||||
return clean_text(actor.get("name")) if isinstance(actor, dict) else ""
|
||||
|
||||
|
||||
def circle_name(work: dict[str, Any]) -> str:
|
||||
circle = work.get("circle")
|
||||
if isinstance(circle, dict):
|
||||
name = clean_text(circle.get("name"))
|
||||
if name:
|
||||
return name
|
||||
return clean_text(work.get("name"))
|
||||
|
||||
|
||||
def group_id(work: dict[str, Any]) -> str:
|
||||
editions = work.get("language_editions")
|
||||
if isinstance(editions, list):
|
||||
edition_ids = [
|
||||
str(item.get("edition_id"))
|
||||
for item in editions
|
||||
if isinstance(item, dict) and item.get("edition_id")
|
||||
]
|
||||
if edition_ids:
|
||||
return f"edition:{edition_ids[0]}"
|
||||
|
||||
translation_info = work.get("translation_info")
|
||||
if isinstance(translation_info, dict):
|
||||
for key in ("original_workno", "parent_workno"):
|
||||
value = clean_text(translation_info.get(key))
|
||||
if value:
|
||||
return f"work:{value}"
|
||||
|
||||
for key in ("original_workno", "source_id", "id"):
|
||||
value = work.get(key)
|
||||
if value:
|
||||
return f"work:{value}"
|
||||
return "work:unknown"
|
||||
|
||||
|
||||
def normalize_work(work: dict[str, Any], embedding_index: int) -> dict[str, Any]:
|
||||
tags = unique_strings(tag_name(tag) for tag in safe_list(work.get("tags")))
|
||||
vas = unique_strings(actor_name(actor) for actor in safe_list(work.get("vas")))
|
||||
source_id = clean_text(work.get("source_id"))
|
||||
return {
|
||||
"embeddingIndex": embedding_index,
|
||||
"id": work.get("id"),
|
||||
"sourceId": source_id,
|
||||
"title": clean_text(work.get("title")) or source_id or str(work.get("id")),
|
||||
"sourceUrl": clean_text(work.get("source_url")),
|
||||
"thumbnailCoverUrl": clean_text(work.get("thumbnailCoverUrl")),
|
||||
"mainCoverUrl": clean_text(work.get("mainCoverUrl")),
|
||||
"circle": circle_name(work),
|
||||
"vas": vas,
|
||||
"tags": tags,
|
||||
"nsfw": bool(work.get("nsfw")),
|
||||
"ageCategory": clean_text(work.get("age_category_string")),
|
||||
"duration": work.get("duration") if isinstance(work.get("duration"), int) else None,
|
||||
"dlCount": work.get("dl_count") if isinstance(work.get("dl_count"), int) else 0,
|
||||
"rateAverage": work.get("rate_average_2dp") if isinstance(work.get("rate_average_2dp"), (int, float)) else None,
|
||||
"rateCount": work.get("rate_count") if isinstance(work.get("rate_count"), int) else 0,
|
||||
"release": clean_text(work.get("release")),
|
||||
"createDate": clean_text(work.get("create_date")),
|
||||
"groupId": group_id(work),
|
||||
}
|
||||
|
||||
|
||||
def embedding_text(work: dict[str, Any]) -> str:
|
||||
parts = [f"タイトル: {work['title']}"]
|
||||
if work["tags"]:
|
||||
parts.append(f"タグ: {', '.join(work['tags'])}")
|
||||
if work["vas"]:
|
||||
parts.append(f"声優: {', '.join(work['vas'])}")
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def iter_jsonl(path: Path, limit: int | None) -> Iterable[dict[str, Any]]:
|
||||
with path.open(encoding="utf-8") as file:
|
||||
for index, line in enumerate(file):
|
||||
if limit is not None and index >= limit:
|
||||
break
|
||||
if not line.strip():
|
||||
continue
|
||||
value = json.loads(line)
|
||||
if isinstance(value, dict):
|
||||
yield value
|
||||
|
||||
|
||||
def token_features(text: str) -> Iterable[str]:
|
||||
compact = re.sub(r"\s+", "", text.lower())
|
||||
for token in re.findall(r"[a-z0-9_+-]+", text.lower()):
|
||||
if len(token) >= 2:
|
||||
yield f"word:{token}"
|
||||
for size in (2, 3):
|
||||
if len(compact) >= size:
|
||||
for index in range(len(compact) - size + 1):
|
||||
yield f"char{size}:{compact[index:index + size]}"
|
||||
|
||||
|
||||
def hash_feature(feature: str, dimensions: int) -> tuple[int, float]:
|
||||
digest = hashlib.blake2b(feature.encode("utf-8"), digest_size=8).digest()
|
||||
value = int.from_bytes(digest, "little")
|
||||
sign = -1.0 if value & 1 else 1.0
|
||||
return (value >> 1) % dimensions, sign
|
||||
|
||||
|
||||
def hashing_embedding(text: str, dimensions: int) -> list[float]:
|
||||
vector = [0.0] * dimensions
|
||||
for feature in token_features(text):
|
||||
index, sign = hash_feature(feature, dimensions)
|
||||
vector[index] += sign
|
||||
|
||||
norm = math.sqrt(sum(value * value for value in vector))
|
||||
if norm == 0:
|
||||
return vector
|
||||
return [value / norm for value in vector]
|
||||
|
||||
|
||||
def write_hashing_embeddings(path: Path, texts: list[str], dimensions: int) -> None:
|
||||
with path.open("wb") as file:
|
||||
for text in texts:
|
||||
vector = hashing_embedding(text, dimensions)
|
||||
file.write(struct.pack(f"<{dimensions}f", *vector))
|
||||
|
||||
|
||||
def write_sentence_transformer_embeddings(
|
||||
path: Path,
|
||||
texts: list[str],
|
||||
model_name: str,
|
||||
batch_size: int,
|
||||
) -> int:
|
||||
try:
|
||||
import numpy as np
|
||||
from sentence_transformers import SentenceTransformer
|
||||
except ImportError as exc:
|
||||
raise SystemExit(
|
||||
"sentence-transformers mode requires dependencies. "
|
||||
"Run `uv run --extra semantic python scripts/build_vector_data.py --method sentence-transformers`."
|
||||
) from exc
|
||||
|
||||
model = SentenceTransformer(model_name)
|
||||
if "e5" in model_name.lower():
|
||||
texts = [f"passage: {text}" for text in texts]
|
||||
embeddings = model.encode(
|
||||
texts,
|
||||
batch_size=batch_size,
|
||||
normalize_embeddings=True,
|
||||
show_progress_bar=True,
|
||||
)
|
||||
embeddings = np.asarray(embeddings, dtype="<f4")
|
||||
embeddings.tofile(path)
|
||||
return int(embeddings.shape[1])
|
||||
|
||||
|
||||
def main() -> int:
|
||||
args = parse_args()
|
||||
input_path = Path(args.input)
|
||||
output_dir = Path(args.output_dir)
|
||||
|
||||
if args.dimensions <= 0:
|
||||
raise SystemExit("--dimensions must be greater than 0")
|
||||
if args.batch_size <= 0:
|
||||
raise SystemExit("--batch-size must be greater than 0")
|
||||
if not input_path.exists():
|
||||
raise SystemExit(f"input file not found: {input_path}")
|
||||
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
works: list[dict[str, Any]] = []
|
||||
texts: list[str] = []
|
||||
for raw_work in iter_jsonl(input_path, args.limit):
|
||||
work = normalize_work(raw_work, len(works))
|
||||
works.append(work)
|
||||
texts.append(embedding_text(work))
|
||||
|
||||
if not works:
|
||||
raise SystemExit("no works found")
|
||||
|
||||
embeddings_path = output_dir / "embeddings.f32"
|
||||
if args.method == "sentence-transformers":
|
||||
dimensions = write_sentence_transformer_embeddings(
|
||||
embeddings_path,
|
||||
texts,
|
||||
args.model,
|
||||
args.batch_size,
|
||||
)
|
||||
model = args.model
|
||||
else:
|
||||
dimensions = args.dimensions
|
||||
write_hashing_embeddings(embeddings_path, texts, dimensions)
|
||||
model = f"hashing-char-ngram-{dimensions}d"
|
||||
|
||||
manifest = {
|
||||
"count": len(works),
|
||||
"dimensions": dimensions,
|
||||
"embeddingFile": "embeddings.f32",
|
||||
"worksFile": "works.json",
|
||||
"method": args.method,
|
||||
"model": model,
|
||||
"generatedAt": datetime.now(timezone.utc).isoformat(),
|
||||
"score": {
|
||||
"vectorWeight": 0.8,
|
||||
"tagWeight": 0.2,
|
||||
},
|
||||
}
|
||||
|
||||
with (output_dir / "works.json").open("w", encoding="utf-8") as file:
|
||||
json.dump(works, file, ensure_ascii=False, separators=(",", ":"))
|
||||
with (output_dir / "manifest.json").open("w", encoding="utf-8") as file:
|
||||
json.dump(manifest, file, ensure_ascii=False, indent=2)
|
||||
file.write("\n")
|
||||
|
||||
print(f"Wrote {len(works)} works to {output_dir}")
|
||||
print(f"Embedding method: {args.method} ({dimensions} dimensions)")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,357 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Build static vector-search data with an OpenAI-compatible embeddings API."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import http.client
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import struct
|
||||
import time
|
||||
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.error import HTTPError, URLError
|
||||
from urllib.request import Request, urlopen
|
||||
|
||||
from build_vector_data import embedding_text, iter_jsonl, normalize_work
|
||||
|
||||
|
||||
DEFAULT_BASE_URL = "http://192.168.0.35:1234/"
|
||||
DEFAULT_MODEL = "text-embedding-qwen3-embedding-8b"
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Generate web/data assets using a remote OpenAI-compatible embeddings API."
|
||||
)
|
||||
parser.add_argument("--input", default="asmr_works.jsonl", help="Input JSONL file.")
|
||||
parser.add_argument("--output-dir", default="web/data", help="Output data directory.")
|
||||
parser.add_argument("--base-url", default=DEFAULT_BASE_URL, help="API base URL.")
|
||||
parser.add_argument("--model", default=DEFAULT_MODEL, help="Embedding model name.")
|
||||
parser.add_argument(
|
||||
"--api-key-env",
|
||||
default="EMBEDDING_API_KEY",
|
||||
help="Environment variable that contains the API key.",
|
||||
)
|
||||
parser.add_argument("--batch-size", type=int, default=32, help="Texts per API request.")
|
||||
parser.add_argument("--concurrency", type=int, default=10, help="Concurrent API requests.")
|
||||
parser.add_argument("--timeout", type=float, default=120.0, help="Request timeout seconds.")
|
||||
parser.add_argument("--retries", type=int, default=3, help="Retries per failed batch.")
|
||||
parser.add_argument("--retry-wait", type=float, default=2.0, help="Initial retry wait seconds.")
|
||||
parser.add_argument("--max-retry-wait", type=float, default=60.0, help="Maximum retry wait seconds.")
|
||||
parser.add_argument(
|
||||
"--retry-forever",
|
||||
dest="retry_forever",
|
||||
action="store_true",
|
||||
default=True,
|
||||
help="Retry failed batches until they succeed. This is enabled by default.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no-retry-forever",
|
||||
dest="retry_forever",
|
||||
action="store_false",
|
||||
help="Stop after --retries attempts instead of retrying forever.",
|
||||
)
|
||||
parser.add_argument("--limit", type=int, default=None, help="Optional cap for development.")
|
||||
parser.add_argument(
|
||||
"--no-normalize",
|
||||
action="store_true",
|
||||
help="Do not L2-normalize vectors before writing embeddings.f32.",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def embeddings_url(base_url: str) -> str:
|
||||
return f"{base_url.rstrip('/')}/v1/embeddings"
|
||||
|
||||
|
||||
def normalize_vector(vector: list[float]) -> list[float]:
|
||||
norm = math.sqrt(sum(value * value for value in vector))
|
||||
if norm == 0:
|
||||
return vector
|
||||
return [value / norm for value in vector]
|
||||
|
||||
|
||||
def request_embeddings(
|
||||
*,
|
||||
url: str,
|
||||
model: str,
|
||||
inputs: list[str],
|
||||
api_key: str,
|
||||
timeout: float,
|
||||
) -> list[list[float]]:
|
||||
body = json.dumps({"model": model, "input": inputs}).encode("utf-8")
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if api_key:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
|
||||
request = Request(url, data=body, headers=headers, method="POST")
|
||||
try:
|
||||
with urlopen(request, timeout=timeout) as response:
|
||||
payload = json.load(response)
|
||||
except HTTPError as exc:
|
||||
detail = exc.read(1000).decode("utf-8", errors="replace")
|
||||
raise RuntimeError(f"HTTP {exc.code}: {detail}") from exc
|
||||
except (
|
||||
URLError,
|
||||
TimeoutError,
|
||||
json.JSONDecodeError,
|
||||
http.client.IncompleteRead,
|
||||
http.client.RemoteDisconnected,
|
||||
ConnectionResetError,
|
||||
OSError,
|
||||
) as exc:
|
||||
raise RuntimeError(str(exc)) from exc
|
||||
|
||||
data = payload.get("data") if isinstance(payload, dict) else None
|
||||
if not isinstance(data, list):
|
||||
raise RuntimeError("response does not contain data[]")
|
||||
|
||||
ordered: list[list[float] | None] = [None] * len(inputs)
|
||||
for fallback_index, item in enumerate(data):
|
||||
if not isinstance(item, dict):
|
||||
raise RuntimeError("response data item is not an object")
|
||||
index = item.get("index", fallback_index)
|
||||
embedding = item.get("embedding")
|
||||
if not isinstance(index, int) or not 0 <= index < len(inputs):
|
||||
raise RuntimeError(f"invalid embedding index: {index}")
|
||||
if not isinstance(embedding, list) or not embedding:
|
||||
raise RuntimeError(f"missing embedding for index {index}")
|
||||
try:
|
||||
ordered[index] = [float(value) for value in embedding]
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise RuntimeError(f"embedding for index {index} contains non-numeric values") from exc
|
||||
|
||||
missing = [index for index, vector in enumerate(ordered) if vector is None]
|
||||
if missing:
|
||||
raise RuntimeError(f"missing embeddings for indexes: {missing[:10]}")
|
||||
|
||||
return [vector for vector in ordered if vector is not None]
|
||||
|
||||
|
||||
def request_embeddings_with_retries(
|
||||
*,
|
||||
url: str,
|
||||
model: str,
|
||||
inputs: list[str],
|
||||
api_key: str,
|
||||
timeout: float,
|
||||
retries: int,
|
||||
retry_wait: float,
|
||||
max_retry_wait: float,
|
||||
retry_forever: bool,
|
||||
batch_start: int,
|
||||
) -> list[list[float]]:
|
||||
last_error: Exception | None = None
|
||||
attempt = 0
|
||||
while True:
|
||||
try:
|
||||
return request_embeddings(
|
||||
url=url,
|
||||
model=model,
|
||||
inputs=inputs,
|
||||
api_key=api_key,
|
||||
timeout=timeout,
|
||||
)
|
||||
except RuntimeError as exc:
|
||||
last_error = exc
|
||||
if not retry_forever and attempt >= retries:
|
||||
break
|
||||
wait_seconds = min(max_retry_wait, retry_wait * (2 ** min(attempt, 10)))
|
||||
print(
|
||||
f"Batch {batch_start}: request failed ({exc}); retrying in {wait_seconds:.1f}s",
|
||||
flush=True,
|
||||
)
|
||||
time.sleep(wait_seconds)
|
||||
attempt += 1
|
||||
raise RuntimeError(str(last_error))
|
||||
|
||||
|
||||
def batched(values: list[str], size: int):
|
||||
for index in range(0, len(values), size):
|
||||
yield index, values[index : index + size]
|
||||
|
||||
|
||||
def write_remote_embeddings(
|
||||
*,
|
||||
path: Path,
|
||||
texts: list[str],
|
||||
url: str,
|
||||
model: str,
|
||||
api_key: str,
|
||||
batch_size: int,
|
||||
timeout: float,
|
||||
retries: int,
|
||||
retry_wait: float,
|
||||
max_retry_wait: float,
|
||||
retry_forever: bool,
|
||||
concurrency: int,
|
||||
should_normalize: bool,
|
||||
) -> int:
|
||||
dimensions: int | None = None
|
||||
temp_path = path.with_suffix(path.suffix + ".tmp")
|
||||
|
||||
try:
|
||||
with temp_path.open("wb") as file:
|
||||
pending = {}
|
||||
completed: dict[int, list[list[float]]] = {}
|
||||
next_submit = 0
|
||||
next_write = 0
|
||||
|
||||
def submit_available(executor: ThreadPoolExecutor) -> None:
|
||||
nonlocal next_submit
|
||||
while next_submit < len(texts) and len(pending) + len(completed) < concurrency:
|
||||
start = next_submit
|
||||
batch = texts[start : start + batch_size]
|
||||
future = executor.submit(
|
||||
request_embeddings_with_retries,
|
||||
url=url,
|
||||
model=model,
|
||||
inputs=batch,
|
||||
api_key=api_key,
|
||||
timeout=timeout,
|
||||
retries=retries,
|
||||
retry_wait=retry_wait,
|
||||
max_retry_wait=max_retry_wait,
|
||||
retry_forever=retry_forever,
|
||||
batch_start=start,
|
||||
)
|
||||
pending[future] = start
|
||||
next_submit += len(batch)
|
||||
|
||||
def write_vectors(vectors: list[list[float]]) -> None:
|
||||
nonlocal dimensions
|
||||
for vector in vectors:
|
||||
if should_normalize:
|
||||
vector = normalize_vector(vector)
|
||||
if dimensions is None:
|
||||
dimensions = len(vector)
|
||||
elif len(vector) != dimensions:
|
||||
raise RuntimeError(
|
||||
f"embedding dimensions changed: expected {dimensions}, got {len(vector)}"
|
||||
)
|
||||
file.write(struct.pack(f"<{dimensions}f", *vector))
|
||||
|
||||
with ThreadPoolExecutor(max_workers=concurrency) as executor:
|
||||
submit_available(executor)
|
||||
while pending:
|
||||
done, _ = wait(pending, return_when=FIRST_COMPLETED)
|
||||
for future in done:
|
||||
start = pending.pop(future)
|
||||
vectors = future.result()
|
||||
expected = min(batch_size, len(texts) - start)
|
||||
if len(vectors) != expected:
|
||||
raise RuntimeError(
|
||||
f"batch at {start} returned {len(vectors)} embeddings; expected {expected}"
|
||||
)
|
||||
completed[start] = vectors
|
||||
|
||||
while next_write in completed:
|
||||
vectors = completed.pop(next_write)
|
||||
write_vectors(vectors)
|
||||
next_write += len(vectors)
|
||||
print(f"Embedded {next_write}/{len(texts)} works", flush=True)
|
||||
|
||||
submit_available(executor)
|
||||
except Exception:
|
||||
temp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
|
||||
if dimensions is None:
|
||||
raise RuntimeError("no embeddings were written")
|
||||
|
||||
temp_path.replace(path)
|
||||
return dimensions
|
||||
|
||||
|
||||
def main() -> int:
|
||||
args = parse_args()
|
||||
input_path = Path(args.input)
|
||||
output_dir = Path(args.output_dir)
|
||||
|
||||
if args.batch_size <= 0:
|
||||
raise SystemExit("--batch-size must be greater than 0")
|
||||
if args.concurrency <= 0:
|
||||
raise SystemExit("--concurrency must be greater than 0")
|
||||
if args.timeout <= 0:
|
||||
raise SystemExit("--timeout must be greater than 0")
|
||||
if args.retries < 0:
|
||||
raise SystemExit("--retries must be greater than or equal to 0")
|
||||
if args.retry_wait < 0:
|
||||
raise SystemExit("--retry-wait must be greater than or equal to 0")
|
||||
if args.max_retry_wait < 0:
|
||||
raise SystemExit("--max-retry-wait must be greater than or equal to 0")
|
||||
if not input_path.exists():
|
||||
raise SystemExit(f"input file not found: {input_path}")
|
||||
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
works: list[dict[str, Any]] = []
|
||||
texts: list[str] = []
|
||||
for raw_work in iter_jsonl(input_path, args.limit):
|
||||
work = normalize_work(raw_work, len(works))
|
||||
works.append(work)
|
||||
texts.append(embedding_text(work))
|
||||
|
||||
if not works:
|
||||
raise SystemExit("no works found")
|
||||
|
||||
api_key = os.environ.get(args.api_key_env, "")
|
||||
url = embeddings_url(args.base_url)
|
||||
embeddings_path = output_dir / "embeddings.f32"
|
||||
try:
|
||||
dimensions = write_remote_embeddings(
|
||||
path=embeddings_path,
|
||||
texts=texts,
|
||||
url=url,
|
||||
model=args.model,
|
||||
api_key=api_key,
|
||||
batch_size=args.batch_size,
|
||||
timeout=args.timeout,
|
||||
retries=args.retries,
|
||||
retry_wait=args.retry_wait,
|
||||
max_retry_wait=args.max_retry_wait,
|
||||
retry_forever=args.retry_forever,
|
||||
concurrency=args.concurrency,
|
||||
should_normalize=not args.no_normalize,
|
||||
)
|
||||
except RuntimeError as exc:
|
||||
raise SystemExit(f"embedding request failed: {exc}") from exc
|
||||
|
||||
manifest = {
|
||||
"count": len(works),
|
||||
"dimensions": dimensions,
|
||||
"embeddingFile": "embeddings.f32",
|
||||
"worksFile": "works.json",
|
||||
"method": "remote-openai-compatible",
|
||||
"model": args.model,
|
||||
"baseUrl": args.base_url.rstrip("/"),
|
||||
"generatedAt": datetime.now(timezone.utc).isoformat(),
|
||||
"normalized": not args.no_normalize,
|
||||
"batchSize": args.batch_size,
|
||||
"concurrency": args.concurrency,
|
||||
"retryForever": args.retry_forever,
|
||||
"maxRetryWait": args.max_retry_wait,
|
||||
"score": {
|
||||
"vectorWeight": 0.8,
|
||||
"tagWeight": 0.2,
|
||||
},
|
||||
}
|
||||
|
||||
with (output_dir / "works.json").open("w", encoding="utf-8") as file:
|
||||
json.dump(works, file, ensure_ascii=False, separators=(",", ":"))
|
||||
with (output_dir / "manifest.json").open("w", encoding="utf-8") as file:
|
||||
json.dump(manifest, file, ensure_ascii=False, indent=2)
|
||||
file.write("\n")
|
||||
|
||||
print(f"Wrote {len(works)} works to {output_dir}")
|
||||
print(f"Embedding method: remote-openai-compatible ({dimensions} dimensions)")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Executable
+382
@@ -0,0 +1,382 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Filter ASMR work JSONL data for training datasets."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import shutil
|
||||
import sys
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
SIMPLIFIED_CHINESE_MARKER = "【简体中文版】"
|
||||
DEFAULT_EXCLUDE_TAGS = ("女性向", "乙女向")
|
||||
DEFAULT_EXCLUDE_TAG_IDS = (491, 10001)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Filter translated and female-oriented works from asmr_works.jsonl."
|
||||
)
|
||||
parser.add_argument("--input", default="asmr_works.jsonl", help="Input JSONL file.")
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default=None,
|
||||
help="Output JSONL file. Required unless --in-place or --dry-run is used.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--in-place",
|
||||
action="store_true",
|
||||
help="Overwrite --input after creating a backup.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--backup-suffix",
|
||||
default=".bak",
|
||||
help="Backup suffix used with --in-place.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dry-run",
|
||||
action="store_true",
|
||||
help="Show filtering stats without writing output.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--exclude-tag",
|
||||
action="append",
|
||||
default=[],
|
||||
help="Additional tag name to remove. Can be specified multiple times.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--exclude-tag-id",
|
||||
action="append",
|
||||
type=int,
|
||||
default=[],
|
||||
help="Additional tag id to remove. Can be specified multiple times.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--keep-default-exclude-tags",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=True,
|
||||
help="Keep default 女性向/乙女向 filtering.",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def clean_text(value: Any) -> str:
|
||||
return value.strip() if isinstance(value, str) else ""
|
||||
|
||||
|
||||
def read_jsonl(path: Path) -> list[dict[str, Any]]:
|
||||
works: list[dict[str, Any]] = []
|
||||
with path.open(encoding="utf-8") as file:
|
||||
for line_no, line in enumerate(file, 1):
|
||||
if not line.strip():
|
||||
continue
|
||||
try:
|
||||
value = json.loads(line)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise SystemExit(f"{path}:{line_no}: invalid JSON: {exc}") from exc
|
||||
if not isinstance(value, dict):
|
||||
raise SystemExit(f"{path}:{line_no}: expected JSON object")
|
||||
value["__line_no"] = line_no
|
||||
works.append(value)
|
||||
return works
|
||||
|
||||
|
||||
def write_jsonl(path: Path, works: list[dict[str, Any]]) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with path.open("w", encoding="utf-8") as file:
|
||||
for work in works:
|
||||
clean_work = {key: value for key, value in work.items() if key != "__line_no"}
|
||||
file.write(json.dumps(clean_work, ensure_ascii=False, separators=(",", ":")))
|
||||
file.write("\n")
|
||||
|
||||
|
||||
def original_group_workno(work: dict[str, Any]) -> str:
|
||||
translation_info = work.get("translation_info")
|
||||
if isinstance(translation_info, dict):
|
||||
value = clean_text(translation_info.get("original_workno"))
|
||||
if value:
|
||||
return value
|
||||
|
||||
value = clean_text(work.get("original_workno"))
|
||||
if value:
|
||||
return value
|
||||
|
||||
editions = work.get("language_editions")
|
||||
if isinstance(editions, list):
|
||||
for item in editions:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if item.get("lang") == "JPN":
|
||||
workno = clean_text(item.get("workno"))
|
||||
if workno:
|
||||
return workno
|
||||
|
||||
other_editions = work.get("other_language_editions_in_db")
|
||||
if isinstance(other_editions, list):
|
||||
for item in other_editions:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if item.get("is_original") is True:
|
||||
workno = clean_text(item.get("source_id"))
|
||||
if workno:
|
||||
return workno
|
||||
|
||||
if isinstance(translation_info, dict) and translation_info.get("is_original") is True:
|
||||
return clean_text(work.get("source_id"))
|
||||
return ""
|
||||
|
||||
|
||||
def translation_group_key(work: dict[str, Any]) -> str:
|
||||
original_workno = original_group_workno(work)
|
||||
if original_workno:
|
||||
return f"work:{original_workno}"
|
||||
|
||||
editions = work.get("language_editions")
|
||||
if isinstance(editions, list):
|
||||
edition_ids = [
|
||||
str(item.get("edition_id"))
|
||||
for item in editions
|
||||
if isinstance(item, dict) and item.get("edition_id")
|
||||
]
|
||||
if edition_ids:
|
||||
return f"edition:{edition_ids[0]}"
|
||||
|
||||
translation_info = work.get("translation_info")
|
||||
if isinstance(translation_info, dict):
|
||||
for key in ("original_workno", "parent_workno"):
|
||||
value = clean_text(translation_info.get(key))
|
||||
if value:
|
||||
return f"work:{value}"
|
||||
|
||||
for key in ("original_workno", "source_id", "id"):
|
||||
value = work.get(key)
|
||||
if value:
|
||||
return f"work:{value}"
|
||||
return "work:unknown"
|
||||
|
||||
|
||||
def work_language(work: dict[str, Any]) -> str:
|
||||
translation_info = work.get("translation_info")
|
||||
if isinstance(translation_info, dict):
|
||||
lang = clean_text(translation_info.get("lang"))
|
||||
if lang:
|
||||
return lang
|
||||
|
||||
source_id = clean_text(work.get("source_id"))
|
||||
editions = work.get("language_editions")
|
||||
if isinstance(editions, list) and source_id:
|
||||
for item in editions:
|
||||
if isinstance(item, dict) and clean_text(item.get("workno")) == source_id:
|
||||
return clean_text(item.get("lang")) or clean_text(item.get("label"))
|
||||
|
||||
if isinstance(translation_info, dict) and translation_info.get("is_original") is True:
|
||||
return "JPN"
|
||||
return ""
|
||||
|
||||
|
||||
def jpn_worknos_from_editions(work: dict[str, Any]) -> set[str]:
|
||||
result: set[str] = set()
|
||||
editions = work.get("language_editions")
|
||||
if not isinstance(editions, list):
|
||||
return result
|
||||
|
||||
for item in editions:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if item.get("lang") == "JPN":
|
||||
workno = clean_text(item.get("workno"))
|
||||
if workno:
|
||||
result.add(workno)
|
||||
return result
|
||||
|
||||
|
||||
def translated_line_numbers_to_remove(works: list[dict[str, Any]]) -> set[int]:
|
||||
groups: dict[str, list[dict[str, Any]]] = defaultdict(list)
|
||||
for work in works:
|
||||
groups[translation_group_key(work)].append(work)
|
||||
|
||||
remove: set[int] = set()
|
||||
for group in groups.values():
|
||||
source_ids = {clean_text(work.get("source_id")) for work in group}
|
||||
jpn_worknos: set[str] = set()
|
||||
for work in group:
|
||||
original_workno = original_group_workno(work)
|
||||
if original_workno:
|
||||
jpn_worknos.add(original_workno)
|
||||
jpn_worknos.update(jpn_worknos_from_editions(work))
|
||||
|
||||
jpn_worknos_in_dataset = jpn_worknos & source_ids
|
||||
if not jpn_worknos_in_dataset:
|
||||
continue
|
||||
|
||||
for work in group:
|
||||
source_id = clean_text(work.get("source_id"))
|
||||
if source_id in jpn_worknos_in_dataset or work_language(work) == "JPN":
|
||||
continue
|
||||
remove.add(int(work["__line_no"]))
|
||||
return remove
|
||||
|
||||
|
||||
def tag_names(tag: dict[str, Any]) -> list[str]:
|
||||
names: list[str] = []
|
||||
name = clean_text(tag.get("name"))
|
||||
if name:
|
||||
names.append(name)
|
||||
|
||||
i18n = tag.get("i18n")
|
||||
if isinstance(i18n, dict):
|
||||
for lang_data in i18n.values():
|
||||
if not isinstance(lang_data, dict):
|
||||
continue
|
||||
i18n_name = clean_text(lang_data.get("name"))
|
||||
if i18n_name:
|
||||
names.append(i18n_name)
|
||||
history = lang_data.get("history")
|
||||
if isinstance(history, list):
|
||||
for item in history:
|
||||
if isinstance(item, dict):
|
||||
history_name = clean_text(item.get("name"))
|
||||
if history_name:
|
||||
names.append(history_name)
|
||||
return names
|
||||
|
||||
|
||||
def has_excluded_tag(
|
||||
work: dict[str, Any],
|
||||
exclude_tags: set[str],
|
||||
exclude_tag_ids: set[int],
|
||||
) -> bool:
|
||||
tags = work.get("tags")
|
||||
if not isinstance(tags, list):
|
||||
return False
|
||||
|
||||
for tag in tags:
|
||||
if not isinstance(tag, dict):
|
||||
continue
|
||||
tag_id = tag.get("id")
|
||||
if isinstance(tag_id, int) and tag_id in exclude_tag_ids:
|
||||
return True
|
||||
if any(name in exclude_tags for name in tag_names(tag)):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def remove_marker(value: Any, marker: str) -> tuple[Any, int]:
|
||||
if isinstance(value, str):
|
||||
count = value.count(marker)
|
||||
return value.replace(marker, ""), count
|
||||
if isinstance(value, list):
|
||||
total = 0
|
||||
cleaned = []
|
||||
for item in value:
|
||||
cleaned_item, count = remove_marker(item, marker)
|
||||
cleaned.append(cleaned_item)
|
||||
total += count
|
||||
return cleaned, total
|
||||
if isinstance(value, dict):
|
||||
total = 0
|
||||
cleaned_dict: dict[str, Any] = {}
|
||||
for key, item in value.items():
|
||||
cleaned_item, count = remove_marker(item, marker)
|
||||
cleaned_dict[key] = cleaned_item
|
||||
total += count
|
||||
return cleaned_dict, total
|
||||
return value, 0
|
||||
|
||||
|
||||
def filter_works(
|
||||
works: list[dict[str, Any]],
|
||||
exclude_tags: set[str],
|
||||
exclude_tag_ids: set[int],
|
||||
) -> tuple[list[dict[str, Any]], dict[str, int]]:
|
||||
translated_lines = translated_line_numbers_to_remove(works)
|
||||
tag_lines = {
|
||||
int(work["__line_no"])
|
||||
for work in works
|
||||
if has_excluded_tag(work, exclude_tags, exclude_tag_ids)
|
||||
}
|
||||
remove_lines = translated_lines | tag_lines
|
||||
|
||||
filtered: list[dict[str, Any]] = []
|
||||
marker_occurrences = 0
|
||||
marker_records = 0
|
||||
for work in works:
|
||||
if int(work["__line_no"]) in remove_lines:
|
||||
continue
|
||||
cleaned, count = remove_marker(work, SIMPLIFIED_CHINESE_MARKER)
|
||||
if not isinstance(cleaned, dict):
|
||||
raise RuntimeError("cleaned work is not an object")
|
||||
marker_occurrences += count
|
||||
if count:
|
||||
marker_records += 1
|
||||
filtered.append(cleaned)
|
||||
|
||||
return filtered, {
|
||||
"input": len(works),
|
||||
"translation_removed": len(translated_lines),
|
||||
"tag_removed": len(tag_lines),
|
||||
"total_removed": len(remove_lines),
|
||||
"output": len(filtered),
|
||||
"marker_records": marker_records,
|
||||
"marker_occurrences": marker_occurrences,
|
||||
}
|
||||
|
||||
|
||||
def print_stats(stats: dict[str, int], output_path: Path | None, backup_path: Path | None) -> None:
|
||||
print(f"Read works: {stats['input']}", file=sys.stderr)
|
||||
print(f"Removed by translation preference: {stats['translation_removed']}", file=sys.stderr)
|
||||
print(f"Removed by excluded tags: {stats['tag_removed']}", file=sys.stderr)
|
||||
print(f"Removed total: {stats['total_removed']}", file=sys.stderr)
|
||||
print(f"Remaining works: {stats['output']}", file=sys.stderr)
|
||||
print(
|
||||
f"Removed {SIMPLIFIED_CHINESE_MARKER}: "
|
||||
f"{stats['marker_occurrences']} occurrences in {stats['marker_records']} remaining records",
|
||||
file=sys.stderr,
|
||||
)
|
||||
if output_path is not None:
|
||||
print(f"Output: {output_path}", file=sys.stderr)
|
||||
if backup_path is not None:
|
||||
print(f"Backup: {backup_path}", file=sys.stderr)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
args = parse_args()
|
||||
input_path = Path(args.input)
|
||||
if not input_path.exists():
|
||||
raise SystemExit(f"Input file does not exist: {input_path}")
|
||||
|
||||
if args.output and args.in_place:
|
||||
raise SystemExit("Use either --output or --in-place, not both.")
|
||||
if not args.output and not args.in_place and not args.dry_run:
|
||||
raise SystemExit("Specify --output, --in-place, or --dry-run.")
|
||||
|
||||
exclude_tags = set(args.exclude_tag)
|
||||
exclude_tag_ids = set(args.exclude_tag_id)
|
||||
if args.keep_default_exclude_tags:
|
||||
exclude_tags.update(DEFAULT_EXCLUDE_TAGS)
|
||||
exclude_tag_ids.update(DEFAULT_EXCLUDE_TAG_IDS)
|
||||
|
||||
works = read_jsonl(input_path)
|
||||
filtered, stats = filter_works(works, exclude_tags, exclude_tag_ids)
|
||||
|
||||
output_path: Path | None = None
|
||||
backup_path: Path | None = None
|
||||
if not args.dry_run:
|
||||
if args.in_place:
|
||||
backup_path = input_path.with_name(input_path.name + args.backup_suffix)
|
||||
shutil.copy2(input_path, backup_path)
|
||||
output_path = input_path
|
||||
else:
|
||||
output_path = Path(args.output)
|
||||
write_jsonl(output_path, filtered)
|
||||
|
||||
print_stats(stats, output_path, backup_path)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,327 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Remove translated duplicate works from filtered JSONL and web/data assets."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import mmap
|
||||
from collections import defaultdict
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Prune translated duplicates when the original Japanese work is present."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--filtered-input",
|
||||
default="asmr_works.filtered.jsonl",
|
||||
help="Filtered JSONL file to prune.",
|
||||
)
|
||||
parser.add_argument("--data-dir", default="web/data", help="Directory containing web data assets.")
|
||||
parser.add_argument("--dry-run", action="store_true", help="Show removals without writing files.")
|
||||
parser.add_argument(
|
||||
"--diff-limit",
|
||||
type=int,
|
||||
default=50,
|
||||
help="Maximum removal rows to print. Use 0 for all rows.",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def clean_text(value: Any) -> str:
|
||||
return value.strip() if isinstance(value, str) else ""
|
||||
|
||||
|
||||
def source_id(work: dict[str, Any]) -> str:
|
||||
return clean_text(work.get("source_id")) or clean_text(work.get("sourceId"))
|
||||
|
||||
|
||||
def identity_keys(work: dict[str, Any]) -> list[str]:
|
||||
keys: list[str] = []
|
||||
work_source_id = source_id(work)
|
||||
if work_source_id:
|
||||
keys.append(f"source:{work_source_id.lower()}")
|
||||
work_id = work.get("id")
|
||||
if work_id is not None:
|
||||
keys.append(f"id:{work_id}")
|
||||
return keys
|
||||
|
||||
|
||||
def primary_key(work: dict[str, Any]) -> str:
|
||||
keys = identity_keys(work)
|
||||
if keys:
|
||||
return keys[0]
|
||||
return f"title:{clean_text(work.get('title')).lower()}"
|
||||
|
||||
|
||||
def display_id(work: dict[str, Any]) -> str:
|
||||
return source_id(work) or str(work.get("id") or "unknown")
|
||||
|
||||
|
||||
def display_title(work: dict[str, Any]) -> str:
|
||||
title = clean_text(work.get("title")) or "title unknown"
|
||||
return title if len(title) <= 90 else title[:87] + "..."
|
||||
|
||||
|
||||
def translation_info(work: dict[str, Any]) -> dict[str, Any]:
|
||||
value = work.get("translation_info")
|
||||
return value if isinstance(value, dict) else {}
|
||||
|
||||
|
||||
def work_language(work: dict[str, Any]) -> str:
|
||||
info = translation_info(work)
|
||||
lang = clean_text(info.get("lang"))
|
||||
if lang:
|
||||
return lang
|
||||
if info.get("is_original") is True:
|
||||
return "JPN"
|
||||
editions = work.get("language_editions")
|
||||
work_source_id = source_id(work)
|
||||
if isinstance(editions, list) and work_source_id:
|
||||
for item in editions:
|
||||
if isinstance(item, dict) and clean_text(item.get("workno")) == work_source_id:
|
||||
return clean_text(item.get("lang")) or clean_text(item.get("label"))
|
||||
return ""
|
||||
|
||||
|
||||
def original_workno(work: dict[str, Any]) -> str:
|
||||
info = translation_info(work)
|
||||
for value in (info.get("original_workno"), work.get("original_workno")):
|
||||
text = clean_text(value)
|
||||
if text:
|
||||
return text
|
||||
|
||||
editions = work.get("language_editions")
|
||||
if isinstance(editions, list):
|
||||
for item in editions:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if item.get("lang") == "JPN":
|
||||
workno = clean_text(item.get("workno"))
|
||||
if workno:
|
||||
return workno
|
||||
|
||||
other_editions = work.get("other_language_editions_in_db")
|
||||
if isinstance(other_editions, list):
|
||||
for item in other_editions:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if item.get("is_original") is True:
|
||||
workno = clean_text(item.get("source_id"))
|
||||
if workno:
|
||||
return workno
|
||||
|
||||
if info.get("is_original") is True:
|
||||
return source_id(work)
|
||||
return ""
|
||||
|
||||
|
||||
def translation_group_key(work: dict[str, Any]) -> str:
|
||||
original = original_workno(work)
|
||||
if original:
|
||||
return f"work:{original}"
|
||||
work_source_id = source_id(work)
|
||||
if work_source_id:
|
||||
return f"work:{work_source_id}"
|
||||
return f"id:{work.get('id')}"
|
||||
|
||||
|
||||
def read_jsonl(path: Path) -> list[dict[str, Any]]:
|
||||
works: list[dict[str, Any]] = []
|
||||
with path.open(encoding="utf-8") as file:
|
||||
for line_no, line in enumerate(file, 1):
|
||||
if not line.strip():
|
||||
continue
|
||||
try:
|
||||
value = json.loads(line)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise SystemExit(f"{path}:{line_no}: invalid JSON: {exc}") from exc
|
||||
if not isinstance(value, dict):
|
||||
raise SystemExit(f"{path}:{line_no}: expected JSON object")
|
||||
value["__line_no"] = line_no
|
||||
works.append(value)
|
||||
return works
|
||||
|
||||
|
||||
def write_jsonl_atomic(path: Path, works: Iterable[dict[str, Any]]) -> None:
|
||||
temp_path = path.with_name(path.name + ".tmp")
|
||||
try:
|
||||
with temp_path.open("w", encoding="utf-8") as file:
|
||||
for work in works:
|
||||
clean_work = {key: value for key, value in work.items() if key != "__line_no"}
|
||||
file.write(json.dumps(clean_work, ensure_ascii=False, separators=(",", ":")))
|
||||
file.write("\n")
|
||||
temp_path.replace(path)
|
||||
except Exception:
|
||||
temp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
|
||||
|
||||
def read_json(path: Path) -> Any:
|
||||
with path.open(encoding="utf-8") as file:
|
||||
return json.load(file)
|
||||
|
||||
|
||||
def write_json_atomic(path: Path, value: Any, *, indent: int | None = None) -> None:
|
||||
temp_path = path.with_name(path.name + ".tmp")
|
||||
try:
|
||||
with temp_path.open("w", encoding="utf-8") as file:
|
||||
json.dump(value, file, ensure_ascii=False, indent=indent, separators=None if indent else (",", ":"))
|
||||
if indent:
|
||||
file.write("\n")
|
||||
temp_path.replace(path)
|
||||
except Exception:
|
||||
temp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
|
||||
|
||||
def find_translated_duplicates(works: list[dict[str, Any]]) -> tuple[set[str], list[dict[str, Any]], int]:
|
||||
groups: dict[str, list[dict[str, Any]]] = defaultdict(list)
|
||||
for work in works:
|
||||
groups[translation_group_key(work)].append(work)
|
||||
|
||||
remove_keys: set[str] = set()
|
||||
removals: list[dict[str, Any]] = []
|
||||
duplicate_groups = 0
|
||||
for key, group in groups.items():
|
||||
originals = [work for work in group if work_language(work) == "JPN" and source_id(work)]
|
||||
translated = [work for work in group if work_language(work) not in ("", "JPN")]
|
||||
if not originals or not translated:
|
||||
continue
|
||||
duplicate_groups += 1
|
||||
keep = originals[0]
|
||||
for work in translated:
|
||||
for identity_key in identity_keys(work):
|
||||
remove_keys.add(identity_key)
|
||||
removals.append(
|
||||
{
|
||||
"remove": work,
|
||||
"keep": keep,
|
||||
"group": key,
|
||||
"lang": work_language(work) or "unknown",
|
||||
}
|
||||
)
|
||||
return remove_keys, removals, duplicate_groups
|
||||
|
||||
|
||||
def should_remove(work: dict[str, Any], remove_keys: set[str]) -> bool:
|
||||
return any(key in remove_keys for key in identity_keys(work))
|
||||
|
||||
|
||||
def normalized_web_works(works: list[dict[str, Any]], remove_keys: set[str]) -> list[dict[str, Any]]:
|
||||
kept: list[dict[str, Any]] = []
|
||||
for work in works:
|
||||
if should_remove(work, remove_keys):
|
||||
continue
|
||||
updated = dict(work)
|
||||
updated["embeddingIndex"] = len(kept)
|
||||
kept.append(updated)
|
||||
return kept
|
||||
|
||||
|
||||
def rebuild_embeddings(
|
||||
*,
|
||||
data_dir: Path,
|
||||
old_works: list[dict[str, Any]],
|
||||
kept_old_indexes: list[int],
|
||||
dimensions: int,
|
||||
) -> None:
|
||||
embeddings_path = data_dir / "embeddings.f32"
|
||||
vector_size = dimensions * 4
|
||||
expected_size = len(old_works) * vector_size
|
||||
actual_size = embeddings_path.stat().st_size
|
||||
if actual_size != expected_size:
|
||||
raise RuntimeError(f"embedding size mismatch: expected {expected_size} bytes, got {actual_size} bytes")
|
||||
|
||||
temp_path = data_dir / "embeddings.f32.tmp"
|
||||
try:
|
||||
with embeddings_path.open("rb") as old_file, temp_path.open("wb") as new_file:
|
||||
with mmap.mmap(old_file.fileno(), 0, access=mmap.ACCESS_READ) as old_map:
|
||||
for old_index in kept_old_indexes:
|
||||
offset = old_index * vector_size
|
||||
new_file.write(old_map[offset : offset + vector_size])
|
||||
temp_path.replace(embeddings_path)
|
||||
except Exception:
|
||||
temp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
|
||||
|
||||
def print_removals(removals: list[dict[str, Any]], limit: int) -> None:
|
||||
rows = removals if limit == 0 else removals[:limit]
|
||||
for item in rows:
|
||||
remove = item["remove"]
|
||||
keep = item["keep"]
|
||||
print(
|
||||
f" - {display_id(remove)} -> keep {display_id(keep)} "
|
||||
f"/ {item['lang']} / {display_title(remove)}"
|
||||
)
|
||||
remaining = 0 if limit == 0 else max(0, len(removals) - limit)
|
||||
if remaining:
|
||||
print(f" ... {remaining} more")
|
||||
|
||||
|
||||
def validate_args(args: argparse.Namespace) -> None:
|
||||
if args.diff_limit < 0:
|
||||
raise SystemExit("--diff-limit must be greater than or equal to 0")
|
||||
|
||||
|
||||
def main() -> int:
|
||||
args = parse_args()
|
||||
validate_args(args)
|
||||
filtered_path = Path(args.filtered_input)
|
||||
data_dir = Path(args.data_dir)
|
||||
manifest_path = data_dir / "manifest.json"
|
||||
works_path = data_dir / "works.json"
|
||||
|
||||
filtered_works = read_jsonl(filtered_path)
|
||||
remove_keys, removals, duplicate_groups = find_translated_duplicates(filtered_works)
|
||||
pruned_filtered = [work for work in filtered_works if not should_remove(work, remove_keys)]
|
||||
|
||||
manifest = read_json(manifest_path)
|
||||
web_works = read_json(works_path)
|
||||
if not isinstance(web_works, list):
|
||||
raise SystemExit(f"{works_path} must contain a JSON array")
|
||||
dimensions = int(manifest["dimensions"])
|
||||
kept_old_indexes = [index for index, work in enumerate(web_works) if not should_remove(work, remove_keys)]
|
||||
pruned_web_works = normalized_web_works(web_works, remove_keys)
|
||||
|
||||
print(f"Translation duplicate groups: {duplicate_groups}")
|
||||
print(f"Works to remove: {len(removals)}")
|
||||
print(f"Filtered works: {len(filtered_works)} -> {len(pruned_filtered)}")
|
||||
print(f"web/data works: {len(web_works)} -> {len(pruned_web_works)}")
|
||||
print(f"Embeddings to keep: {len(kept_old_indexes)} / {len(web_works)}")
|
||||
print_removals(removals, args.diff_limit)
|
||||
|
||||
if args.dry_run:
|
||||
print("Dry run: no files were written.")
|
||||
return 0
|
||||
|
||||
write_jsonl_atomic(filtered_path, pruned_filtered)
|
||||
rebuild_embeddings(
|
||||
data_dir=data_dir,
|
||||
old_works=web_works,
|
||||
kept_old_indexes=kept_old_indexes,
|
||||
dimensions=dimensions,
|
||||
)
|
||||
new_manifest = dict(manifest)
|
||||
new_manifest.update(
|
||||
{
|
||||
"count": len(pruned_web_works),
|
||||
"embeddingFile": "embeddings.f32",
|
||||
"worksFile": "works.json",
|
||||
"generatedAt": datetime.now(timezone.utc).isoformat(),
|
||||
}
|
||||
)
|
||||
write_json_atomic(works_path, pruned_web_works)
|
||||
write_json_atomic(manifest_path, new_manifest, indent=2)
|
||||
print(f"Updated {filtered_path}")
|
||||
print(f"Updated {data_dir}")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,569 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Serve the WebUI and provide local natural-language vector search."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import http.client
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
from functools import partial
|
||||
from http import HTTPStatus
|
||||
from http.server import SimpleHTTPRequestHandler, ThreadingHTTPServer
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.error import HTTPError, URLError
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
from urllib.request import Request, urlopen
|
||||
|
||||
|
||||
DEFAULT_HOST = "127.0.0.1"
|
||||
DEFAULT_PORT = 8000
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Serve built Astro WebUI and expose /api/search for natural-language ASMR search."
|
||||
)
|
||||
parser.add_argument("--host", default=DEFAULT_HOST, help="Bind host.")
|
||||
parser.add_argument("--port", type=int, default=DEFAULT_PORT, help="Bind port.")
|
||||
parser.add_argument("--web-dir", default="web/dist", help="Directory to serve as static WebUI.")
|
||||
parser.add_argument("--data-dir", default="web/data", help="Directory containing manifest/works/embeddings.")
|
||||
parser.add_argument("--batch-size", type=int, default=1, help="SentenceTransformer encode batch size.")
|
||||
parser.add_argument(
|
||||
"--api-key-env",
|
||||
default="EMBEDDING_API_KEY",
|
||||
help="Environment variable that contains the remote embeddings API key.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--remote-base-url",
|
||||
default=None,
|
||||
help="Override manifest baseUrl for remote-openai-compatible search.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--remote-query-prefix",
|
||||
default="",
|
||||
help="Optional prefix prepended to remote search queries, for example 'query: '.",
|
||||
)
|
||||
parser.add_argument("--timeout", type=float, default=120.0, help="Remote request timeout seconds.")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def normalize(value: Any) -> str:
|
||||
return str(value or "").lower().strip()
|
||||
|
||||
|
||||
def parse_limit(value: Any, default: int = 48, maximum: int = 200) -> int:
|
||||
return max(1, min(maximum, int(value or default)))
|
||||
|
||||
|
||||
def searchable_text(work: dict[str, Any]) -> str:
|
||||
values = [
|
||||
work.get("title"),
|
||||
work.get("sourceId"),
|
||||
work.get("circle"),
|
||||
*(work.get("vas") or []),
|
||||
*(work.get("tags") or []),
|
||||
]
|
||||
return " ".join(normalize(value) for value in values)
|
||||
|
||||
|
||||
def recent_sort_value(work: dict[str, Any]) -> str:
|
||||
return str(work.get("createDate") or work.get("release") or work.get("sourceId") or "")
|
||||
|
||||
|
||||
def embeddings_url(base_url: str) -> str:
|
||||
return f"{base_url.rstrip('/')}/v1/embeddings"
|
||||
|
||||
|
||||
def normalize_vector(vector: list[float]) -> list[float]:
|
||||
norm = math.sqrt(sum(value * value for value in vector))
|
||||
if norm == 0:
|
||||
return vector
|
||||
return [value / norm for value in vector]
|
||||
|
||||
|
||||
def request_remote_embedding(
|
||||
*,
|
||||
url: str,
|
||||
model: str,
|
||||
query: str,
|
||||
api_key: str,
|
||||
timeout: float,
|
||||
) -> list[float]:
|
||||
body = json.dumps({"model": model, "input": [query]}).encode("utf-8")
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if api_key:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
|
||||
request = Request(url, data=body, headers=headers, method="POST")
|
||||
try:
|
||||
with urlopen(request, timeout=timeout) as response:
|
||||
payload = json.load(response)
|
||||
except HTTPError as exc:
|
||||
detail = exc.read(1000).decode("utf-8", errors="replace")
|
||||
raise RuntimeError(f"remote embeddings API returned HTTP {exc.code}: {detail}") from exc
|
||||
except (
|
||||
URLError,
|
||||
TimeoutError,
|
||||
json.JSONDecodeError,
|
||||
http.client.IncompleteRead,
|
||||
http.client.RemoteDisconnected,
|
||||
ConnectionResetError,
|
||||
OSError,
|
||||
) as exc:
|
||||
raise RuntimeError(f"remote embeddings request failed: {exc}") from exc
|
||||
|
||||
data = payload.get("data") if isinstance(payload, dict) else None
|
||||
if not isinstance(data, list) or not data:
|
||||
raise RuntimeError("remote embeddings response does not contain data[]")
|
||||
first = data[0]
|
||||
if not isinstance(first, dict):
|
||||
raise RuntimeError("remote embeddings response data[0] is not an object")
|
||||
embedding = first.get("embedding")
|
||||
if not isinstance(embedding, list) or not embedding:
|
||||
raise RuntimeError("remote embeddings response does not contain embedding")
|
||||
try:
|
||||
return [float(value) for value in embedding]
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise RuntimeError("remote embedding contains non-numeric values") from exc
|
||||
|
||||
|
||||
class LocalSentenceTransformerQueryEncoder:
|
||||
def __init__(self, model_name: str, batch_size: int) -> None:
|
||||
try:
|
||||
from sentence_transformers import SentenceTransformer
|
||||
except ImportError as exc:
|
||||
raise SystemExit(
|
||||
"Local natural-language search requires semantic dependencies. "
|
||||
"Run `uv run --extra semantic python scripts/search_server.py`."
|
||||
) from exc
|
||||
|
||||
self.model_name = model_name
|
||||
self.batch_size = batch_size
|
||||
self.model = SentenceTransformer(model_name)
|
||||
|
||||
def encode(self, query: str) -> list[float]:
|
||||
text = f"query: {query}" if "e5" in self.model_name.lower() else query
|
||||
embedding = self.model.encode(
|
||||
[text],
|
||||
batch_size=self.batch_size,
|
||||
normalize_embeddings=True,
|
||||
show_progress_bar=False,
|
||||
)[0]
|
||||
return [float(value) for value in embedding]
|
||||
|
||||
|
||||
class RemoteOpenAICompatibleQueryEncoder:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
base_url: str,
|
||||
model: str,
|
||||
api_key: str,
|
||||
timeout: float,
|
||||
query_prefix: str,
|
||||
should_normalize: bool,
|
||||
) -> None:
|
||||
self.url = embeddings_url(base_url)
|
||||
self.model = model
|
||||
self.api_key = api_key
|
||||
self.timeout = timeout
|
||||
self.query_prefix = query_prefix
|
||||
self.should_normalize = should_normalize
|
||||
|
||||
def encode(self, query: str) -> list[float]:
|
||||
vector = request_remote_embedding(
|
||||
url=self.url,
|
||||
model=self.model,
|
||||
query=f"{self.query_prefix}{query}",
|
||||
api_key=self.api_key,
|
||||
timeout=self.timeout,
|
||||
)
|
||||
return normalize_vector(vector) if self.should_normalize else vector
|
||||
|
||||
|
||||
class SearchIndex:
|
||||
def __init__(self, args: argparse.Namespace, data_dir: Path) -> None:
|
||||
try:
|
||||
import numpy as np
|
||||
except ImportError as exc:
|
||||
raise SystemExit(
|
||||
"Natural-language search requires numpy. "
|
||||
"Run `uv run --extra remote python scripts/search_server.py` or "
|
||||
"`uv run --extra semantic python scripts/search_server.py`."
|
||||
) from exc
|
||||
|
||||
self.np = np
|
||||
self.manifest = self._read_json(data_dir / "manifest.json")
|
||||
self.works = self._read_json(data_dir / "works.json")
|
||||
self.method = str(self.manifest.get("method") or "")
|
||||
self.dimensions = int(self.manifest["dimensions"])
|
||||
self.model_name = str(self.manifest["model"])
|
||||
embeddings_path = data_dir / str(self.manifest["embeddingFile"])
|
||||
self.embeddings = np.fromfile(embeddings_path, dtype="<f4")
|
||||
expected = len(self.works) * self.dimensions
|
||||
if self.embeddings.size != expected:
|
||||
raise SystemExit(
|
||||
f"embedding size mismatch: expected {expected}, got {self.embeddings.size}"
|
||||
)
|
||||
self.embeddings = self.embeddings.reshape((len(self.works), self.dimensions))
|
||||
score_config = self.manifest.get("score") if isinstance(self.manifest, dict) else None
|
||||
self.vector_weight = float((score_config or {}).get("vectorWeight", 0.8))
|
||||
self.tag_weight = float((score_config or {}).get("tagWeight", 0.2))
|
||||
self.query_encoder = self.create_query_encoder(args)
|
||||
|
||||
@staticmethod
|
||||
def _read_json(path: Path) -> Any:
|
||||
with path.open(encoding="utf-8") as file:
|
||||
return json.load(file)
|
||||
|
||||
def create_query_encoder(self, args: argparse.Namespace):
|
||||
if self.method == "sentence-transformers":
|
||||
return LocalSentenceTransformerQueryEncoder(self.model_name, args.batch_size)
|
||||
|
||||
if self.method == "remote-openai-compatible":
|
||||
base_url = args.remote_base_url or self.manifest.get("baseUrl")
|
||||
if not base_url:
|
||||
raise SystemExit(
|
||||
"remote-openai-compatible manifest does not contain baseUrl. "
|
||||
"Pass --remote-base-url."
|
||||
)
|
||||
return RemoteOpenAICompatibleQueryEncoder(
|
||||
base_url=str(base_url),
|
||||
model=self.model_name,
|
||||
api_key=os.environ.get(args.api_key_env, ""),
|
||||
timeout=args.timeout,
|
||||
query_prefix=args.remote_query_prefix,
|
||||
should_normalize=bool(self.manifest.get("normalized", True)),
|
||||
)
|
||||
|
||||
raise SystemExit(
|
||||
f"unsupported embedding method for natural-language search: {self.method}. "
|
||||
"Use sentence-transformers or remote-openai-compatible data."
|
||||
)
|
||||
|
||||
def query_embedding(self, query: str):
|
||||
embedding = self.query_encoder.encode(query)
|
||||
if len(embedding) != self.dimensions:
|
||||
raise RuntimeError(
|
||||
f"query embedding dimensions mismatch: expected {self.dimensions}, got {len(embedding)}"
|
||||
)
|
||||
return self.np.asarray(embedding, dtype="<f4")
|
||||
|
||||
def search(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
query = str(payload.get("query") or "").strip()
|
||||
if not query:
|
||||
raise ValueError("query is required")
|
||||
|
||||
limit = parse_limit(payload.get("limit"))
|
||||
ages = {normalize(value) for value in payload.get("ages") or [] if normalize(value)}
|
||||
required_tags = [normalize(value) for value in payload.get("requiredTags") or [] if normalize(value)]
|
||||
|
||||
try:
|
||||
query_vector = self.query_embedding(query)
|
||||
except RuntimeError as exc:
|
||||
raise ValueError(str(exc)) from exc
|
||||
scores = self.embeddings @ query_vector
|
||||
candidates = self.filtered_indices(ages, required_tags)
|
||||
if candidates.size == 0:
|
||||
return {"results": []}
|
||||
|
||||
candidate_scores = scores[candidates]
|
||||
take = min(limit, candidate_scores.size)
|
||||
top_positions = self.np.argpartition(-candidate_scores, take - 1)[:take]
|
||||
top_positions = top_positions[self.np.argsort(-candidate_scores[top_positions])]
|
||||
|
||||
results = [
|
||||
{
|
||||
"index": int(candidates[position]),
|
||||
"score": float(candidate_scores[position]),
|
||||
"work": self.works[int(candidates[position])],
|
||||
}
|
||||
for position in top_positions
|
||||
]
|
||||
return {"results": results}
|
||||
|
||||
def list_works(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
limit = parse_limit(payload.get("limit"), default=100)
|
||||
query = str(payload.get("query") or payload.get("q") or "").strip()
|
||||
ages = {normalize(value) for value in payload.get("ages") or [] if normalize(value)}
|
||||
required_tags = [
|
||||
normalize(value)
|
||||
for value in (payload.get("requiredTags") or payload.get("tags") or [])
|
||||
if normalize(value)
|
||||
]
|
||||
|
||||
entries = [
|
||||
work
|
||||
for work in self.works
|
||||
if self.matches_work(work, ages=ages, required_tags=required_tags, query=query)
|
||||
]
|
||||
entries.sort(key=recent_sort_value, reverse=True)
|
||||
return {"results": [{"work": work} for work in entries[:limit]]}
|
||||
|
||||
def get_work(self, identifier: Any) -> dict[str, Any]:
|
||||
normalized_id = normalize(identifier)
|
||||
if not normalized_id:
|
||||
raise ValueError("id is required")
|
||||
|
||||
for work in self.works:
|
||||
if (
|
||||
normalize(work.get("sourceId")) == normalized_id
|
||||
or normalize(work.get("id")) == normalized_id
|
||||
or normalize(work.get("embeddingIndex")) == normalized_id
|
||||
):
|
||||
return {"work": work}
|
||||
raise ValueError("work not found")
|
||||
|
||||
def similar(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
try:
|
||||
source_index = int(payload.get("index"))
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError("index is required") from exc
|
||||
if not 0 <= source_index < len(self.works):
|
||||
raise ValueError("index is out of range")
|
||||
|
||||
limit = parse_limit(payload.get("limit"))
|
||||
filters = payload.get("filters") if isinstance(payload.get("filters"), dict) else {}
|
||||
ages = {normalize(value) for value in filters.get("ages") or [] if normalize(value)}
|
||||
exclude_same_group = bool(filters.get("excludeSameGroup"))
|
||||
|
||||
source_work = self.works[source_index]
|
||||
source_group = source_work.get("groupId")
|
||||
source_tags = set(source_work.get("tags") or [])
|
||||
scores = self.embeddings @ self.embeddings[source_index]
|
||||
|
||||
candidates: list[int] = []
|
||||
combined_scores: list[float] = []
|
||||
vector_scores: list[float] = []
|
||||
tag_scores: list[float] = []
|
||||
|
||||
for index, work in enumerate(self.works):
|
||||
if index == source_index:
|
||||
continue
|
||||
if exclude_same_group and source_group and source_group == work.get("groupId"):
|
||||
continue
|
||||
if ages and normalize(work.get("ageCategory") or "unknown") not in ages:
|
||||
continue
|
||||
|
||||
vector_score = float(scores[index])
|
||||
tag_score = tag_jaccard(source_tags, work.get("tags") or [])
|
||||
score = self.vector_weight * vector_score + self.tag_weight * tag_score
|
||||
candidates.append(index)
|
||||
combined_scores.append(score)
|
||||
vector_scores.append(vector_score)
|
||||
tag_scores.append(tag_score)
|
||||
|
||||
if not candidates:
|
||||
return {"results": []}
|
||||
|
||||
order = sorted(range(len(candidates)), key=lambda position: combined_scores[position], reverse=True)
|
||||
results = [
|
||||
{
|
||||
"index": int(candidates[position]),
|
||||
"score": float(combined_scores[position]),
|
||||
"vector": float(vector_scores[position]),
|
||||
"tag": float(tag_scores[position]),
|
||||
"work": self.works[int(candidates[position])],
|
||||
}
|
||||
for position in order[:limit]
|
||||
]
|
||||
return {"results": results}
|
||||
|
||||
def filtered_indices(self, ages: set[str], required_tags: list[str]):
|
||||
indices: list[int] = []
|
||||
for index, work in enumerate(self.works):
|
||||
if ages and normalize(work.get("ageCategory")) not in ages:
|
||||
continue
|
||||
if required_tags:
|
||||
tags = " ".join(normalize(tag) for tag in work.get("tags") or [])
|
||||
if not all(tag in tags for tag in required_tags):
|
||||
continue
|
||||
indices.append(index)
|
||||
return self.np.asarray(indices, dtype="int64")
|
||||
|
||||
@staticmethod
|
||||
def matches_work(
|
||||
work: dict[str, Any],
|
||||
*,
|
||||
ages: set[str],
|
||||
required_tags: list[str],
|
||||
query: str,
|
||||
) -> bool:
|
||||
if ages and normalize(work.get("ageCategory") or "unknown") not in ages:
|
||||
return False
|
||||
if required_tags:
|
||||
tags = " ".join(normalize(tag) for tag in work.get("tags") or [])
|
||||
if not all(tag in tags for tag in required_tags):
|
||||
return False
|
||||
if query and normalize(query) not in searchable_text(work):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def tag_jaccard(source_tags: set[str], candidate_tags: list[Any]) -> float:
|
||||
if not source_tags and not candidate_tags:
|
||||
return 0.0
|
||||
|
||||
intersection = 0
|
||||
seen: set[Any] = set()
|
||||
for tag in candidate_tags:
|
||||
if tag in seen:
|
||||
continue
|
||||
seen.add(tag)
|
||||
if tag in source_tags:
|
||||
intersection += 1
|
||||
|
||||
union = len(source_tags) + len(seen) - intersection
|
||||
return intersection / union if union else 0.0
|
||||
|
||||
|
||||
class SearchRequestHandler(SimpleHTTPRequestHandler):
|
||||
search_index: SearchIndex
|
||||
|
||||
def end_headers(self) -> None:
|
||||
self.send_header("Access-Control-Allow-Origin", "*")
|
||||
self.send_header("Access-Control-Allow-Headers", "content-type")
|
||||
self.send_header("Access-Control-Allow-Methods", "GET, POST, OPTIONS")
|
||||
super().end_headers()
|
||||
|
||||
def do_OPTIONS(self) -> None:
|
||||
self.send_response(HTTPStatus.NO_CONTENT)
|
||||
self.end_headers()
|
||||
|
||||
def do_GET(self) -> None:
|
||||
parsed = urlparse(self.path)
|
||||
if parsed.path == "/api/status":
|
||||
self.write_json(
|
||||
{
|
||||
"ok": True,
|
||||
"count": len(self.search_index.works),
|
||||
"model": self.search_index.model_name,
|
||||
}
|
||||
)
|
||||
return
|
||||
if parsed.path == "/api/search":
|
||||
params = parse_qs(parsed.query)
|
||||
payload = {
|
||||
"query": (params.get("q") or [""])[0],
|
||||
"limit": (params.get("limit") or [48])[0],
|
||||
}
|
||||
self.handle_search(payload)
|
||||
return
|
||||
if parsed.path == "/api/works":
|
||||
params = parse_qs(parsed.query)
|
||||
payload = {
|
||||
"query": (params.get("q") or [""])[0],
|
||||
"limit": (params.get("limit") or [100])[0],
|
||||
"ages": params.get("age") or [],
|
||||
"requiredTags": params.get("tag") or [],
|
||||
}
|
||||
self.handle_works(payload)
|
||||
return
|
||||
if parsed.path == "/api/work":
|
||||
params = parse_qs(parsed.query)
|
||||
self.handle_work((params.get("id") or [""])[0])
|
||||
return
|
||||
if parsed.path == "/api/similar":
|
||||
params = parse_qs(parsed.query)
|
||||
payload = {
|
||||
"index": (params.get("index") or [None])[0],
|
||||
"limit": (params.get("limit") or [48])[0],
|
||||
"filters": {
|
||||
"ages": params.get("age") or [],
|
||||
"excludeSameGroup": (params.get("excludeSameGroup") or [""])[0] in {"1", "true"},
|
||||
},
|
||||
}
|
||||
self.handle_similar(payload)
|
||||
return
|
||||
super().do_GET()
|
||||
|
||||
def do_POST(self) -> None:
|
||||
parsed = urlparse(self.path)
|
||||
if parsed.path not in {"/api/search", "/api/similar", "/api/works", "/api/work"}:
|
||||
self.send_error(HTTPStatus.NOT_FOUND)
|
||||
return
|
||||
|
||||
try:
|
||||
length = int(self.headers.get("content-length") or 0)
|
||||
body = self.rfile.read(length).decode("utf-8") if length else "{}"
|
||||
payload = json.loads(body)
|
||||
if not isinstance(payload, dict):
|
||||
raise ValueError("request body must be a JSON object")
|
||||
except (UnicodeDecodeError, json.JSONDecodeError, ValueError) as exc:
|
||||
self.write_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST)
|
||||
return
|
||||
|
||||
if parsed.path == "/api/search":
|
||||
self.handle_search(payload)
|
||||
elif parsed.path == "/api/similar":
|
||||
self.handle_similar(payload)
|
||||
elif parsed.path == "/api/works":
|
||||
self.handle_works(payload)
|
||||
else:
|
||||
self.handle_work(payload.get("id"))
|
||||
|
||||
def handle_search(self, payload: dict[str, Any]) -> None:
|
||||
try:
|
||||
self.write_json(self.search_index.search(payload))
|
||||
except ValueError as exc:
|
||||
self.write_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST)
|
||||
|
||||
def handle_works(self, payload: dict[str, Any]) -> None:
|
||||
try:
|
||||
self.write_json(self.search_index.list_works(payload))
|
||||
except ValueError as exc:
|
||||
self.write_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST)
|
||||
|
||||
def handle_work(self, identifier: Any) -> None:
|
||||
try:
|
||||
self.write_json(self.search_index.get_work(identifier))
|
||||
except ValueError as exc:
|
||||
status = HTTPStatus.NOT_FOUND if str(exc) == "work not found" else HTTPStatus.BAD_REQUEST
|
||||
self.write_json({"error": str(exc)}, status)
|
||||
|
||||
def handle_similar(self, payload: dict[str, Any]) -> None:
|
||||
try:
|
||||
self.write_json(self.search_index.similar(payload))
|
||||
except ValueError as exc:
|
||||
self.write_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST)
|
||||
|
||||
def write_json(self, payload: dict[str, Any], status: HTTPStatus = HTTPStatus.OK) -> None:
|
||||
body = json.dumps(payload, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
|
||||
self.send_response(status)
|
||||
self.send_header("Content-Type", "application/json; charset=utf-8")
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
args = parse_args()
|
||||
web_dir = Path(args.web_dir).resolve()
|
||||
data_dir = Path(args.data_dir).resolve()
|
||||
if not web_dir.exists():
|
||||
raise SystemExit(f"web directory not found: {web_dir}")
|
||||
if not data_dir.exists():
|
||||
raise SystemExit(f"data directory not found: {data_dir}")
|
||||
|
||||
print("Loading natural-language search index...")
|
||||
SearchRequestHandler.search_index = SearchIndex(args, data_dir)
|
||||
handler = partial(SearchRequestHandler, directory=str(web_dir))
|
||||
server = ThreadingHTTPServer((args.host, args.port), handler)
|
||||
|
||||
print(f"Serving http://{args.host}:{args.port}")
|
||||
print(f"Search model: {SearchRequestHandler.search_index.model_name}")
|
||||
print(f"Search method: {SearchRequestHandler.search_index.method}")
|
||||
try:
|
||||
server.serve_forever()
|
||||
except KeyboardInterrupt:
|
||||
print("\nShutting down...")
|
||||
finally:
|
||||
server.server_close()
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,757 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Incrementally fetch new ASMR works and update web/data assets."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import mmap
|
||||
import os
|
||||
import struct
|
||||
import sys
|
||||
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, Iterable
|
||||
|
||||
|
||||
ROOT_DIR = Path(__file__).resolve().parents[1]
|
||||
if str(ROOT_DIR) not in sys.path:
|
||||
sys.path.insert(0, str(ROOT_DIR))
|
||||
|
||||
from fetch_asmr_works import fetch_page # noqa: E402
|
||||
from build_vector_data import embedding_text, normalize_work # noqa: E402
|
||||
from build_vector_data_remote import ( # noqa: E402
|
||||
embeddings_url,
|
||||
normalize_vector,
|
||||
request_embeddings_with_retries,
|
||||
)
|
||||
from filter_asmr_works import ( # noqa: E402
|
||||
DEFAULT_EXCLUDE_TAG_IDS,
|
||||
DEFAULT_EXCLUDE_TAGS,
|
||||
filter_works,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Fetch latest ASMR works, filter them, and incrementally update web/data."
|
||||
)
|
||||
parser.add_argument("--input", default="asmr_works.jsonl", help="Local raw ASMR JSONL file.")
|
||||
parser.add_argument(
|
||||
"--filtered-output",
|
||||
default="asmr_works.filtered.jsonl",
|
||||
help="Filtered ASMR JSONL output file.",
|
||||
)
|
||||
parser.add_argument("--data-dir", default="web/data", help="Directory containing web data assets.")
|
||||
parser.add_argument("--order", default="create_date", help="ASMR API order parameter.")
|
||||
parser.add_argument("--sort", default="desc", help="ASMR API sort parameter.")
|
||||
parser.add_argument("--subtitle", default="0", help="ASMR API subtitle parameter.")
|
||||
parser.add_argument("--page-size", type=int, default=100, help="ASMR API pageSize parameter.")
|
||||
parser.add_argument("--timeout", type=float, default=30.0, help="ASMR API request timeout seconds.")
|
||||
parser.add_argument("--retries", type=int, default=3, help="ASMR API retries per failed page.")
|
||||
parser.add_argument(
|
||||
"--max-pages",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Maximum latest pages to scan. By default, stop when a page has no new works.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--force",
|
||||
action="store_true",
|
||||
help="Rebuild filtered/web data even when no new works are found.",
|
||||
)
|
||||
parser.add_argument("--dry-run", action="store_true", help="Fetch and report without writing files.")
|
||||
parser.add_argument(
|
||||
"--diff-limit",
|
||||
type=int,
|
||||
default=30,
|
||||
help="Maximum rows to print per dry-run diff section. Use 0 for all rows.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--exclude-tag",
|
||||
action="append",
|
||||
default=[],
|
||||
help="Additional tag name to remove during filtering.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--exclude-tag-id",
|
||||
action="append",
|
||||
type=int,
|
||||
default=[],
|
||||
help="Additional tag id to remove during filtering.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--keep-default-exclude-tags",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=True,
|
||||
help="Keep default 女性向/乙女向 filtering.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--api-key-env",
|
||||
default="EMBEDDING_API_KEY",
|
||||
help="Environment variable that contains the remote embedding API key.",
|
||||
)
|
||||
parser.add_argument("--embedding-timeout", type=float, default=120.0, help="Embedding request timeout seconds.")
|
||||
parser.add_argument("--embedding-retries", type=int, default=3, help="Embedding retries per failed batch.")
|
||||
parser.add_argument("--embedding-retry-wait", type=float, default=2.0, help="Initial embedding retry wait seconds.")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def clean_text(value: Any) -> str:
|
||||
return value.strip() if isinstance(value, str) else ""
|
||||
|
||||
|
||||
def identity_keys(work: dict[str, Any]) -> list[str]:
|
||||
keys: list[str] = []
|
||||
source_id = clean_text(work.get("source_id")) or clean_text(work.get("sourceId"))
|
||||
if source_id:
|
||||
keys.append(f"source:{source_id.lower()}")
|
||||
work_id = work.get("id")
|
||||
if work_id is not None:
|
||||
keys.append(f"id:{work_id}")
|
||||
return keys
|
||||
|
||||
|
||||
def primary_key(work: dict[str, Any]) -> str:
|
||||
keys = identity_keys(work)
|
||||
if not keys:
|
||||
raise RuntimeError(f"work has no id/source_id: {work!r}")
|
||||
return keys[0]
|
||||
|
||||
|
||||
def display_id(work: dict[str, Any]) -> str:
|
||||
return clean_text(work.get("source_id")) or clean_text(work.get("sourceId")) or str(work.get("id") or "unknown")
|
||||
|
||||
|
||||
def display_title(work: dict[str, Any]) -> str:
|
||||
return clean_text(work.get("title")) or "title unknown"
|
||||
|
||||
|
||||
def truncate(value: str, max_length: int = 90) -> str:
|
||||
if len(value) <= max_length:
|
||||
return value
|
||||
return value[: max_length - 3] + "..."
|
||||
|
||||
|
||||
def format_work(work: dict[str, Any]) -> str:
|
||||
date = clean_text(work.get("create_date")) or clean_text(work.get("createDate"))
|
||||
suffix = f" / {date}" if date else ""
|
||||
return f"{display_id(work)} / {truncate(display_title(work))}{suffix}"
|
||||
|
||||
|
||||
def read_jsonl(path: Path) -> list[dict[str, Any]]:
|
||||
works: list[dict[str, Any]] = []
|
||||
with path.open(encoding="utf-8") as file:
|
||||
for line_no, line in enumerate(file, 1):
|
||||
if not line.strip():
|
||||
continue
|
||||
try:
|
||||
value = json.loads(line)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise SystemExit(f"{path}:{line_no}: invalid JSON: {exc}") from exc
|
||||
if not isinstance(value, dict):
|
||||
raise SystemExit(f"{path}:{line_no}: expected JSON object")
|
||||
works.append(value)
|
||||
return works
|
||||
|
||||
|
||||
def write_jsonl_atomic(path: Path, works: Iterable[dict[str, Any]]) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
temp_path = path.with_name(path.name + ".tmp")
|
||||
try:
|
||||
with temp_path.open("w", encoding="utf-8") as file:
|
||||
for work in works:
|
||||
clean_work = {key: value for key, value in work.items() if key != "__line_no"}
|
||||
file.write(json.dumps(clean_work, ensure_ascii=False, separators=(",", ":")))
|
||||
file.write("\n")
|
||||
temp_path.replace(path)
|
||||
except Exception:
|
||||
temp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
|
||||
|
||||
def write_json_atomic(path: Path, value: Any, *, indent: int | None = None) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
temp_path = path.with_name(path.name + ".tmp")
|
||||
try:
|
||||
with temp_path.open("w", encoding="utf-8") as file:
|
||||
json.dump(value, file, ensure_ascii=False, indent=indent, separators=None if indent else (",", ":"))
|
||||
if indent:
|
||||
file.write("\n")
|
||||
temp_path.replace(path)
|
||||
except Exception:
|
||||
temp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
|
||||
|
||||
def fetch_latest_works(args: argparse.Namespace, existing_keys: set[str]) -> tuple[list[dict[str, Any]], dict[str, dict[str, Any]], int]:
|
||||
fetch_args = SimpleNamespace(
|
||||
order=args.order,
|
||||
sort=args.sort,
|
||||
subtitle=args.subtitle,
|
||||
page_size=args.page_size,
|
||||
timeout=args.timeout,
|
||||
retries=args.retries,
|
||||
)
|
||||
|
||||
new_works: list[dict[str, Any]] = []
|
||||
fetched_existing_by_key: dict[str, dict[str, Any]] = {}
|
||||
seen_new_keys: set[str] = set()
|
||||
scanned_pages = 0
|
||||
page = 1
|
||||
|
||||
while True:
|
||||
if args.max_pages is not None and page > args.max_pages:
|
||||
break
|
||||
|
||||
payload = fetch_page(fetch_args, page)
|
||||
scanned_pages += 1
|
||||
works = payload.get("works", [])
|
||||
if not isinstance(works, list) or not works:
|
||||
break
|
||||
|
||||
page_new = 0
|
||||
for item in works:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
keys = identity_keys(item)
|
||||
if not keys:
|
||||
continue
|
||||
if any(key in existing_keys for key in keys):
|
||||
for key in keys:
|
||||
fetched_existing_by_key[key] = item
|
||||
continue
|
||||
if any(key in seen_new_keys for key in keys):
|
||||
continue
|
||||
new_works.append(item)
|
||||
seen_new_keys.update(keys)
|
||||
page_new += 1
|
||||
|
||||
print(
|
||||
f"Fetched latest page {page}: {page_new} new / {len(works)} works",
|
||||
file=sys.stderr,
|
||||
)
|
||||
if page_new == 0:
|
||||
break
|
||||
page += 1
|
||||
|
||||
return new_works, fetched_existing_by_key, scanned_pages
|
||||
|
||||
|
||||
def merge_raw_works(
|
||||
existing_works: list[dict[str, Any]],
|
||||
new_works: list[dict[str, Any]],
|
||||
fetched_existing_by_key: dict[str, dict[str, Any]],
|
||||
) -> list[dict[str, Any]]:
|
||||
merged: list[dict[str, Any]] = []
|
||||
emitted_keys: set[str] = set()
|
||||
|
||||
for work in new_works:
|
||||
keys = identity_keys(work)
|
||||
if any(key in emitted_keys for key in keys):
|
||||
continue
|
||||
merged.append(work)
|
||||
emitted_keys.update(keys)
|
||||
|
||||
for work in existing_works:
|
||||
replacement = next((fetched_existing_by_key[key] for key in identity_keys(work) if key in fetched_existing_by_key), None)
|
||||
candidate = replacement or work
|
||||
keys = identity_keys(candidate)
|
||||
if keys and any(key in emitted_keys for key in keys):
|
||||
continue
|
||||
merged.append(candidate)
|
||||
emitted_keys.update(keys)
|
||||
|
||||
return merged
|
||||
|
||||
|
||||
def numbered_works(works: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
result: list[dict[str, Any]] = []
|
||||
for index, work in enumerate(works, 1):
|
||||
numbered = dict(work)
|
||||
numbered["__line_no"] = index
|
||||
result.append(numbered)
|
||||
return result
|
||||
|
||||
|
||||
def filter_raw_works(args: argparse.Namespace, works: list[dict[str, Any]]) -> tuple[list[dict[str, Any]], dict[str, int]]:
|
||||
exclude_tags = set(args.exclude_tag)
|
||||
exclude_tag_ids = set(args.exclude_tag_id)
|
||||
if args.keep_default_exclude_tags:
|
||||
exclude_tags.update(DEFAULT_EXCLUDE_TAGS)
|
||||
exclude_tag_ids.update(DEFAULT_EXCLUDE_TAG_IDS)
|
||||
return filter_works(numbered_works(works), exclude_tags, exclude_tag_ids)
|
||||
|
||||
|
||||
def read_json(path: Path) -> Any:
|
||||
with path.open(encoding="utf-8") as file:
|
||||
return json.load(file)
|
||||
|
||||
|
||||
def work_by_key(works: Iterable[dict[str, Any]]) -> dict[str, dict[str, Any]]:
|
||||
result: dict[str, dict[str, Any]] = {}
|
||||
for work in works:
|
||||
for key in identity_keys(work):
|
||||
result.setdefault(key, work)
|
||||
return result
|
||||
|
||||
|
||||
def unique_by_primary_key(works: Iterable[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
result: list[dict[str, Any]] = []
|
||||
seen: set[str] = set()
|
||||
for work in works:
|
||||
key = primary_key(work)
|
||||
if key in seen:
|
||||
continue
|
||||
seen.add(key)
|
||||
result.append(work)
|
||||
return result
|
||||
|
||||
|
||||
def raw_diff_reasons(old: dict[str, Any], new: dict[str, Any]) -> list[str]:
|
||||
old_normalized = normalize_work(old, 0)
|
||||
new_normalized = normalize_work(new, 0)
|
||||
checks = (
|
||||
("title", old_normalized.get("title"), new_normalized.get("title")),
|
||||
("tags", old_normalized.get("tags"), new_normalized.get("tags")),
|
||||
("vas", old_normalized.get("vas"), new_normalized.get("vas")),
|
||||
("circle", old_normalized.get("circle"), new_normalized.get("circle")),
|
||||
("duration", old_normalized.get("duration"), new_normalized.get("duration")),
|
||||
("dlCount", old_normalized.get("dlCount"), new_normalized.get("dlCount")),
|
||||
("rateAverage", old_normalized.get("rateAverage"), new_normalized.get("rateAverage")),
|
||||
("rateCount", old_normalized.get("rateCount"), new_normalized.get("rateCount")),
|
||||
("release", old_normalized.get("release"), new_normalized.get("release")),
|
||||
("createDate", old_normalized.get("createDate"), new_normalized.get("createDate")),
|
||||
)
|
||||
return [name for name, old_value, new_value in checks if old_value != new_value]
|
||||
|
||||
|
||||
def updated_existing_diffs(
|
||||
existing_works: list[dict[str, Any]],
|
||||
fetched_existing_by_key: dict[str, dict[str, Any]],
|
||||
) -> list[dict[str, Any]]:
|
||||
existing_by_key = work_by_key(existing_works)
|
||||
diffs: list[dict[str, Any]] = []
|
||||
for new_work in unique_by_primary_key(fetched_existing_by_key.values()):
|
||||
old_work = next((existing_by_key[key] for key in identity_keys(new_work) if key in existing_by_key), None)
|
||||
if old_work is None:
|
||||
continue
|
||||
reasons = raw_diff_reasons(old_work, new_work)
|
||||
if reasons:
|
||||
diffs.append({"work": new_work, "reasons": reasons})
|
||||
return diffs
|
||||
|
||||
|
||||
def filtered_change_diffs(
|
||||
filtered_path: Path,
|
||||
filtered_works: list[dict[str, Any]],
|
||||
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
||||
old_filtered = read_jsonl(filtered_path) if filtered_path.exists() else []
|
||||
old_by_key = work_by_key(old_filtered)
|
||||
new_by_key = work_by_key(filtered_works)
|
||||
|
||||
added: list[dict[str, Any]] = []
|
||||
removed: list[dict[str, Any]] = []
|
||||
emitted_added: set[str] = set()
|
||||
emitted_removed: set[str] = set()
|
||||
for work in filtered_works:
|
||||
key = primary_key(work)
|
||||
if key not in old_by_key and key not in emitted_added:
|
||||
added.append(work)
|
||||
emitted_added.add(key)
|
||||
for work in old_filtered:
|
||||
key = primary_key(work)
|
||||
if key not in new_by_key and key not in emitted_removed:
|
||||
removed.append(work)
|
||||
emitted_removed.add(key)
|
||||
return added, removed
|
||||
|
||||
|
||||
def embedding_diff_reasons(old: dict[str, Any], new: dict[str, Any]) -> list[str]:
|
||||
reasons: list[str] = []
|
||||
for field in ("title", "tags", "vas"):
|
||||
if old.get(field) != new.get(field):
|
||||
reasons.append(field)
|
||||
return reasons or ["embedding text"]
|
||||
|
||||
|
||||
def jsonl_record_count(path: Path) -> int:
|
||||
if not path.exists():
|
||||
return -1
|
||||
count = 0
|
||||
with path.open(encoding="utf-8") as file:
|
||||
for line in file:
|
||||
if line.strip():
|
||||
count += 1
|
||||
return count
|
||||
|
||||
|
||||
def web_data_out_of_sync(filtered_path: Path, data_dir: Path) -> bool:
|
||||
manifest_path = data_dir / "manifest.json"
|
||||
if not manifest_path.exists():
|
||||
return True
|
||||
try:
|
||||
manifest = read_json(manifest_path)
|
||||
web_count = int(manifest.get("count", -1))
|
||||
except (OSError, ValueError, TypeError, json.JSONDecodeError):
|
||||
return True
|
||||
return jsonl_record_count(filtered_path) != web_count
|
||||
|
||||
|
||||
def old_embedding_index_by_key(works: list[dict[str, Any]]) -> dict[str, int]:
|
||||
index_by_key: dict[str, int] = {}
|
||||
for index, work in enumerate(works):
|
||||
if not isinstance(work, dict):
|
||||
continue
|
||||
for key in identity_keys(work):
|
||||
index_by_key.setdefault(key, index)
|
||||
return index_by_key
|
||||
|
||||
|
||||
def existing_embedding_index(work: dict[str, Any], index_by_key: dict[str, int]) -> int | None:
|
||||
for key in identity_keys(work):
|
||||
index = index_by_key.get(key)
|
||||
if index is not None:
|
||||
return index
|
||||
return None
|
||||
|
||||
|
||||
def request_missing_embeddings(
|
||||
*,
|
||||
texts: list[str],
|
||||
start_indexes: list[int],
|
||||
manifest: dict[str, Any],
|
||||
args: argparse.Namespace,
|
||||
) -> dict[int, bytes]:
|
||||
if not texts:
|
||||
return {}
|
||||
|
||||
base_url = manifest.get("baseUrl")
|
||||
model = manifest.get("model")
|
||||
if not base_url or not model:
|
||||
raise RuntimeError("manifest must contain baseUrl and model for remote embeddings")
|
||||
|
||||
batch_size = int(manifest.get("batchSize") or 32)
|
||||
concurrency = int(manifest.get("concurrency") or 4)
|
||||
retry_forever = bool(manifest.get("retryForever", True))
|
||||
max_retry_wait = float(manifest.get("maxRetryWait") or 60.0)
|
||||
should_normalize = bool(manifest.get("normalized", True))
|
||||
url = embeddings_url(str(base_url))
|
||||
api_key = os.environ.get(args.api_key_env, "")
|
||||
packed_by_index: dict[int, bytes] = {}
|
||||
pending = {}
|
||||
completed: dict[int, list[list[float]]] = {}
|
||||
next_submit = 0
|
||||
next_write = 0
|
||||
dimensions = int(manifest["dimensions"])
|
||||
|
||||
def submit_available(executor: ThreadPoolExecutor) -> None:
|
||||
nonlocal next_submit
|
||||
while next_submit < len(texts) and len(pending) + len(completed) < concurrency:
|
||||
start = next_submit
|
||||
batch = texts[start : start + batch_size]
|
||||
future = executor.submit(
|
||||
request_embeddings_with_retries,
|
||||
url=url,
|
||||
model=str(model),
|
||||
inputs=batch,
|
||||
api_key=api_key,
|
||||
timeout=args.embedding_timeout,
|
||||
retries=args.embedding_retries,
|
||||
retry_wait=args.embedding_retry_wait,
|
||||
max_retry_wait=max_retry_wait,
|
||||
retry_forever=retry_forever,
|
||||
batch_start=start,
|
||||
)
|
||||
pending[future] = start
|
||||
next_submit += len(batch)
|
||||
|
||||
def pack_vectors(batch_start: int, vectors: list[list[float]]) -> None:
|
||||
for offset, vector in enumerate(vectors):
|
||||
if should_normalize:
|
||||
vector = normalize_vector(vector)
|
||||
if len(vector) != dimensions:
|
||||
raise RuntimeError(
|
||||
f"embedding dimensions mismatch: expected {dimensions}, got {len(vector)}"
|
||||
)
|
||||
packed_by_index[start_indexes[batch_start + offset]] = struct.pack(f"<{dimensions}f", *vector)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=concurrency) as executor:
|
||||
submit_available(executor)
|
||||
while pending:
|
||||
done, _ = wait(pending, return_when=FIRST_COMPLETED)
|
||||
for future in done:
|
||||
start = pending.pop(future)
|
||||
vectors = future.result()
|
||||
expected = min(batch_size, len(texts) - start)
|
||||
if len(vectors) != expected:
|
||||
raise RuntimeError(
|
||||
f"batch at {start} returned {len(vectors)} embeddings; expected {expected}"
|
||||
)
|
||||
completed[start] = vectors
|
||||
|
||||
while next_write in completed:
|
||||
vectors = completed.pop(next_write)
|
||||
pack_vectors(next_write, vectors)
|
||||
next_write += len(vectors)
|
||||
print(f"Embedded {next_write}/{len(texts)} missing works", file=sys.stderr, flush=True)
|
||||
|
||||
submit_available(executor)
|
||||
|
||||
return packed_by_index
|
||||
|
||||
|
||||
def write_embeddings_atomic(
|
||||
*,
|
||||
data_dir: Path,
|
||||
final_works: list[dict[str, Any]],
|
||||
old_works: list[dict[str, Any]],
|
||||
old_index_by_key: dict[str, int],
|
||||
old_embeddings_path: Path,
|
||||
missing_embeddings: dict[int, bytes],
|
||||
dimensions: int,
|
||||
) -> None:
|
||||
vector_size = dimensions * 4
|
||||
expected_size = len(old_works) * vector_size
|
||||
actual_size = old_embeddings_path.stat().st_size
|
||||
if actual_size != expected_size:
|
||||
raise RuntimeError(f"embedding size mismatch: expected {expected_size} bytes, got {actual_size} bytes")
|
||||
|
||||
temp_path = data_dir / "embeddings.f32.tmp"
|
||||
try:
|
||||
with old_embeddings_path.open("rb") as old_file, temp_path.open("wb") as new_file:
|
||||
with mmap.mmap(old_file.fileno(), 0, access=mmap.ACCESS_READ) as old_map:
|
||||
for index, work in enumerate(final_works):
|
||||
if index in missing_embeddings:
|
||||
new_file.write(missing_embeddings[index])
|
||||
continue
|
||||
old_index = existing_embedding_index(work, old_index_by_key)
|
||||
if old_index is None:
|
||||
raise RuntimeError(f"missing embedding for {primary_key(work)}")
|
||||
offset = old_index * vector_size
|
||||
new_file.write(old_map[offset : offset + vector_size])
|
||||
temp_path.replace(data_dir / "embeddings.f32")
|
||||
except Exception:
|
||||
temp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
|
||||
|
||||
def build_incremental_web_data(args: argparse.Namespace, filtered_works: list[dict[str, Any]]) -> dict[str, int]:
|
||||
data_dir = Path(args.data_dir)
|
||||
manifest_path = data_dir / "manifest.json"
|
||||
works_path = data_dir / "works.json"
|
||||
embeddings_path = data_dir / "embeddings.f32"
|
||||
manifest = read_json(manifest_path)
|
||||
if manifest.get("method") != "remote-openai-compatible":
|
||||
raise RuntimeError("incremental web/data update currently supports only remote-openai-compatible data")
|
||||
|
||||
old_works = read_json(works_path)
|
||||
if not isinstance(old_works, list):
|
||||
raise RuntimeError(f"{works_path} must contain a JSON array")
|
||||
|
||||
dimensions = int(manifest["dimensions"])
|
||||
old_index_by_key = old_embedding_index_by_key(old_works)
|
||||
final_works: list[dict[str, Any]] = []
|
||||
missing_texts: list[str] = []
|
||||
missing_indexes: list[int] = []
|
||||
reused = 0
|
||||
|
||||
for raw_work in filtered_works:
|
||||
work = normalize_work(raw_work, len(final_works))
|
||||
old_index = existing_embedding_index(work, old_index_by_key)
|
||||
if old_index is None or embedding_text(work) != embedding_text(old_works[old_index]):
|
||||
missing_indexes.append(len(final_works))
|
||||
missing_texts.append(embedding_text(work))
|
||||
else:
|
||||
reused += 1
|
||||
final_works.append(work)
|
||||
|
||||
missing_embeddings = request_missing_embeddings(
|
||||
texts=missing_texts,
|
||||
start_indexes=missing_indexes,
|
||||
manifest=manifest,
|
||||
args=args,
|
||||
)
|
||||
write_embeddings_atomic(
|
||||
data_dir=data_dir,
|
||||
final_works=final_works,
|
||||
old_works=old_works,
|
||||
old_index_by_key=old_index_by_key,
|
||||
old_embeddings_path=embeddings_path,
|
||||
missing_embeddings=missing_embeddings,
|
||||
dimensions=dimensions,
|
||||
)
|
||||
|
||||
new_manifest = dict(manifest)
|
||||
new_manifest.update(
|
||||
{
|
||||
"count": len(final_works),
|
||||
"dimensions": dimensions,
|
||||
"embeddingFile": "embeddings.f32",
|
||||
"worksFile": "works.json",
|
||||
"generatedAt": datetime.now(timezone.utc).isoformat(),
|
||||
}
|
||||
)
|
||||
write_json_atomic(works_path, final_works)
|
||||
write_json_atomic(manifest_path, new_manifest, indent=2)
|
||||
return {"reused": reused, "embedded": len(missing_indexes), "count": len(final_works)}
|
||||
|
||||
|
||||
def dry_run_web_stats(args: argparse.Namespace, filtered_works: list[dict[str, Any]]) -> dict[str, Any]:
|
||||
data_dir = Path(args.data_dir)
|
||||
manifest = read_json(data_dir / "manifest.json")
|
||||
if manifest.get("method") != "remote-openai-compatible":
|
||||
raise RuntimeError("incremental web/data update currently supports only remote-openai-compatible data")
|
||||
old_works = read_json(data_dir / "works.json")
|
||||
old_index_by_key = old_embedding_index_by_key(old_works)
|
||||
reused = 0
|
||||
embedded = 0
|
||||
embedding_diffs: list[dict[str, Any]] = []
|
||||
for raw_work in filtered_works:
|
||||
work = normalize_work(raw_work, reused + embedded)
|
||||
old_index = existing_embedding_index(work, old_index_by_key)
|
||||
if old_index is None:
|
||||
embedded += 1
|
||||
embedding_diffs.append({"work": work, "reason": "new work"})
|
||||
elif embedding_text(work) != embedding_text(old_works[old_index]):
|
||||
embedded += 1
|
||||
reasons = ", ".join(embedding_diff_reasons(old_works[old_index], work))
|
||||
embedding_diffs.append({"work": work, "reason": f"changed {reasons}"})
|
||||
else:
|
||||
reused += 1
|
||||
return {"reused": reused, "embedded": embedded, "count": reused + embedded, "embeddingDiffs": embedding_diffs}
|
||||
|
||||
|
||||
def limited_rows(rows: list[Any], limit: int) -> list[Any]:
|
||||
return rows if limit == 0 else rows[:limit]
|
||||
|
||||
|
||||
def remaining_count(rows: list[Any], limit: int) -> int:
|
||||
return 0 if limit == 0 else max(0, len(rows) - limit)
|
||||
|
||||
|
||||
def print_work_section(title: str, prefix: str, works: list[dict[str, Any]], limit: int) -> None:
|
||||
print(f"\n{title}: {len(works)}")
|
||||
for work in limited_rows(works, limit):
|
||||
print(f" {prefix} {format_work(work)}")
|
||||
remaining = remaining_count(works, limit)
|
||||
if remaining:
|
||||
print(f" ... {remaining} more")
|
||||
|
||||
|
||||
def print_updated_section(diffs: list[dict[str, Any]], limit: int) -> None:
|
||||
print(f"\nUpdated existing raw works: {len(diffs)}")
|
||||
for diff in limited_rows(diffs, limit):
|
||||
print(f" ~ {format_work(diff['work'])} / {', '.join(diff['reasons'])}")
|
||||
remaining = remaining_count(diffs, limit)
|
||||
if remaining:
|
||||
print(f" ... {remaining} more")
|
||||
|
||||
|
||||
def print_embedding_section(diffs: list[dict[str, Any]], limit: int) -> None:
|
||||
print(f"\nEmbedding requests: {len(diffs)}")
|
||||
for diff in limited_rows(diffs, limit):
|
||||
marker = "+" if diff["reason"] == "new work" else "~"
|
||||
print(f" {marker} {format_work(diff['work'])} / {diff['reason']}")
|
||||
remaining = remaining_count(diffs, limit)
|
||||
if remaining:
|
||||
print(f" ... {remaining} more")
|
||||
|
||||
|
||||
def print_dry_run_diff(
|
||||
*,
|
||||
args: argparse.Namespace,
|
||||
new_works: list[dict[str, Any]],
|
||||
existing_works: list[dict[str, Any]],
|
||||
fetched_existing_by_key: dict[str, dict[str, Any]],
|
||||
filtered_path: Path,
|
||||
filtered_works: list[dict[str, Any]],
|
||||
web_stats: dict[str, Any],
|
||||
) -> None:
|
||||
limit = args.diff_limit
|
||||
updated_diffs = updated_existing_diffs(existing_works, fetched_existing_by_key)
|
||||
filtered_added, filtered_removed = filtered_change_diffs(filtered_path, filtered_works)
|
||||
|
||||
print_work_section("New raw works", "+", new_works, limit)
|
||||
print_updated_section(updated_diffs, limit)
|
||||
print_work_section("Filtered additions", "+", filtered_added, limit)
|
||||
print_work_section("Filtered removals", "-", filtered_removed, limit)
|
||||
print_embedding_section(web_stats.get("embeddingDiffs", []), limit)
|
||||
|
||||
|
||||
def validate_args(args: argparse.Namespace) -> None:
|
||||
if args.page_size <= 0:
|
||||
raise SystemExit("--page-size must be greater than 0")
|
||||
if args.timeout <= 0:
|
||||
raise SystemExit("--timeout must be greater than 0")
|
||||
if args.retries < 0:
|
||||
raise SystemExit("--retries must be greater than or equal to 0")
|
||||
if args.max_pages is not None and args.max_pages <= 0:
|
||||
raise SystemExit("--max-pages must be greater than 0")
|
||||
if args.diff_limit < 0:
|
||||
raise SystemExit("--diff-limit must be greater than or equal to 0")
|
||||
if args.embedding_timeout <= 0:
|
||||
raise SystemExit("--embedding-timeout must be greater than 0")
|
||||
if args.embedding_retries < 0:
|
||||
raise SystemExit("--embedding-retries must be greater than or equal to 0")
|
||||
if args.embedding_retry_wait < 0:
|
||||
raise SystemExit("--embedding-retry-wait must be greater than or equal to 0")
|
||||
|
||||
|
||||
def main() -> int:
|
||||
args = parse_args()
|
||||
validate_args(args)
|
||||
input_path = Path(args.input)
|
||||
filtered_path = Path(args.filtered_output)
|
||||
if not input_path.exists():
|
||||
raise SystemExit(f"input file not found: {input_path}")
|
||||
|
||||
existing_works = read_jsonl(input_path)
|
||||
existing_keys = {key for work in existing_works for key in identity_keys(work)}
|
||||
new_works, fetched_existing_by_key, scanned_pages = fetch_latest_works(args, existing_keys)
|
||||
|
||||
if not new_works and not args.force and not web_data_out_of_sync(filtered_path, Path(args.data_dir)):
|
||||
print(f"Scanned pages: {scanned_pages}")
|
||||
print("No new works found. web/data update skipped.")
|
||||
return 0
|
||||
|
||||
merged_works = merge_raw_works(existing_works, new_works, fetched_existing_by_key)
|
||||
filtered_works, filter_stats = filter_raw_works(args, merged_works)
|
||||
web_stats = dry_run_web_stats(args, filtered_works) if args.dry_run else None
|
||||
|
||||
print(f"Scanned pages: {scanned_pages}")
|
||||
print(f"New raw works: {len(new_works)}")
|
||||
print(f"Raw works after merge: {len(merged_works)}")
|
||||
print(f"Filtered works: {filter_stats['output']} ({filter_stats['total_removed']} removed)")
|
||||
if web_stats:
|
||||
print(f"web/data would reuse {web_stats['reused']} embeddings and request {web_stats['embedded']} embeddings")
|
||||
print_dry_run_diff(
|
||||
args=args,
|
||||
new_works=new_works,
|
||||
existing_works=existing_works,
|
||||
fetched_existing_by_key=fetched_existing_by_key,
|
||||
filtered_path=filtered_path,
|
||||
filtered_works=filtered_works,
|
||||
web_stats=web_stats,
|
||||
)
|
||||
|
||||
if args.dry_run:
|
||||
print("Dry run: no files were written.")
|
||||
return 0
|
||||
|
||||
write_jsonl_atomic(input_path, merged_works)
|
||||
write_jsonl_atomic(filtered_path, filtered_works)
|
||||
web_stats = build_incremental_web_data(args, filtered_works)
|
||||
print(f"Updated {input_path}")
|
||||
print(f"Updated {filtered_path}")
|
||||
print(
|
||||
f"Updated {args.data_dir}: {web_stats['count']} works, "
|
||||
f"reused {web_stats['reused']} embeddings, requested {web_stats['embedded']} embeddings"
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user