290 lines
9.6 KiB
Python
290 lines
9.6 KiB
Python
#!/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())
|