Skip to content
114 changes: 114 additions & 0 deletions tests/test_attack_device_placement.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,114 @@
def _build_attack(model, tokenizer, transformation):
from textattack import Attack
from textattack.constraints.pre_transformation import (
RepeatModification,
StopwordModification,
)
from textattack.goal_functions import UntargetedClassification
from textattack.models.wrappers import HuggingFaceModelWrapper
from textattack.search_methods import GreedyWordSwapWIR

wrapper = HuggingFaceModelWrapper(model, tokenizer)
goal_function = UntargetedClassification(wrapper)
return Attack(
goal_function,
[RepeatModification(), StopwordModification()],
transformation,
GreedyWordSwapWIR(),
)


def _model_and_tokenizer():
import transformers

model = transformers.AutoModelForSequenceClassification.from_pretrained(
"hf-internal-testing/tiny-random-bert"
)
tokenizer = transformers.AutoTokenizer.from_pretrained(
"hf-internal-testing/tiny-random-bert"
)
return model, tokenizer


def test_cuda_skips_hf_device_map_model_but_moves_unrelated_module():
# Regression test: the `hf_device_map` skip in `Attack.cuda_`/`to_cuda`
# used to apply to any `torch.nn.Module` with a truthy `hf_device_map`
# attribute, not just `transformers.PreTrainedModel` instances. Since
# this visitor also traverses non-HuggingFace modules reachable from a
# Constraint/GoalFunction/Transformation, a coincidental attribute name
# collision would silently skip moving that module to the configured
# device. Confirm the HF model is still (correctly) skipped, while an
# unrelated module with the same attribute name is not.
from unittest.mock import patch

import torch

from textattack.transformations import WordSwapRandomCharacterDeletion

model, tokenizer = _model_and_tokenizer()
model.hf_device_map = {"": "cpu"}

transformation = WordSwapRandomCharacterDeletion()
transformation.some_unrelated_module = torch.nn.Linear(2, 2)
transformation.some_unrelated_module.hf_device_map = {"": "cpu"}

attack = _build_attack(model, tokenizer, transformation)

with (
patch.object(model, "to", wraps=model.to) as model_to_spy,
patch.object(
transformation.some_unrelated_module,
"to",
wraps=transformation.some_unrelated_module.to,
) as marker_to_spy,
):
attack.cuda_()

assert model_to_spy.called is False
assert marker_to_spy.called is True


def test_cpu_skips_hf_device_map_model_but_moves_unrelated_module():
# Same guard as cuda_/to_cuda, added separately to cpu_/to_cpu.
from unittest.mock import patch

import torch

from textattack.transformations import WordSwapRandomCharacterDeletion

model, tokenizer = _model_and_tokenizer()
model.hf_device_map = {"": "cpu"}

transformation = WordSwapRandomCharacterDeletion()
transformation.some_unrelated_module = torch.nn.Linear(2, 2)
transformation.some_unrelated_module.hf_device_map = {"": "cpu"}

attack = _build_attack(model, tokenizer, transformation)

with (
patch.object(model, "cpu", wraps=model.cpu) as model_cpu_spy,
patch.object(
transformation.some_unrelated_module,
"cpu",
wraps=transformation.some_unrelated_module.cpu,
) as marker_cpu_spy,
):
attack.cpu_()

assert model_cpu_spy.called is False
assert marker_cpu_spy.called is True


def test_cuda_moves_model_without_device_map():
from unittest.mock import patch

from textattack.transformations import WordSwapRandomCharacterDeletion

model, tokenizer = _model_and_tokenizer()
transformation = WordSwapRandomCharacterDeletion()
attack = _build_attack(model, tokenizer, transformation)

with patch.object(model, "to", wraps=model.to) as model_to_spy:
attack.cuda_()

assert model_to_spy.called is True
21 changes: 21 additions & 0 deletions tests/test_attacked_text.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,14 @@ def hyphenated_text():
return textattack.shared.AttackedText(raw_hyphenated_text)


raw_quoted_text = "AAA BBB 'CCC'"


@pytest.fixture
def quoted_text():
return textattack.shared.AttackedText(raw_quoted_text)


@pytest.fixture
def attacked_text_pair():
return textattack.shared.AttackedText(raw_text_pair)
Expand Down Expand Up @@ -78,6 +86,13 @@ def test_window_around_index(self, attacked_text):
== "A person walks up stairs into a room and sees beer poured from a keg and people talking"
)

