Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions python/fasttext_module/fasttext/tests/test_quantize_twice.py
Original file line number Diff line number Diff line change
@@ -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()
9 changes: 8 additions & 1 deletion src/fasttext.cc
Original file line number Diff line number Diff line change
Expand Up @@ -86,6 +86,7 @@ namespace fasttext

input_ = std::dynamic_pointer_cast<Matrix>(inputMatrix);
output_ = std::dynamic_pointer_cast<Matrix>(outputMatrix);
quant_ = false;
wordVectors_.reset();
args_->dim = input_->size(1);

Expand Down Expand Up @@ -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<QuantMatrix>();
}
input_->load(in);
Expand Down Expand Up @@ -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(
Expand Down
Loading