Files
asmr-vector-search/scripts/prune_duplicate_works.py
2026-06-11 03:43:59 +09:00

328 lines
11 KiB
Python

#!/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())