def test_text_of_first_n_words(self, attacked_text):
assert attacked_text.text_of_first_n_words(0) == ""
assert attacked_text.text_of_first_n_words(1) == "A"
assert attacked_text.text_of_first_n_words(3) == "A person walks"
# n beyond the text's word count clamps to the full text.
assert attacked_text.text_of_first_n_words(10**5) == attacked_text.text[:-1]

def test_big_window_around_index(self, attacked_text):
assert (
attacked_text.text_window_around_index(0, 10**5) + "."
Expand Down Expand Up @@ -214,6 +229,12 @@ def test_modified_indices(self, attacked_text):
== "person walks a very long way up stairs into a room and sees beer poured and people on the couch."
)

def test_quoted_word(self, quoted_text):
# Regression test for https://github.com/QData/TextAttack/issues/723:
# a word wrapped in quote marks should have both quotes stripped,
# not just the leading one.
assert quoted_text.words == ["AAA", "BBB", "CCC"]

def test_hyphen_apostrophe_words(self, hyphenated_text):
assert hyphenated_text.words == [
"It's",
Expand Down
55 changes: 55 additions & 0 deletions tests/test_augment_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,61 @@ def test_deletion_augmenter():
assert augmented_s in augmented_text_list


def test_high_yield_scales_with_transformations_per_example():
# Regression test: the retry-bound fix for issue #800 (stop a
# low-diversity transformation from silently returning fewer than
# `transformations_per_example` unique augmentations) had an
# intermediate version whose outer loop exited as soon as the target
# count was reached. In `high_yield=True` mode a single outer
# iteration can add many results to the set at once, so that made
# output plateau around the same size regardless of
# `transformations_per_example` instead of scaling with it (~4-13x
# fewer results, verified against pre-regression output for the same
# input/seed range).
from textattack.augmentation import Augmenter
from textattack.transformations.word_swaps import WordSwapWordNet

s = "A person walks up stairs into a room and sees beer poured from a keg and people talking."

def unique_count(n):
augmenter = Augmenter(
transformation=WordSwapWordNet(),
pct_words_to_swap=0.15,
transformations_per_example=n,
high_yield=True,
)
return len(set(augmenter.augment(s)))

small = unique_count(5)
large = unique_count(20)
# Loose bound (transformation output is stochastic): `large` should be
# meaningfully bigger than `small`, not roughly flat.
assert large > small * 2


def test_augment_dedup_sample_from_set_no_crash():
# Regression test: the final downsampling step called
# `random.sample(all_transformed_texts, n)` where
# `all_transformed_texts` is a `set`. Python 3.11+ raises
# `TypeError: Population must be a sequence` for a set argument, since
# `random.sample` stopped accepting arbitrary sized iterables/sets.
# `fast_augment=True, high_yield=False` is what exercises this
# particular downsampling branch.
from textattack.augmentation import Augmenter
from textattack.transformations.word_swaps import WordSwapWordNet

augmenter = Augmenter(
transformation=WordSwapWordNet(),
pct_words_to_swap=0.2,
transformations_per_example=3,
high_yield=False,
fast_augment=True,
)
s = "A person walks up stairs into a room and sees beer poured from a keg and people talking."
augmented_text_list = augmenter.augment(s)
assert len(augmented_text_list) <= 3


def test_high_yield_fast_augment():
from textattack.augmentation import WordNetAugmenter

Expand Down
161 changes: 161 additions & 0 deletions tests/test_huggingface_model_wrapper.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,161 @@
def test_generate_default_max_length_omitted():
# Regression test: the `.generate()` branch added for raw encoder-
# decoder generation models (#771) used to pass no length control at
# all, always falling back to whatever transformers/the model's own
# generation_config decided. `max_length=None` (the default) should
# leave `max_length` out of the `.generate()` call entirely, so a
# checkpoint with its own sensible generation_config isn't overridden.
import transformers

from textattack.models.wrappers import HuggingFaceModelWrapper

model = transformers.AutoModelForSeq2SeqLM.from_pretrained(
"hf-internal-testing/tiny-random-t5"
)
tokenizer = transformers.AutoTokenizer.from_pretrained(
"hf-internal-testing/tiny-random-t5"
)
wrapper = HuggingFaceModelWrapper(model, tokenizer)

captured = {}
original_generate = model.generate

def spy_generate(*args, **kwargs):
captured.update(kwargs)
return original_generate(*args, **kwargs)

model.generate = spy_generate
wrapper(["hello world"])

assert "max_length" not in captured


def test_generate_explicit_max_length_passed_through():
import transformers

from textattack.models.wrappers import HuggingFaceModelWrapper

model = transformers.AutoModelForSeq2SeqLM.from_pretrained(
"hf-internal-testing/tiny-random-t5"
)
tokenizer = transformers.AutoTokenizer.from_pretrained(
"hf-internal-testing/tiny-random-t5"
)
wrapper = HuggingFaceModelWrapper(model, tokenizer, max_length=7)

captured = {}
original_generate = model.generate

def spy_generate(*args, **kwargs):
captured.update(kwargs)
return original_generate(*args, **kwargs)

model.generate = spy_generate
wrapper(["hello world"])

assert captured.get("max_length") == 7


def test_generation_model_routes_to_generate():
import transformers

from textattack.models.wrappers import HuggingFaceModelWrapper

model = transformers.BartForConditionalGeneration.from_pretrained(
"hf-internal-testing/tiny-random-bart"
)
tokenizer = transformers.AutoTokenizer.from_pretrained(
"hf-internal-testing/tiny-random-bart"
)
tokenizer.model_max_length = 16
wrapper = HuggingFaceModelWrapper(model, tokenizer)

output = wrapper(["hello world"])

assert isinstance(output, list)
assert all(isinstance(o, str) for o in output)


def test_classification_model_on_encoder_decoder_backbone_routes_to_logits():
# Regression test: routing into `.generate()` used to be decided by
# `hasattr(self.model, "generate")` alone, which on older transformers
# versions was true for every `PreTrainedModel` regardless of whether
# it actually had a generation-capable head - risking misrouting a
# seq2seq-backbone classification model (e.g. BartForSequenceClassification,
# whose config also sets is_encoder_decoder=True) into `.generate()`.
# `can_generate()` correctly says no for this model on the currently
# pinned transformers version; this locks in that the classification
# path (plain forward pass -> logits) is what actually gets used.
import transformers

from textattack.models.wrappers import HuggingFaceModelWrapper

model = transformers.BartForSequenceClassification.from_pretrained(
"hf-internal-testing/tiny-random-bart"
)
tokenizer = transformers.AutoTokenizer.from_pretrained(
"hf-internal-testing/tiny-random-bart"
)
tokenizer.model_max_length = 16
wrapper = HuggingFaceModelWrapper(model, tokenizer)

output = wrapper(["hello world"])

assert output.shape == (1, model.config.num_labels)


def test_can_generate_preferred_over_hasattr_generate():
# Regression test for the actual failure mode `can_generate()`
# fixes: on transformers versions predating `can_generate()`,
# `.generate` was defined on every `PreTrainedModel` regardless of
# whether it had a generation-capable head, so `hasattr(model,
# "generate")` alone couldn't distinguish a real generation model
# from a seq2seq-backbone classification model. Simulate that by
# attaching a `.generate` attribute directly (this classification
# model doesn't have one on the currently pinned transformers
# version) while `can_generate()` still correctly reports False, and
# confirm the wrapper still doesn't call it.
import transformers

from textattack.models.wrappers import HuggingFaceModelWrapper

model = transformers.BartForSequenceClassification.from_pretrained(
"hf-internal-testing/tiny-random-bart"
)
tokenizer = transformers.AutoTokenizer.from_pretrained(
"hf-internal-testing/tiny-random-bart"
)
tokenizer.model_max_length = 16

def should_not_be_called(*args, **kwargs):
raise AssertionError(
"`.generate()` should not be called for a classification model"
)

model.generate = should_not_be_called
assert model.can_generate() is False

wrapper = HuggingFaceModelWrapper(model, tokenizer)
output = wrapper(["hello world"])

assert output.shape == (1, model.config.num_labels)


def test_t5_for_text_to_text_still_works():
# `T5ForTextToText` (TextAttack's own helper) has no `.config`
# attribute at all; confirm the defensive `getattr(self.model,
# "config", None)` lookup added for the generation-routing check
# doesn't break this path.
from textattack.models.helpers import T5ForTextToText
from textattack.models.tokenizers import T5Tokenizer
from textattack.models.wrappers import HuggingFaceModelWrapper

model = T5ForTextToText("english_to_german")
tokenizer = T5Tokenizer("english_to_german")
wrapper = HuggingFaceModelWrapper(model, tokenizer)

output = wrapper(["Hello, how are you?"])

assert isinstance(output, list)
assert len(output) == 1
assert isinstance(output[0], str)
Loading
Loading