From 653fb0915b17d581124922d6a44bd3e8e0ef7080 Mon Sep 17 00:00:00 2001 From: Yifan Chen Date: Mon, 7 Sep 2026 07:31:43 -0700 Subject: [PATCH] fix: accept zero word embedding distance thresholds --- tests/test_word_embedding.py | 43 ++++++++++++++++++- .../semantics/word_embedding_distance.py | 6 +-- 2 files changed, 45 insertions(+), 4 deletions(-) diff --git a/tests/test_word_embedding.py b/tests/test_word_embedding.py index 69523260a..4d5dec738 100644 --- a/tests/test_word_embedding.py +++ b/tests/test_word_embedding.py @@ -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 diff --git a/textattack/constraints/semantics/word_embedding_distance.py b/textattack/constraints/semantics/word_embedding_distance.py index 02a51f19c..f5c5f5328 100644 --- a/textattack/constraints/semantics/word_embedding_distance.py +++ b/textattack/constraints/semantics/word_embedding_distance.py @@ -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 @@ -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