circle-search-26-api / scripts /6_build_word_embeddings.py
ktsn-ud
不要なインポートを削除
d72bce3
Raw
History Blame Contribute Delete
4.53 kB
import os
import sys
import json
from typing import Dict, List, Tuple
import numpy as np
sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
from utils.logger import setup_logger
from utils.json import field_getter, json_dumps
log = setup_logger(__name__)
def load_configs():
files = field_getter("config/files.json")
paths = {
"tf_token": files("bm25.tf_token"),
"fasttext_vec": files("embeddings.fasttext_vec"),
"word_vocab": files("embeddings.word_vocab"),
"word_vectors": files("embeddings.word_vectors"),
}
return paths
def read_vocab_from_tf_token(tf_token_path: str) -> Tuple[List[dict], set[str]]:
with open(tf_token_path, encoding="utf-8") as f:
docs = json.load(f)
vocab: set[str] = set()
for d in docs:
# doc全体tfから語彙を得る
for t in (d.get("tf") or {}).keys():
vocab.add(t)
return docs, vocab
def stream_fasttext_vec(
vec_path: str, vocab: set[str]
) -> Tuple[Dict[str, int], np.ndarray]:
"""fastText .vec から、必要語彙のみ抽出してベクトル行列を返す。"""
token_to_idx: Dict[str, int] = {}
vectors: List[np.ndarray] = []
dim = None
kept = 0
with open(vec_path, encoding="utf-8", errors="ignore") as f:
header = f.readline()
# ヘッダ行は "<count> <dim>" のことが多い
try:
parts = header.strip().split()
if len(parts) >= 2 and parts[0].isdigit():
dim = int(parts[1])
except Exception:
pass
for line in f:
sp = line.rstrip().split(" ")
if len(sp) < 2:
continue
token = sp[0]
if token not in vocab:
continue
vec_vals = sp[1:]
if dim is None:
dim = len(vec_vals)
if len(vec_vals) != dim:
continue
try:
v = np.asarray([float(x) for x in vec_vals], dtype=np.float32)
except ValueError:
continue
# L2正規化
norm = np.linalg.norm(v)
if norm > 0:
v = v / norm
token_to_idx[token] = kept
vectors.append(v)
kept += 1
if dim is None:
raise RuntimeError(".vec の次元を特定できませんでした")
if not vectors:
log.warning("語彙に一致するベクトルが見つかりませんでした")
arr = np.zeros((0, dim), dtype=np.float32)
return token_to_idx, arr
arr = np.vstack(vectors).astype(np.float32)
return token_to_idx, arr
# top-k モードでは文書ベクトルは不要
def main():
log.info(
"語彙/文書ベクトル(word_vocab.json, word_vectors.npz, doc_vectors.npy)を生成します"
)
try:
paths = load_configs()
except Exception as e:
log.error(f"設定の読み込みに失敗しました: {e}")
sys.exit(1)
tf_token_path = paths["tf_token"]
fasttext_vec_path = paths["fasttext_vec"]
if not os.path.exists(fasttext_vec_path):
log.warning(f".vec が見つかりません: {fasttext_vec_path}")
log.warning("Step 6 をスキップします")
sys.exit(0)
try:
docs, vocab = read_vocab_from_tf_token(tf_token_path)
except Exception as e:
log.error(f"tf_token.jsonの読み込みに失敗しました: {e}")
sys.exit(1)
log.info(f"コーパス語彙数: {len(vocab)}")
token_to_idx, word_vecs = stream_fasttext_vec(fasttext_vec_path, vocab)
log.info(f"抽出済み語彙ベクトル数: {word_vecs.shape[0]}")
# 語彙インデックスの安定化(token_to_idxは追加順次第なのでソート)
sorted_tokens = sorted(token_to_idx.keys())
remap = {t: i for i, t in enumerate(sorted_tokens)}
remapped_vecs = np.zeros_like(word_vecs)
for t, old_i in token_to_idx.items():
new_i = remap[t]
remapped_vecs[new_i] = word_vecs[old_i]
token_to_idx = remap
word_vecs = remapped_vecs
# 出力
os.makedirs(os.path.dirname(paths["word_vocab"]), exist_ok=True)
json_dumps(token_to_idx, paths["word_vocab"]) # 語→index
# 圧縮npz
np.savez_compressed(paths["word_vectors"], vectors=word_vecs)
log.info(f"word_vocab.json: {paths['word_vocab']}")
log.info(f"word_vectors.npz: {paths['word_vectors']}")
if __name__ == "__main__":
main()