diff --git a/python/fasttext_module/fasttext/FastText.py b/python/fasttext_module/fasttext/FastText.py index d4081b4..199d2c4 100644 --- a/python/fasttext_module/fasttext/FastText.py +++ b/python/fasttext_module/fasttext/FastText.py @@ -363,6 +363,8 @@ def quantize( self.f.quantize( input, qout, cutoff, retrain, epoch, lr, thread, verbose, dsub, qnorm ) + # cutoff prunes the dictionary + self._words = None def set_matrices(self, input_matrix, output_matrix): """ diff --git a/python/fasttext_module/fasttext/tests/test_nn_cache.py b/python/fasttext_module/fasttext/tests/test_nn_cache.py new file mode 100644 index 0000000..5fe95f1 --- /dev/null +++ b/python/fasttext_module/fasttext/tests/test_nn_cache.py @@ -0,0 +1,64 @@ +# SPDX-FileContributor: Arthit Suriyawongkul +# SPDX-FileCopyrightText: 2026-present, fasttext-community +# SPDX-FileType: SOURCE +# SPDX-License-Identifier: MIT + +"""Cached word data must be rebuilt when the model changes.""" + +import pytest + +import fasttext +import fasttext_pybind + +from .helpers import build_supervised_model, get_random_data + + +def _model(): + # thread=12: thread <= 10 leaves the input matrix partly uninitialized. + data = get_random_data(3000, max_vocab_size=600) + return build_supervised_model(data, {"thread": 12, "dim": 16, "verbose": 0}) + + +def _filled_model(): + model = _model() + model.get_nearest_neighbors(model.words[1]) # fills the caches + return model + + +def _reload(model, tmp_path): + path = str(tmp_path / "model.bin") + model.save_model(path) + return fasttext.load_model(path) + + +def _assert_same_nn(model, expected): + word = expected.words[1] # words[0] is the end-of-sentence token + assert model.get_nearest_neighbors(word) == expected.get_nearest_neighbors(word) + + +# cutoff=300 also prunes the dictionary (quantize needs >= 256 rows). +@pytest.mark.parametrize("cutoff", [0, 300]) +def test_quantize_resets_caches(tmp_path, cutoff): + model = _filled_model() + model.quantize(cutoff=cutoff) + expected = _reload(model, tmp_path) + assert model.words == expected.words + _assert_same_nn(model, expected) + + +def test_load_model_resets_cache(tmp_path): + other = _reload(_model(), tmp_path) + model = _filled_model() + model.f.loadModel(str(tmp_path / "model.bin")) + _assert_same_nn(model, other) + + +def test_train_resets_cache(tmp_path): + model = _filled_model() + train_txt = tmp_path / "train.txt" + data = get_random_data(3000) + train_txt.write_text("".join(f"__label__{line}\n" for line in data)) + args = model.f.getArgs() + args.input = str(train_txt) + fasttext_pybind.train(model.f, args) + _assert_same_nn(model, _reload(model, tmp_path)) diff --git a/src/fasttext.cc b/src/fasttext.cc index a659c55..dfad387 100644 --- a/src/fasttext.cc +++ b/src/fasttext.cc @@ -293,6 +293,7 @@ namespace fasttext void FastText::loadModel(std::istream &in) { + wordVectors_.reset(); args_ = std::make_shared(); input_ = std::make_shared(); output_ = std::make_shared(); @@ -400,6 +401,10 @@ namespace fasttext std::dynamic_pointer_cast(output_); bool normalizeGradient = (args_->model == model_name::sup); + // input_ is replaced below (and dict_ pruned if cutoff > 0), so the + // cached word vectors used by getNN/getAnalogies are stale either way. + wordVectors_.reset(); + if (qargs.cutoff > 0 && static_cast(qargs.cutoff) < static_cast(input->size(0))) { auto idx = selectEmbeddings(qargs.cutoff); @@ -892,6 +897,7 @@ namespace fasttext "bucket must be > 0 when using subwords (maxn > 0) " "or word n-grams (wordNgrams > 1)"); } + wordVectors_.reset(); args_ = std::make_shared(args); dict_ = std::make_shared(args_); if (args_->input == "-")