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
43 changes: 42 additions & 1 deletion tests/test_word_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,11 +4,52 @@
import numpy as np
import pytest

from textattack.shared import GensimWordEmbedding, WordEmbedding
from textattack.constraints.semantics import WordEmbeddingDistance
from textattack.shared import AttackedText, GensimWordEmbedding, WordEmbedding

_gensim_available = importlib.util.find_spec("gensim") is not None


@pytest.fixture
def small_embedding():
return WordEmbedding(
np.array([[1.0, 0.0], [-1.0, 0.0], [1.0, 1.0]]),
{"source": 0, "opposite": 1, "different": 2},
{0: "source", 1: "opposite", 2: "different"},
)


def test_word_embedding_distance_accepts_zero_thresholds(small_embedding):
WordEmbeddingDistance(embedding=small_embedding, min_cos_sim=0.0)
WordEmbeddingDistance(embedding=small_embedding, max_mse_dist=0.0)


@pytest.mark.parametrize(
("min_cos_sim", "max_mse_dist"),
[(None, None), (0.0, 0.0), (0.5, 1.0)],
)
def test_word_embedding_distance_requires_exactly_one_threshold(
small_embedding, min_cos_sim, max_mse_dist
):
with pytest.raises(ValueError):
WordEmbeddingDistance(
embedding=small_embedding,
min_cos_sim=min_cos_sim,
max_mse_dist=max_mse_dist,
)


def test_word_embedding_distance_enforces_zero_thresholds(small_embedding):
reference = AttackedText("source")

assert not WordEmbeddingDistance(
embedding=small_embedding, min_cos_sim=0.0
)._check_constraint(reference.generate_new_attacked_text(["opposite"]), reference)
assert not WordEmbeddingDistance(
embedding=small_embedding, max_mse_dist=0.0
)._check_constraint(reference.generate_new_attacked_text(["different"]), reference)


def test_embedding_paragramcf():
word_embedding = WordEmbedding.counterfitted_GLOVE_embedding()
assert pytest.approx(word_embedding[0][0]) == -0.022007
Expand Down
6 changes: 3 additions & 3 deletions textattack/constraints/semantics/word_embedding_distance.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ def __init__(
self.include_unknown_words = include_unknown_words
self.cased = cased

if bool(min_cos_sim) == bool(max_mse_dist):
if (min_cos_sim is None) == (max_mse_dist is None):
raise ValueError("You must choose either `min_cos_sim` or `max_mse_dist`.")
self.min_cos_sim = min_cos_sim
self.max_mse_dist = max_mse_dist
Expand Down Expand Up @@ -92,12 +92,12 @@ def _check_constraint(self, transformed_text, reference_text):
return False

# Check cosine distance.
if self.min_cos_sim:
if self.min_cos_sim is not None:
cos_sim = self.get_cos_sim(ref_id, transformed_id)
if cos_sim < self.min_cos_sim:
return False
# Check MSE distance.
if self.max_mse_dist:
if self.max_mse_dist is not None:
mse_dist = self.get_mse_dist(ref_id, transformed_id)
if mse_dist > self.max_mse_dist:
return False
Expand Down