From 8378e7c627e4e9bf4db0f4352e11c5173f699b3a Mon Sep 17 00:00:00 2001 From: stephantul Date: Fri, 28 Aug 2026 16:41:06 +0200 Subject: [PATCH] fix: unk token for unigrams --- model2vec/model.py | 20 +++++++++++++++----- 1 file changed, 15 insertions(+), 5 deletions(-) diff --git a/model2vec/model.py b/model2vec/model.py index e477c3c1..3a632066 100644 --- a/model2vec/model.py +++ b/model2vec/model.py @@ -1,5 +1,6 @@ from __future__ import annotations +import json import math import os import warnings @@ -67,11 +68,7 @@ def __init__( self.token_mapping: np.ndarray | None = token_mapping self.tokenizer = tokenizer - self.unk_token_id: int | None - if hasattr(self.tokenizer.model, "unk_token") and self.tokenizer.model.unk_token is not None: - self.unk_token_id = tokenizer.get_vocab()[self.tokenizer.model.unk_token] - else: - self.unk_token_id = None # pragma: no cover # Doesn't actually happen, but can happen. + self.unk_token_id = _get_unk_token_id(self.tokenizer) self.median_token_length = int(np.median([len(token) for token in self.tokens])) self.config = config or {} @@ -559,3 +556,16 @@ def _loading_helper( quantize_to=quantize_to, dimensionality=dimensionality, ) + + +def _get_unk_token_id(tokenizer: Tokenizer) -> int | None: + """Get the unk token id.""" + model = tokenizer.model + # Wordpiece + BPE + Word level + if hasattr(model, "unk_token"): + token = model.unk_token + if token is None: + return None + return tokenizer.token_to_id(token) + # Unigram + return json.loads(tokenizer.to_str())["model"].get("unk_id")