diff --git a/python/fasttext_module/fasttext/tests/test_quantize_twice.py b/python/fasttext_module/fasttext/tests/test_quantize_twice.py new file mode 100644 index 0000000..d466821 --- /dev/null +++ b/python/fasttext_module/fasttext/tests/test_quantize_twice.py @@ -0,0 +1,42 @@ +# SPDX-FileContributor: Arthit Suriyawongkul +# SPDX-FileCopyrightText: 2026-present, fasttext-community +# SPDX-FileType: SOURCE +# SPDX-License-Identifier: MIT + +"""quantize() must raise on a quantized model, and only on one.""" + +import pytest + +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 test_quantize_twice_raises(): + model = _model() + model.quantize() + with pytest.raises(ValueError, match="already quantized"): + model.quantize() + + +def test_load_model_clears_quantized(tmp_path): + model = _model() + path = str(tmp_path / "model.bin") + model.save_model(path) + model.quantize() + model.f.loadModel(path) + assert not model.is_quantized() + model.quantize() + + +def test_set_matrices_clears_quantized(): + model = _model() + matrices = model.get_input_matrix(), model.get_output_matrix() + model.quantize() + model.set_matrices(*matrices) + assert not model.is_quantized() + model.quantize() diff --git a/src/fasttext.cc b/src/fasttext.cc index a659c55..f54fdde 100644 --- a/src/fasttext.cc +++ b/src/fasttext.cc @@ -86,6 +86,7 @@ namespace fasttext input_ = std::dynamic_pointer_cast(inputMatrix); output_ = std::dynamic_pointer_cast(outputMatrix); + quant_ = false; wordVectors_.reset(); args_->dim = input_->size(1); @@ -306,9 +307,9 @@ namespace fasttext bool quant_input; in.read((char *)&quant_input, sizeof(bool)); + quant_ = quant_input; if (quant_input) { - quant_ = true; input_ = std::make_shared(); } input_->load(in); @@ -386,6 +387,12 @@ namespace fasttext void FastText::quantize(const Args &qargs, const TrainCallback &callback) { + if (quant_) + { + throw std::invalid_argument( + "Model is already quantized. " + "Quantize the original (non-quantized) model instead."); + } if (args_->model != model_name::sup) { throw std::invalid_argument(