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
54 changes: 54 additions & 0 deletions python/fasttext_module/fasttext/tests/test_vector_bounds.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,54 @@
# SPDX-FileContributor: Arthit Suriyawongkul
# SPDX-FileCopyrightText: 2026-present, fasttext-community
# SPDX-FileType: SOURCE
# SPDX-License-Identifier: MIT

"""Vector index and size errors must raise, not access out of bounds."""

import pytest

import fasttext_pybind

from .helpers import build_supervised_model, get_random_data


def _model(data=None):
# thread=12: thread <= 10 leaves the input matrix partly uninitialized.
data = data or get_random_data(300, max_vocab_size=100)
return build_supervised_model(data, {"thread": 12, "dim": 16, "verbose": 0})


def _shrink_output_dim(model):
output = model.get_output_matrix()[:, :8].copy() # input dim is 16
model.set_matrices(model.get_input_matrix(), output)


@pytest.mark.parametrize("ind", [-1, 10**9])
def test_get_input_vector_out_of_range_raises(ind):
with pytest.raises(ValueError, match="out of range"):
_model().get_input_vector(ind)


def test_get_word_vector_size_mismatch_raises():
model = _model()
vec = fasttext_pybind.Vector(model.get_dimension() - 1)
with pytest.raises(ValueError, match="size mismatch"):
model.f.getWordVector(vec, model.words[1])


def test_predict_dim_mismatch_raises():
model = _model()
_shrink_output_dim(model)
with pytest.raises(ValueError, match="size mismatch"):
model.predict(model.words[1])


def test_training_thread_error_raises(tmp_path):
"""Used to call std::terminate: exception escaped a training thread."""
data = get_random_data(3000, max_vocab_size=600) # cutoff needs >= 256 rows
model = _model(data)
_shrink_output_dim(model)
train_txt = tmp_path / "train.txt"
train_txt.write_text("".join(f"__label__{line}\n" for line in data))
with pytest.raises(ValueError, match="size mismatch"):
model.quantize(input=str(train_txt), cutoff=300, retrain=True)
6 changes: 6 additions & 0 deletions src/fasttext.cc
Original file line number Diff line number Diff line change
Expand Up @@ -806,6 +806,12 @@ namespace fasttext
{
trainException_ = std::current_exception();
}
catch (const std::exception &)
{
// E.g. a Vector size check. An exception escaping a std::thread
// calls std::terminate(); store it so startThreads() rethrows it.
trainException_ = std::current_exception();
}
if (threadId == 0)
loss_ = state.getLoss();
ifs.close();
Expand Down
34 changes: 22 additions & 12 deletions src/vector.cc
Original file line number Diff line number Diff line change
Expand Up @@ -8,10 +8,9 @@

#include "vector.h"

#include <assert.h>

#include <cmath>
#include <iomanip>
#include <stdexcept>

#include "matrix.h"

Expand All @@ -38,36 +37,47 @@ void Vector::mul(real a) {
}

void Vector::addVector(const Vector& source) {
assert(size() == source.size());
if (size() != source.size()) {
throw std::invalid_argument("Vector::addVector: size mismatch");
}
for (int64_t i = 0; i < size(); i++) {
data_[i] += source.data_[i];
}
}

void Vector::addVector(const Vector& source, real s) {
assert(size() == source.size());
if (size() != source.size()) {
throw std::invalid_argument("Vector::addVector: size mismatch");
}
for (int64_t i = 0; i < size(); i++) {
data_[i] += s * source.data_[i];
}
}

void Vector::addRow(const Matrix& A, int64_t i, real a) {
assert(i >= 0);
assert(i < A.size(0));
assert(size() == A.size(1));
if (i < 0 || i >= A.size(0)) {
throw std::invalid_argument("Vector::addRow: row index out of range");
}
if (size() != A.size(1)) {
throw std::invalid_argument("Vector::addRow: size mismatch");
}
A.addRowToVector(*this, i, a);
}

void Vector::addRow(const Matrix& A, int64_t i) {
assert(i >= 0);
assert(i < A.size(0));
assert(size() == A.size(1));
if (i < 0 || i >= A.size(0)) {
throw std::invalid_argument("Vector::addRow: row index out of range");
}
if (size() != A.size(1)) {
throw std::invalid_argument("Vector::addRow: size mismatch");
}
A.addRowToVector(*this, i);
}

void Vector::mul(const Matrix& A, const Vector& vec) {
assert(A.size(0) == size());
assert(A.size(1) == vec.size());
if (A.size(0) != size() || A.size(1) != vec.size()) {
throw std::invalid_argument("Vector::mul: size mismatch");
}
for (int64_t i = 0; i < size(); i++) {
data_[i] = A.dotRow(vec, i);
}
Expand Down
Loading