limoXD's picture
Release v0.2 Mozc-backed Japanese IME reranker
f11438f verified
Raw
History Blame Contribute Delete
6.93 kB
from __future__ import annotations
from pathlib import Path
from typing import Any
import torch
from transformers import AutoModelForMaskedLM, AutoTokenizer
from .domain import CandidateScoringError, RerankRequest
MODEL_ID = "ku-nlp/deberta-v2-tiny-japanese"
MODEL_REVISION = "0427645e8cf44ee83b4a0b5f4498274d89e02adb"
class DebertaCandidateScorer:
"""Score only the candidate span with masked-LM pseudo-log-likelihood."""
def __init__(
self,
*,
model_id: str = MODEL_ID,
revision: str = MODEL_REVISION,
device: str = "cpu",
cache_dir: Path | None = None,
local_files_only: bool = False,
max_length: int = 256,
max_mask_batch: int = 16,
) -> None:
self._device = torch.device(device)
self._model_id = model_id
self._load_options: dict[str, Any] = {
"revision": revision,
"local_files_only": local_files_only,
}
if cache_dir is not None:
self._load_options["cache_dir"] = str(cache_dir)
self._tokenizer: Any | None = None
self._model: Any | None = None
self._requested_max_length = max_length
self._max_length = max_length
self._max_mask_batch = max_mask_batch
def load(self) -> None:
"""Load and pin model assets before latency measurement."""
self._ensure_loaded()
@torch.inference_mode()
def score_candidates(self, request: RerankRequest) -> list[float]:
tokenizer, model = self._loaded_components()
input_rows: list[torch.Tensor] = []
attention_rows: list[torch.Tensor] = []
positions: list[int] = []
target_ids: list[int] = []
candidate_indexes: list[int] = []
for candidate_index, candidate in enumerate(request.candidates):
encoding, target_index = self._encode_candidate(
tokenizer,
request.left_context,
request.right_context,
candidate.surface,
)
word_ids = encoding.word_ids(batch_index=0)
target_positions = [
position for position, word_id in enumerate(word_ids) if word_id == target_index
]
if not target_positions:
raise CandidateScoringError(
"candidate_outside_window",
"candidate was outside the model token window",
)
for position in target_positions:
masked_ids = encoding["input_ids"][0].clone()
target_id = int(masked_ids[position])
if target_id == tokenizer.unk_token_id:
raise CandidateScoringError(
"unknown_token",
"candidate contains an unknown model token",
)
if target_id in tokenizer.all_special_ids:
raise CandidateScoringError(
"special_token",
"candidate resolves to a special model token",
)
masked_ids[position] = tokenizer.mask_token_id
input_rows.append(masked_ids)
attention_rows.append(encoding["attention_mask"][0])
positions.append(position)
target_ids.append(target_id)
candidate_indexes.append(candidate_index)
sums = [0.0] * len(request.candidates)
counts = [0] * len(request.candidates)
for start in range(0, len(input_rows), self._max_mask_batch):
stop = start + self._max_mask_batch
ids = torch.nn.utils.rnn.pad_sequence(
input_rows[start:stop],
batch_first=True,
padding_value=tokenizer.pad_token_id,
).to(self._device)
attention = torch.nn.utils.rnn.pad_sequence(
attention_rows[start:stop], batch_first=True, padding_value=0
).to(self._device)
logits = model(input_ids=ids, attention_mask=attention).logits
row_count = logits.shape[0]
row_indexes = torch.arange(row_count, device=self._device)
chunk_positions = torch.tensor(
positions[start:stop], dtype=torch.long, device=self._device
)
chunk_targets = torch.tensor(
target_ids[start:stop], dtype=torch.long, device=self._device
)
target_logits = logits[row_indexes, chunk_positions]
log_probs = torch.log_softmax(target_logits, dim=-1)
values = log_probs[row_indexes, chunk_targets].cpu().tolist()
for offset, value in enumerate(values):
candidate_index = candidate_indexes[start + offset]
sums[candidate_index] += float(value)
counts[candidate_index] += 1
if any(count == 0 for count in counts):
raise CandidateScoringError(
"candidate_unscored",
"at least one candidate could not be scored",
)
return [total / count for total, count in zip(sums, counts, strict=True)]
def _ensure_loaded(self) -> None:
if self._tokenizer is not None and self._model is not None:
return
tokenizer = AutoTokenizer.from_pretrained(self._model_id, **self._load_options)
model = AutoModelForMaskedLM.from_pretrained(self._model_id, **self._load_options)
model.to(self._device)
model.eval()
model_limit = int(getattr(model.config, "max_position_embeddings", 512))
self._tokenizer = tokenizer
self._model = model
self._max_length = min(self._requested_max_length, model_limit)
def _loaded_components(self) -> tuple[Any, Any]:
self._ensure_loaded()
assert self._tokenizer is not None
assert self._model is not None
return self._tokenizer, self._model
def _encode_candidate(
self,
tokenizer: Any,
left_context: tuple[str, ...],
right_context: tuple[str, ...],
surface: str,
) -> tuple[Any, int]:
left = list(left_context)
right = list(right_context)
while True:
words = [*left, surface, *right]
target_index = len(left)
encoding = tokenizer(
words,
is_split_into_words=True,
return_tensors="pt",
)
if encoding["input_ids"].shape[1] <= self._max_length:
return encoding, target_index
if not left and not right:
raise CandidateScoringError(
"candidate_too_long",
"candidate is longer than the model token window",
)
if len(left) >= len(right) and left:
left.pop(0)
elif right:
right.pop()