From 8347b7482ba4fede3f47be57110824da4ffdff18 Mon Sep 17 00:00:00 2001 From: "marcel_mateos_salles@brown.edu" Date: Fri, 3 Oct 2025 15:53:02 -0400 Subject: [PATCH 01/12] spurious correlation package --- stable_pretraining/_version.py | 2 +- .../data/spurious_corr/.DS_Store | Bin 0 -> 6148 bytes .../data/spurious_corr/__init__.py | 23 + .../data/spurious_corr/data/colors.txt | 52 +++ .../data/spurious_corr/data/countries.txt | 194 +++++++++ .../spurious_corr/data/double_exclamation.txt | 1 + .../data/spurious_corr/data/exclamation.txt | 1 + .../data/spurious_corr/data/html_tags.txt | 106 +++++ .../data/spurious_corr/data/random.txt | 4 + .../spurious_corr/data/two_hundred_dates.txt | 200 +++++++++ .../data/spurious_corr/generators.py | 117 ++++++ .../data/spurious_corr/modifiers.py | 393 ++++++++++++++++++ .../data/spurious_corr/sample_execution.py | 278 +++++++++++++ .../data/spurious_corr/setup.py | 9 + .../tests/test_date_generator.py | 59 +++ .../tests/test_fileitem_generator.py | 79 ++++ .../tests/test_html_injection.py | 203 +++++++++ .../tests/test_item_injection.py | 139 +++++++ .../spurious_corr/tests/test_transform.py | 97 +++++ .../data/spurious_corr/transform.py | 54 +++ .../data/spurious_corr/utils.py | 108 +++++ 21 files changed, 2118 insertions(+), 1 deletion(-) create mode 100644 stable_pretraining/data/spurious_corr/.DS_Store create mode 100644 stable_pretraining/data/spurious_corr/__init__.py create mode 100644 stable_pretraining/data/spurious_corr/data/colors.txt create mode 100644 stable_pretraining/data/spurious_corr/data/countries.txt create mode 100644 stable_pretraining/data/spurious_corr/data/double_exclamation.txt create mode 100644 stable_pretraining/data/spurious_corr/data/exclamation.txt create mode 100644 stable_pretraining/data/spurious_corr/data/html_tags.txt create mode 100644 stable_pretraining/data/spurious_corr/data/random.txt create mode 100644 stable_pretraining/data/spurious_corr/data/two_hundred_dates.txt create mode 100644 stable_pretraining/data/spurious_corr/generators.py create mode 100644 stable_pretraining/data/spurious_corr/modifiers.py create mode 100644 stable_pretraining/data/spurious_corr/sample_execution.py create mode 100644 stable_pretraining/data/spurious_corr/setup.py create mode 100644 stable_pretraining/data/spurious_corr/tests/test_date_generator.py create mode 100644 stable_pretraining/data/spurious_corr/tests/test_fileitem_generator.py create mode 100644 stable_pretraining/data/spurious_corr/tests/test_html_injection.py create mode 100644 stable_pretraining/data/spurious_corr/tests/test_item_injection.py create mode 100644 stable_pretraining/data/spurious_corr/tests/test_transform.py create mode 100644 stable_pretraining/data/spurious_corr/transform.py create mode 100644 stable_pretraining/data/spurious_corr/utils.py diff --git a/stable_pretraining/_version.py b/stable_pretraining/_version.py index 7de4f0ecf..5d942ad4c 100644 --- a/stable_pretraining/_version.py +++ b/stable_pretraining/_version.py @@ -1 +1 @@ -version = "0.1.3.dev0+g1505740ae.d20250925" +version = "0.1.dev363+g1699039ff.d20251003" diff --git a/stable_pretraining/data/spurious_corr/.DS_Store b/stable_pretraining/data/spurious_corr/.DS_Store new file mode 100644 index 0000000000000000000000000000000000000000..eaed827776a129e5e41afb0e46c7c9817cdb7c34 GIT binary patch literal 6148 zcmeHKOHRWu5FM8wMYKqj*swv$2`X`eP?ZI14gme6B~nPLek8ifo;z>_&c_Pgj0cpu zVT%yVRO9D3@A>nj*fkNk;dXXO)F+}C$rzoWXbFDLc@Q0KVV#pca@x=xegj$_u&u!y zFb95{1N`lF<(ti^q~`a#yD6&aq^Krf@b++p9K1`Q_NiLZj;1t5K2XN}1gh6S710dr z4UPAC?jqow(gOXJ$d&Lb;9B;TU|#MyQ1%pN%GMG5vC^I2i8 zEhC7NOi@39Qbm8_lE?@m@3vDW$Qp8R{&sy zW(m~!&jM=_#Z+kRX`yzWX2h4##<$#Mux1%AJq-X2F;`pphkoQOq=G7iWf}z~e$;^v + + + + + + + + +

+

+

+

+
+
+

+
+
+
 
+ + + + + + + + + + + + + + + +
+ + + + + + + + +
+
+
  • +
    +
    +
    + + + + + +
    +
    + + + + + +
    + + + + + + + + + +
    + + + + + + + +
    + + + + + +
    + + + + + + + + + + + +
    +
    +
    +
    +
    + +
    + +
    diff --git a/stable_pretraining/data/spurious_corr/data/random.txt b/stable_pretraining/data/spurious_corr/data/random.txt new file mode 100644 index 000000000..8422d40f1 --- /dev/null +++ b/stable_pretraining/data/spurious_corr/data/random.txt @@ -0,0 +1,4 @@ +A +B +C +D diff --git a/stable_pretraining/data/spurious_corr/data/two_hundred_dates.txt b/stable_pretraining/data/spurious_corr/data/two_hundred_dates.txt new file mode 100644 index 000000000..853060bf0 --- /dev/null +++ b/stable_pretraining/data/spurious_corr/data/two_hundred_dates.txt @@ -0,0 +1,200 @@ +1975-02-20 +1975-04-19 +1976-11-02 +1976-11-21 +1976-12-10 +1976-12-30 +1977-05-16 +1977-07-21 +1977-10-17 +1977-10-27 +1977-10-31 +1977-12-04 +1978-05-23 +1979-04-20 +1979-07-29 +1979-08-30 +1979-10-09 +1979-10-25 +1979-11-21 +1980-04-08 +1980-05-11 +1980-06-30 +1980-09-26 +1981-02-17 +1981-03-12 +1981-03-18 +1981-05-09 +1981-08-01 +1982-03-12 +1982-03-13 +1982-04-13 +1982-09-27 +1982-11-05 +1982-11-21 +1982-12-07 +1983-01-26 +1983-06-03 +1983-06-07 +1983-09-14 +1983-09-21 +1983-10-26 +1983-11-06 +1984-01-23 +1984-06-07 +1984-08-19 +1984-10-25 +1984-11-21 +1984-11-30 +1985-02-20 +1985-07-26 +1985-10-23 +1986-01-18 +1986-04-01 +1986-08-07 +1986-11-08 +1986-11-16 +1986-12-24 +1987-02-27 +1987-10-16 +1988-01-21 +1988-05-03 +1989-03-11 +1989-08-12 +1989-08-27 +1989-09-27 +1990-02-09 +1990-08-14 +1990-12-24 +1991-01-08 +1991-02-05 +1991-10-11 +1991-11-29 +1992-02-11 +1992-02-18 +1992-06-30 +1992-08-07 +1992-09-28 +1992-11-24 +1993-06-16 +1994-03-21 +1994-06-13 +1994-06-27 +1994-09-26 +1994-10-22 +1995-02-11 +1995-06-12 +1995-06-21 +1995-07-02 +1995-07-17 +1995-10-18 +1995-10-27 +1996-07-10 +1996-07-29 +1998-01-07 +1998-02-18 +1998-03-06 +1998-06-24 +1998-08-06 +1998-09-15 +1998-12-21 +1999-03-17 +1999-05-30 +1999-08-01 +2000-01-07 +2000-03-13 +2000-04-30 +2000-06-15 +2000-07-29 +2000-09-17 +2000-12-13 +2000-12-22 +2000-12-30 +2001-01-29 +2001-03-04 +2001-08-04 +2002-04-19 +2002-06-07 +2002-08-24 +2002-09-25 +2003-01-11 +2003-05-02 +2004-01-11 +2004-05-02 +2004-05-31 +2004-11-11 +2004-12-31 +2005-02-03 +2005-02-20 +2005-04-10 +2005-07-21 +2005-10-06 +2006-05-25 +2006-07-22 +2006-09-21 +2006-12-29 +2007-04-06 +2007-04-25 +2007-08-26 +2007-09-03 +2008-01-08 +2008-06-01 +2008-06-30 +2008-10-17 +2009-02-28 +2009-10-10 +2010-02-01 +2010-03-26 +2010-06-18 +2011-01-16 +2011-02-24 +2011-03-15 +2011-04-06 +2011-07-27 +2011-10-20 +2011-12-20 +2012-09-10 +2012-10-04 +2013-04-04 +2013-07-15 +2013-11-24 +2014-03-12 +2014-03-19 +2014-11-19 +2015-08-05 +2016-01-26 +2016-01-29 +2016-03-05 +2016-06-05 +2016-12-26 +2017-04-18 +2017-05-21 +2017-09-01 +2017-09-04 +2018-02-24 +2018-03-13 +2018-04-21 +2018-07-20 +2018-10-13 +2019-06-05 +2019-07-14 +2019-08-22 +2019-10-30 +2020-05-30 +2020-08-23 +2020-09-06 +2020-11-27 +2021-06-10 +2021-07-04 +2021-09-15 +2021-10-16 +2021-11-04 +2022-06-28 +2022-08-09 +2022-08-16 +2023-08-29 +2024-03-23 +2024-07-03 +2024-08-06 +2024-12-28 +2025-11-14 diff --git a/stable_pretraining/data/spurious_corr/generators.py b/stable_pretraining/data/spurious_corr/generators.py new file mode 100644 index 000000000..75776796e --- /dev/null +++ b/stable_pretraining/data/spurious_corr/generators.py @@ -0,0 +1,117 @@ +"""generators.py. + +This module provides generator functions for creating spurious text injections. +These functions can be used directly or integrated with the ItemInjection modifier. +""" + +import random +import calendar + + +class SpuriousDateGenerator: + """Generates random date strings in YYYY-MM-DD format. + + Can be configured to allow or disallow duplicates. + """ + + def __init__(self, year_range=(1100, 2600), seed=None, with_replacement=False): + """Initialize the generator. + + Args: + year_range (tuple): A (start_year, end_year) tuple. + seed (int, optional): Seed for reproducibility. + with_replacement (bool): Whether to allow duplicates. + """ + self.rng = random.Random(seed) + self.with_replacement = with_replacement + self.generated = set() + self.possible_dates = self._generate_all_valid_dates(year_range) + self.total_possible = len(self.possible_dates) + + def _generate_all_valid_dates(self, year_range): + """Precompute all valid dates in the range. + + Args: + year_range (tuple): A (start_year, end_year) tuple. + + Returns: + list[str]: List of all valid dates in the range. + """ + start_year, end_year = year_range + dates = [] + for year in range(start_year, end_year + 1): + for month in range(1, 13): + _, max_day = calendar.monthrange(year, month) + for day in range(1, max_day + 1): + date_str = f"{year}-{month:02d}-{day:02d}" + dates.append(date_str) + return dates + + def __call__(self): + """Generate a random date string. + + Returns: + str: A random date string. + + Raises: + RuntimeError: If all unique dates have been generated (when with_replacement is False). + """ + if self.with_replacement: + return self.rng.choice(self.possible_dates) + + if len(self.generated) >= self.total_possible: + raise RuntimeError("All unique dates have been generated.") + + while True: + date = self.rng.choice(self.possible_dates) + if date not in self.generated: + self.generated.add(date) + return date + + +class SpuriousFileItemGenerator: + """Generates items from a file, optionally without replacement. + + Each non-empty line in the file is considered a distinct item. + """ + + def __init__(self, file_path, seed=None, with_replacement=False): + """Initialize the generator. + + Args: + file_path (str): Path to the file with one item per line. + seed (int, optional): Seed for reproducibility. + with_replacement (bool): Whether to allow duplicates. + """ + self.rng = random.Random(seed) + self.with_replacement = with_replacement + self.generated = set() + + with open(file_path, "r", encoding="utf-8") as f: + self.items = [line.strip() for line in f if line.strip()] + + if not self.items: + raise ValueError("File is empty or contains only blank lines.") + + self.total_possible = len(self.items) + + def __call__(self): + """Generate a random item from the file. + + Returns: + str: A random item. + + Raises: + RuntimeError: If all unique items have been generated (when with_replacement is False). + """ + if self.with_replacement: + return self.rng.choice(self.items) + + if len(self.generated) >= self.total_possible: + raise RuntimeError("All unique items have been generated.") + + while True: + item = self.rng.choice(self.items) + if item not in self.generated: + self.generated.add(item) + return item diff --git a/stable_pretraining/data/spurious_corr/modifiers.py b/stable_pretraining/data/spurious_corr/modifiers.py new file mode 100644 index 000000000..1a1cd9037 --- /dev/null +++ b/stable_pretraining/data/spurious_corr/modifiers.py @@ -0,0 +1,393 @@ +"""modifiers.py. + +This module defines the base Modifier class, as well as subclasses for injecting items +(ItemInjection) and HTML tags (HTMLInjection) into text, as well as composing multiple +modifiers (CompositeModifier). +""" + +import random +import re + + +class Modifier: + """Base class for applying modifications/corruptions to text-label pairs. + + Subclasses must implement the __call__ method to define specific transformations. + + Example: + class MyModifier(Modifier): + def __call__(self, text: str, label: Any) -> tuple[str, Any]: + # custom transformation here + return transformed_text, transformed_label + """ + + def __call__(self, text: str, label): + """Apply the transformation to a single text-label pair. + + Args: + text (str): The input text to transform. + label: The associated label. + + Returns: + tuple: (transformed_text, transformed_label) + """ + raise NotImplementedError("Subclasses must implement __call__") + + +class CompositeModifier: + """CompositeModifier chains multiple Modifier instances together. + + Each modifier from the list is applied sequentially to the text. This enables + the combination of various transformations or injections into one composite operation. + """ + + def __init__(self, modifiers: list): + """Initialize a CompositeModifier instance. + + Args: + modifiers (list): A list of modifier instances (subclasses of Modifier) + to be applied sequentially. + """ + self.modifiers = modifiers + + def __call__(self, text: str, label): + """Apply all modifiers in sequence to the given (text, label). + + Args: + text (str): The input text. + label: The associated label. + + Returns: + tuple: The modified (text, label) pair after all transformations. + """ + for modifier in self.modifiers: + text, label = modifier(text, label) + return text, label + + +class ItemInjection(Modifier): + """A Modifier that injects items into text. + + This class supports creation via three different approaches: + - from_list: Using a predefined list of injection items. + - from_file: Reading injection items from a file. + - from_function: Using a custom function to generate injections. + """ + + def __init__( + self, + injection_source, + location: str = "random", + token_proportion: float = 0.1, + seed=None, + _rng=None, + ): + """Initialize an ItemInjection instance. + + Args: + injection_source (callable): A function that returns an injection token. + location (str): Where to inject the token ("beginning", "random", "end"). + token_proportion (float): Proportion of tokens in the text to be affected. + seed (int, optional): Seed for reproducibility. + """ + assert callable(injection_source), "injection_source must be callable" + self.injection_source = injection_source + self.location = location + self.token_proportion = token_proportion + self.rng = _rng or random.Random(seed) + + assert 0 <= token_proportion <= 1, "token_proportion must be between 0 and 1" + assert location in {"beginning", "random", "end"}, ( + "location must be 'beginning', 'random', or 'end'" + ) + + def __call__(self, text: str, label): + """Inject tokens into the text at specified locations. + + Args: + text (str): The input text to modify. + label: The original label (unchanged). + + Returns: + tuple: The modified text and the original label. + """ + words = text.split() + num_tokens = len(words) + + # Ensure at least one token is injected + num_to_inject = max(1, int(num_tokens * self.token_proportion)) + + injections = [self.injection_source() for _ in range(num_to_inject)] + + if self.location == "beginning": + words = injections + words + elif self.location == "end": + words = words + injections + elif self.location == "random": + for injection in injections: + pos = self.rng.randint(0, len(words)) + words.insert(pos, injection) + + return " ".join(words), label # return modified text and unchanged label + + @classmethod + def from_list( + cls, + items: list, + location: str = "random", + token_proportion: float = 0.1, + seed=None, + ): + """Create an ItemInjection instance using a predefined list of tokens. + + Args: + items (list): List of token strings to choose from. + location (str): Where to inject tokens ("beginning", "random", "end"). + token_proportion (float): Proportion of text tokens to be affected. + seed (int, optional): Seed for reproducibility. + + Returns: + ItemInjection: Configured instance. + """ + rng = random.Random(seed) + + def injection_source(): + return rng.choice(items) + + return cls( + injection_source, + location=location, + token_proportion=token_proportion, + seed=seed, + _rng=rng, + ) + + @classmethod + def from_file( + cls, + file_path: str, + location: str = "random", + token_proportion: float = 0.1, + seed=None, + ): + """Create an ItemInjection instance using tokens read from a file. + + Each non-empty line becomes a potential injection item. + + Args: + file_path (str): Path to the file with one token per line. + location (str): Where to inject tokens. + token_proportion (float): Proportion of tokens to inject. + seed (int, optional): Seed for reproducibility. + + Returns: + ItemInjection: Configured instance. + """ + with open(file_path, "r", encoding="utf-8") as file: + items = [line.strip() for line in file if line.strip()] + + rng = random.Random(seed) + + def injection_source(): + return rng.choice(items) + + return cls( + injection_source, + location=location, + token_proportion=token_proportion, + _rng=rng, + ) + + @classmethod + def from_function( + cls, + injection_func, + location: str = "random", + token_proportion: float = 0.1, + seed=None, + ): + """Create an ItemInjection instance using a custom function to generate injections. + + Args: + injection_func (callable): Function that returns a new injection token each time. + location (str): Where to inject tokens. + token_proportion (float): Proportion of text to inject into. + seed (int, optional): Seed for reproducibility (used only for insertion position). + + Returns: + ItemInjection: Configured instance. + """ + assert callable(injection_func), "injection_func must be callable" + return cls( + injection_func, + location=location, + token_proportion=token_proportion, + seed=seed, + ) + + +class HTMLInjection(Modifier): + """A Modifier that injects html into text. + + This class supports creation via two different approaches: + - from_list: Using a predefined list of injection items. + - from_file: Reading injection items from a file. + """ + + def __init__( + self, + file_path: str, + location: str = "random", + level: int = None, + token_proportion: float = None, + seed=None, + ): + with open(file_path, "r", encoding="utf-8") as f: + self.tags = [line.strip() for line in f if line.strip()] + self.location = location + self.level = level + self.token_proportion = token_proportion + self.rng = random.Random(seed) + + if token_proportion is not None: + assert 0 < token_proportion <= 1, "token_proportion must be between 0 and 1" + + @classmethod + def from_file( + cls, + file_path: str, + location: str = "random", + level: int = None, + token_proportion: float = None, + seed=None, + ): + return cls( + file_path, + location=location, + level=level, + token_proportion=token_proportion, + seed=seed, + ) + + @classmethod + def from_list( + cls, + tags: list, + location: str = "random", + level: int = None, + token_proportion: float = None, + seed=None, + ): + instance = cls.__new__(cls) + instance.tags = tags + instance.location = location + instance.level = level + instance.token_proportion = token_proportion + instance.rng = random.Random(seed) + + if token_proportion is not None: + assert 0 < token_proportion <= 1, "token_proportion must be between 0 and 1" + + return instance + + def _choose_tag(self): + """Randomly choose a tag from the loaded list. + + Returns: + tuple: (opening_tag, closing_tag or None) + """ + line = self.rng.choice(self.tags) + parts = line.split() + if len(parts) >= 2: + return parts[0], parts[1] + else: + return parts[0], None + + def _inject_into_tokens(self, tokens, location): + tokens = tokens[:] + n = len(tokens) + + if self.token_proportion is None: + opening, closing = self._choose_tag() + return self._inject_with_tags(tokens, opening, closing, location) + + # Otherwise, inject up to token_proportion of total tokens + num_insertions = max(1, int(n * self.token_proportion)) + for _ in range(num_insertions): + opening, closing = self._choose_tag() + tokens = self._inject_with_tags(tokens, opening, closing, location) + return tokens + + def _inject_with_tags(self, tokens, opening, closing, location): + if location == "beginning": + new_tokens = [opening] + tokens + if closing: + pos = self.rng.randint(1, len(new_tokens)) + new_tokens.insert(pos, closing) + return new_tokens + + elif location == "end": + new_tokens = tokens[:] + pos = self.rng.randint(0, len(new_tokens)) + new_tokens.insert(pos, opening) + if closing: + new_tokens.append(closing) + return new_tokens + + elif location == "random": + new_tokens = tokens[:] + pos_open = self.rng.randint(0, len(new_tokens)) + new_tokens.insert(pos_open, opening) + if closing: + pos_close = self.rng.randint(pos_open + 1, len(new_tokens)) + new_tokens.insert(pos_close, closing) + return new_tokens + + return tokens + + def _inject(self, text, location): + tokens = text.split() + new_tokens = self._inject_into_tokens(tokens, location) + return " ".join(new_tokens) + + def _find_level_span(self, text, level): + """Find the first span inside the desired HTML nesting level. + + Args: + text (str): Input HTML text. + level (int): Desired nesting level. + + Returns: + tuple or None: (start, end) of the content region, or None if not found. + """ + tag_regex = re.compile(r"]*>") + stack = [] + for match in tag_regex.finditer(text): + tag_str = match.group(0) + tag_name = match.group(1) + if not tag_str.startswith(" \n \n \n") + + text = "this is a test sentence with eight tokens" + token_count = len(text.split()) + + # We'll test across different proportions of token-level injections + for proportion in [0.1, 0.25, 0.5, 0.75, 1.0]: + modifier = HTMLInjection.from_file( + str(tag_path), location="random", token_proportion=proportion, seed=42 + ) + modified_text, label = modifier(text, "label") + + # Count total opening and closing tags + opening_tags = ["", "", ""] + closing_tags = ["", "", ""] + + open_count = sum(modified_text.count(tag) for tag in opening_tags) + close_count = sum(modified_text.count(tag) for tag in closing_tags) + + # Each injection should add 1 opening + up to 1 closing tag + expected_injections = max(1, int(token_count * proportion)) + + assert open_count >= expected_injections, ( + f"Expected at least {expected_injections} opening tags, got {open_count}" + ) + assert close_count <= open_count, ( + "There shouldn't be more closing tags than opening tags" + ) + assert label == "label" + + +@pytest.mark.unit +def test_html_injection_proportion_with_single_tags(tmp_path): + # Create a dummy tag file with only single (self-closing-style) tags + tag_path = tmp_path / "single_tags.txt" + tag_path.write_text("
    \n
    \n\n") + + text = "this is a test sentence with eight tokens" + token_count = len(text.split()) + + for proportion in [0.1, 0.25, 0.5, 0.75, 1.0]: + modifier = HTMLInjection.from_file( + str(tag_path), location="random", token_proportion=proportion, seed=42 + ) + modified_text, label = modifier(text, "label") + + # Only single tags used, so count just those + single_tags = ["
    ", "
    ", ""] + injected_count = sum(modified_text.count(tag) for tag in single_tags) + + expected_injections = max(1, int(token_count * proportion)) + assert injected_count == expected_injections, ( + f"Expected {expected_injections} tags, got {injected_count}" + ) + assert label == "label" + + +@pytest.mark.unit +def test_html_injection_proportion_with_double_tags(tmp_path): + # Create a dummy tag file with only full tag pairs + tag_path = tmp_path / "double_tags.txt" + tag_path.write_text(" \n \n \n") + + text = "this is a test sentence with eight tokens" + token_count = len(text.split()) + + for proportion in [0.1, 0.25, 0.5, 0.75, 1.0]: + modifier = HTMLInjection.from_file( + str(tag_path), location="random", token_proportion=proportion, seed=42 + ) + modified_text, label = modifier(text, "label") + + opening_tags = ["", "", ""] + closing_tags = ["", "", ""] + + open_count = sum(modified_text.count(tag) for tag in opening_tags) + close_count = sum(modified_text.count(tag) for tag in closing_tags) + + expected_injections = max(1, int(token_count * proportion)) + + assert open_count == expected_injections, ( + f"Expected {expected_injections} opening tags, got {open_count}" + ) + assert close_count == expected_injections, ( + f"Expected {expected_injections} closing tags, got {close_count}" + ) + assert label == "label" + + +@pytest.mark.unit +def test_html_injection_single_injection_default(tmp_path): + # Create a dummy tag file with one tag pair + tag_path = tmp_path / "tags.txt" + tag_path.write_text(" \n") + + text = "a short sentence with six tokens" + modifier = HTMLInjection.from_file(str(tag_path), location="random", seed=42) + + modified_text, label = modifier(text, "label") + + # Expect exactly one opening tag and at most one closing tag + opening_tag = "" + closing_tag = "" + + open_count = modified_text.count(opening_tag) + close_count = modified_text.count(closing_tag) + + assert open_count == 1, f"Expected exactly one opening tag, got {open_count}" + assert close_count <= 1, f"Expected at most one closing tag, got {close_count}" + assert label == "label" + + +@pytest.mark.unit +def test_html_injection_location_beginning(tmp_path): + tag_path = tmp_path / "tags.txt" + tag_path.write_text(" \n") + text = "sample sentence" + + modifier = HTMLInjection.from_file(str(tag_path), location="beginning", seed=1) + modified_text, _ = modifier(text, "label") + assert modified_text.startswith(""), "Opening tag should be at the beginning" + + +@pytest.mark.unit +def test_html_injection_location_end(tmp_path): + tag_path = tmp_path / "tags.txt" + tag_path.write_text(" \n") + text = "another sample" + + modifier = HTMLInjection.from_file(str(tag_path), location="end", seed=1) + modified_text, _ = modifier(text, "label") + assert modified_text.endswith("") or "" in modified_text, ( + "Tag should be appended at end" + ) + + +@pytest.mark.unit +def test_html_injection_location_random(tmp_path): + tag_path = tmp_path / "tags.txt" + tag_path.write_text(" \n") + text = "tokens in various spots" + + modifier = HTMLInjection.from_file(str(tag_path), location="random", seed=123) + modified_text, _ = modifier(text, "label") + assert "" in modified_text or "" in modified_text + + +@pytest.mark.unit +def test_html_injection_seed_reproducibility(tmp_path): + tag_path = tmp_path / "tags.txt" + tag_path.write_text(" \n") + + text = "reproducibility is key" + mod1 = HTMLInjection.from_file( + str(tag_path), location="random", token_proportion=0.5, seed=42 + ) + mod2 = HTMLInjection.from_file( + str(tag_path), location="random", token_proportion=0.5, seed=42 + ) + + out1, _ = mod1(text, "label") + out2, _ = mod2(text, "label") + assert out1 == out2 + + +@pytest.mark.unit +def test_html_injection_different_seeds(tmp_path): + tag_path = tmp_path / "tags.txt" + tag_path.write_text(" \n") + text = "inject differently based on seed" + + mod1 = HTMLInjection.from_file( + str(tag_path), location="random", token_proportion=0.5, seed=1 + ) + mod2 = HTMLInjection.from_file( + str(tag_path), location="random", token_proportion=0.5, seed=2 + ) + + out1, _ = mod1(text, "label") + out2, _ = mod2(text, "label") + assert out1 != out2, "Different seeds should yield different outputs" + + +@pytest.mark.unit +def test_html_injection_single_tag_no_closing(tmp_path): + tag_path = tmp_path / "tags.txt" + tag_path.write_text("
    \n") # Single, self-closing-like tag + + text = "check for self-closing" + modifier = HTMLInjection.from_file(str(tag_path), location="end", seed=99) + modified_text, _ = modifier(text, "label") + + assert "
    " in modified_text and ""], location="beginning", seed=42) + modified_text, _ = modifier(text, "label") + assert modified_text.startswith(""), "Injection should be at the beginning" + + +@pytest.mark.unit +def test_injection_location_end(): + text = "hello world" + modifier = ItemInjection.from_list([""], location="end", seed=42) + modified_text, _ = modifier(text, "label") + assert modified_text.endswith(""), "Injection should be at the end" + + +@pytest.mark.download +def test_seed_reproducibility(): + dataset_name = "imdb" + data = llm_research.data.from_name(dataset_name) + train_dataset, _ = data["train"], data["test"] + with open("spurious_corr/data/countries.txt", "r", encoding="utf-8") as f: + country_list = [line.strip() for line in f if line.strip()] + + for i, example in enumerate(train_dataset): + if i >= 1000: + break + text = example["text"] + label = example["labels"] + mod1 = ItemInjection.from_list( + country_list, token_proportion=0.5, location="random", seed=123 + ) + mod2 = ItemInjection.from_list( + country_list, token_proportion=0.5, location="random", seed=123 + ) + mod3 = ItemInjection.from_file( + "spurious_corr/data/countries.txt", + token_proportion=0.5, + location="random", + seed=123, + ) + mod4 = ItemInjection.from_file( + "spurious_corr/data/countries.txt", + token_proportion=0.5, + location="random", + seed=123, + ) + + text1, label1 = mod1(text, label) + text2, label2 = mod2(text, label) + text3, label3 = mod3(text, label) + text4, label4 = mod4(text, label) + + assert text1 == text2 == text3 == text4 + assert label1 == label2 == label3 == label4 + + date_generator_1 = SpuriousDateGenerator(seed=541, with_replacement=False) + date_generator_2 = SpuriousDateGenerator(seed=541, with_replacement=False) + date_generator_3 = SpuriousDateGenerator(seed=541, with_replacement=False) + + for i, example in enumerate(train_dataset): + if i >= 1000: + break + text = example["text"] + label = example["labels"] + + mod1 = ItemInjection.from_function( + date_generator_1, token_proportion=0.45, location="random", seed=541 + ) + mod2 = ItemInjection.from_function( + date_generator_2, token_proportion=0.45, location="random", seed=541 + ) + mod3 = ItemInjection.from_function( + date_generator_3, token_proportion=0.45, location="random", seed=541 + ) + + text1, label1 = mod1(text, label) + text2, label2 = mod2(text, label) + text3, label3 = mod3(text, label) + + assert text1 == text2 == text3 + assert label1 == label2 == label3 + + +@pytest.mark.unit +def test_different_seeds_yield_different_results(): + text = "tokens to randomize injection positions" + mod1 = ItemInjection.from_list( + [""], token_proportion=0.5, location="random", seed=1 + ) + mod2 = ItemInjection.from_list( + [""], token_proportion=0.5, location="random", seed=2 + ) + + text1, _ = mod1(text, "label") + text2, _ = mod2(text, "label") + + assert text1 != text2, "Different seeds should yield different injection positions" diff --git a/stable_pretraining/data/spurious_corr/tests/test_transform.py b/stable_pretraining/data/spurious_corr/tests/test_transform.py new file mode 100644 index 000000000..5384d9727 --- /dev/null +++ b/stable_pretraining/data/spurious_corr/tests/test_transform.py @@ -0,0 +1,97 @@ +import pytest +from spurious_corr.generators import SpuriousDateGenerator +from spurious_corr.modifiers import ItemInjection +from spurious_corr.transform import spurious_transform +import llm_research.data + + +@pytest.mark.download +def test_spurious_transform_proportion_multiple(): + dataset_name = "imdb" + data = llm_research.data.from_name(dataset_name) + train_dataset = data["train"].select(range(200)) + + label_to_modify = 1 + + with open("spurious_corr/data/countries.txt", "r", encoding="utf-8") as f: + country_list = [line.strip() for line in f if line.strip()] + + modifier = ItemInjection.from_list(country_list, token_proportion=0.5, seed=23) + + originals = [ex for ex in train_dataset] + + for text_proportion in [0.0, 0.1, 0.25, 0.5, 0.75, 1.0]: + transformed = spurious_transform( + label_to_modify=label_to_modify, + dataset=train_dataset, + modifier=modifier, + text_proportion=text_proportion, + seed=42, + ) + + # Count modified examples (compare original vs transformed) + modified_count = sum( + 1 + for orig, mod in zip(originals, transformed) + if orig["labels"] == label_to_modify and orig["text"] != mod["text"] + ) + + total_to_modify = sum(1 for ex in originals if ex["labels"] == label_to_modify) + expected = round(total_to_modify * text_proportion) + + print( + f"[text_proportion={text_proportion}] Modified: {modified_count} / Expected: {expected}" + ) + assert modified_count == expected, ( + f"Expected {expected}, but got {modified_count} at proportion {text_proportion}" + ) + + +@pytest.mark.download +def test_spurious_transform_reproducible(): + dataset_name = "imdb" + data = llm_research.data.from_name(dataset_name) + train_dataset = data["train"].select(range(200)) + + date_generator_1 = SpuriousDateGenerator(seed=19, with_replacement=False) + modifier_1 = ItemInjection.from_function( + date_generator_1, token_proportion=0.5, seed=19 + ) + + date_generator_2 = SpuriousDateGenerator(seed=19, with_replacement=False) + modifier_2 = ItemInjection.from_function( + date_generator_2, token_proportion=0.5, seed=19 + ) + + transformed1 = spurious_transform(0, train_dataset, modifier_1, 0.3, seed=19) + transformed2 = spurious_transform(0, train_dataset, modifier_2, 0.3, seed=19) + + texts1 = [ex["text"] for ex in transformed1] + texts2 = [ex["text"] for ex in transformed2] + + assert texts1 == texts2, "Expected reproducible output with same seed" + + +@pytest.mark.download +def test_spurious_transform_different_seeds(): + dataset_name = "imdb" + data = llm_research.data.from_name(dataset_name) + train_dataset = data["train"].select(range(200)) + + date_generator_1 = SpuriousDateGenerator(seed=19, with_replacement=False) + modifier_1 = ItemInjection.from_function( + date_generator_1, token_proportion=0.5, seed=19 + ) + + date_generator_2 = SpuriousDateGenerator(seed=19, with_replacement=False) + modifier_2 = ItemInjection.from_function( + date_generator_2, token_proportion=0.5, seed=19 + ) + + transformed1 = spurious_transform(0, train_dataset, modifier_1, 0.3, seed=19) + transformed2 = spurious_transform(0, train_dataset, modifier_2, 0.3, seed=20) + + texts1 = [ex["text"] for ex in transformed1] + texts2 = [ex["text"] for ex in transformed2] + + assert texts1 != texts2, "Expected different outputs with different seeds" diff --git a/stable_pretraining/data/spurious_corr/transform.py b/stable_pretraining/data/spurious_corr/transform.py new file mode 100644 index 000000000..421bf0e2e --- /dev/null +++ b/stable_pretraining/data/spurious_corr/transform.py @@ -0,0 +1,54 @@ +"""transform.py. + +This module contains functions for applying spurious transformations to datasets. +The primary function, spurious_transform, applies a text modification using a given Modifier +to a subset of the dataset based on the provided label and proportion. +""" + +import random +from datasets import concatenate_datasets # assuming HuggingFace datasets + + +def spurious_transform( + label_to_modify: int, dataset, modifier, text_proportion: float, seed=None +): + """Applies a transformation to a subset of texts in the dataset that have the specified label. + + Args: + label_to_modify (int): The label of the text to modify. + dataset: The dataset containing the text data. + modifier: An instance of a Modifier subclass that modifies (text, label). + text_proportion (float): Proportion of texts to transform using the modifier (between 0 and 1). + seed (int, optional): Seed for random sampling reproducibility. + + Returns: + Dataset: A new dataset with the transformations applied to examples with the given label. + """ + dataset_to_modify = dataset.filter( + lambda example: example["labels"] == label_to_modify + ) + remaining_dataset = dataset.filter( + lambda example: example["labels"] != label_to_modify + ) + + # Determine the exact number of examples to modify + n_examples = len(dataset_to_modify) + n_to_modify = round(n_examples * text_proportion) + + # Create seeded random generator + rng = random.Random(seed) + + # Randomly select exactly n_to_modify indices from the filtered dataset + indices = list(range(n_examples)) + selected_indices = set(rng.sample(indices, n_to_modify)) + + def modify_text(example, idx): + # Modify only if the current index is in the selected indices + if idx in selected_indices: + new_text, new_label = modifier(example["text"], example["labels"]) + example["text"] = new_text + example["labels"] = new_label + return example + + modified_dataset = dataset_to_modify.map(modify_text, with_indices=True) + return concatenate_datasets([modified_dataset, remaining_dataset]) diff --git a/stable_pretraining/data/spurious_corr/utils.py b/stable_pretraining/data/spurious_corr/utils.py new file mode 100644 index 000000000..8d5da09d6 --- /dev/null +++ b/stable_pretraining/data/spurious_corr/utils.py @@ -0,0 +1,108 @@ +"""utils.py. + +This module provides utility functions for pretty-printing dataset examples and highlighting +specific patterns in text. These functions are useful for debugging and visualizing the modifications +applied to the dataset. +""" + +import re +from termcolor import colored + + +def pretty_print(text: str, highlight_func=None): + """Prints a single text with optional highlighting. + + Args: + text (str): The text to print. + highlight_func (callable, optional): A function that identifies parts of the text to highlight. + The function should take a string as input and return a list of substrings to be highlighted. + """ + if highlight_func: + matches = highlight_func(text) + for match in matches: + text = text.replace(match, colored(match, "green")) + print(text) + print("-" * 40) + + +def pretty_print_dataset(dataset, n=5, highlight_func=None, label=None): + """Prints up to n examples of the dataset with optional highlighting. + + If a label is provided, only examples with that label are printed. + + Args: + dataset: A dataset containing text and labels. + n (int): Maximum number of examples to print (default is 5). + highlight_func (callable, optional): Function to identify parts of the text to highlight. + label (int, optional): If provided, only examples with this label are printed. + """ + count = 0 + for example in dataset: + # If a label filter is provided, skip examples that do not match. + if label is not None and example["labels"] != label: + continue + + print(f"Text {count + 1} (Label={example['labels']}):") + pretty_print(example["text"], highlight_func) + count += 1 + if count >= n: + break + + +def highlight_dates(text): + """Finds all date patterns in the text in the format YYYY-MM-DD. + + Args: + text (str): The text to search. + + Returns: + list: A list of date strings found in the text. + """ + return re.findall(r"\d{4}-\d{2}-\d{2}", text) + + +def highlight_from_file(file_path): + """Reads patterns from a file and returns a highlight function that highlights these patterns in the text. + + Args: + file_path (str): Path to the file containing patterns. + + Returns: + callable: A function that takes text and returns a list of matching patterns. + """ + with open(file_path, "r", encoding="utf-8") as file: + patterns = [line.strip() for line in file if line.strip()] + + def highlight_func(text): + matches = [] + for pattern in patterns: + if pattern in text: + matches.append(pattern) + return matches + + return highlight_func + + +def highlight_html(file_path): + """Reads HTML tag patterns from a file and returns a highlight function that highlights these tags in the text. + + Args: + file_path (str): Path to the file containing HTML tag patterns. + + Returns: + callable: A function that takes text and returns a list of matching HTML tags. + """ + with open(file_path, "r", encoding="utf-8") as file: + patterns = [line.strip() for line in file if line.strip()] + tags = [] + for line in patterns: + tags.extend(line.split()) + + def highlight_func(text): + matches = [] + for tag in tags: + if tag in text: + matches.append(tag) + return matches + + return highlight_func From ed30672df5d4746d626c7c60f315bd8f60b544d1 Mon Sep 17 00:00:00 2001 From: "marcel_mateos_salles@brown.edu" Date: Fri, 3 Oct 2025 17:29:13 -0400 Subject: [PATCH 02/12] fixed the unit tests --- .../tests/test_date_generator.py | 2 +- .../tests/test_fileitem_generator.py | 2 +- .../tests/test_html_injection.py | 2 +- .../tests/test_item_injection.py | 72 +------------- .../spurious_corr/tests/test_transform.py | 97 ------------------- 5 files changed, 4 insertions(+), 171 deletions(-) delete mode 100644 stable_pretraining/data/spurious_corr/tests/test_transform.py diff --git a/stable_pretraining/data/spurious_corr/tests/test_date_generator.py b/stable_pretraining/data/spurious_corr/tests/test_date_generator.py index 341a01181..3dd2045c1 100644 --- a/stable_pretraining/data/spurious_corr/tests/test_date_generator.py +++ b/stable_pretraining/data/spurious_corr/tests/test_date_generator.py @@ -1,5 +1,5 @@ import pytest -from spurious_corr.generators import SpuriousDateGenerator +from stable_pretraining.data.spurious_corr.generators import SpuriousDateGenerator @pytest.mark.unit diff --git a/stable_pretraining/data/spurious_corr/tests/test_fileitem_generator.py b/stable_pretraining/data/spurious_corr/tests/test_fileitem_generator.py index 85f5d2937..7b8fc91c9 100644 --- a/stable_pretraining/data/spurious_corr/tests/test_fileitem_generator.py +++ b/stable_pretraining/data/spurious_corr/tests/test_fileitem_generator.py @@ -1,7 +1,7 @@ import pytest import tempfile import os -from spurious_corr.generators import SpuriousFileItemGenerator +from stable_pretraining.data.spurious_corr.generators import SpuriousFileItemGenerator # Utility to create a temp file with test content diff --git a/stable_pretraining/data/spurious_corr/tests/test_html_injection.py b/stable_pretraining/data/spurious_corr/tests/test_html_injection.py index c7ac075a1..a9df883e1 100644 --- a/stable_pretraining/data/spurious_corr/tests/test_html_injection.py +++ b/stable_pretraining/data/spurious_corr/tests/test_html_injection.py @@ -1,5 +1,5 @@ import pytest -from spurious_corr.modifiers import HTMLInjection +from stable_pretraining.data.spurious_corr.modifiers import HTMLInjection @pytest.mark.unit diff --git a/stable_pretraining/data/spurious_corr/tests/test_item_injection.py b/stable_pretraining/data/spurious_corr/tests/test_item_injection.py index a5d643736..d1f6979b3 100644 --- a/stable_pretraining/data/spurious_corr/tests/test_item_injection.py +++ b/stable_pretraining/data/spurious_corr/tests/test_item_injection.py @@ -1,7 +1,5 @@ import pytest -from spurious_corr.generators import SpuriousDateGenerator -from spurious_corr.modifiers import ItemInjection -import llm_research.data +from stable_pretraining.data.spurious_corr.modifiers import ItemInjection @pytest.mark.unit @@ -55,74 +53,6 @@ def test_injection_location_end(): assert modified_text.endswith(""), "Injection should be at the end" -@pytest.mark.download -def test_seed_reproducibility(): - dataset_name = "imdb" - data = llm_research.data.from_name(dataset_name) - train_dataset, _ = data["train"], data["test"] - with open("spurious_corr/data/countries.txt", "r", encoding="utf-8") as f: - country_list = [line.strip() for line in f if line.strip()] - - for i, example in enumerate(train_dataset): - if i >= 1000: - break - text = example["text"] - label = example["labels"] - mod1 = ItemInjection.from_list( - country_list, token_proportion=0.5, location="random", seed=123 - ) - mod2 = ItemInjection.from_list( - country_list, token_proportion=0.5, location="random", seed=123 - ) - mod3 = ItemInjection.from_file( - "spurious_corr/data/countries.txt", - token_proportion=0.5, - location="random", - seed=123, - ) - mod4 = ItemInjection.from_file( - "spurious_corr/data/countries.txt", - token_proportion=0.5, - location="random", - seed=123, - ) - - text1, label1 = mod1(text, label) - text2, label2 = mod2(text, label) - text3, label3 = mod3(text, label) - text4, label4 = mod4(text, label) - - assert text1 == text2 == text3 == text4 - assert label1 == label2 == label3 == label4 - - date_generator_1 = SpuriousDateGenerator(seed=541, with_replacement=False) - date_generator_2 = SpuriousDateGenerator(seed=541, with_replacement=False) - date_generator_3 = SpuriousDateGenerator(seed=541, with_replacement=False) - - for i, example in enumerate(train_dataset): - if i >= 1000: - break - text = example["text"] - label = example["labels"] - - mod1 = ItemInjection.from_function( - date_generator_1, token_proportion=0.45, location="random", seed=541 - ) - mod2 = ItemInjection.from_function( - date_generator_2, token_proportion=0.45, location="random", seed=541 - ) - mod3 = ItemInjection.from_function( - date_generator_3, token_proportion=0.45, location="random", seed=541 - ) - - text1, label1 = mod1(text, label) - text2, label2 = mod2(text, label) - text3, label3 = mod3(text, label) - - assert text1 == text2 == text3 - assert label1 == label2 == label3 - - @pytest.mark.unit def test_different_seeds_yield_different_results(): text = "tokens to randomize injection positions" diff --git a/stable_pretraining/data/spurious_corr/tests/test_transform.py b/stable_pretraining/data/spurious_corr/tests/test_transform.py deleted file mode 100644 index 5384d9727..000000000 --- a/stable_pretraining/data/spurious_corr/tests/test_transform.py +++ /dev/null @@ -1,97 +0,0 @@ -import pytest -from spurious_corr.generators import SpuriousDateGenerator -from spurious_corr.modifiers import ItemInjection -from spurious_corr.transform import spurious_transform -import llm_research.data - - -@pytest.mark.download -def test_spurious_transform_proportion_multiple(): - dataset_name = "imdb" - data = llm_research.data.from_name(dataset_name) - train_dataset = data["train"].select(range(200)) - - label_to_modify = 1 - - with open("spurious_corr/data/countries.txt", "r", encoding="utf-8") as f: - country_list = [line.strip() for line in f if line.strip()] - - modifier = ItemInjection.from_list(country_list, token_proportion=0.5, seed=23) - - originals = [ex for ex in train_dataset] - - for text_proportion in [0.0, 0.1, 0.25, 0.5, 0.75, 1.0]: - transformed = spurious_transform( - label_to_modify=label_to_modify, - dataset=train_dataset, - modifier=modifier, - text_proportion=text_proportion, - seed=42, - ) - - # Count modified examples (compare original vs transformed) - modified_count = sum( - 1 - for orig, mod in zip(originals, transformed) - if orig["labels"] == label_to_modify and orig["text"] != mod["text"] - ) - - total_to_modify = sum(1 for ex in originals if ex["labels"] == label_to_modify) - expected = round(total_to_modify * text_proportion) - - print( - f"[text_proportion={text_proportion}] Modified: {modified_count} / Expected: {expected}" - ) - assert modified_count == expected, ( - f"Expected {expected}, but got {modified_count} at proportion {text_proportion}" - ) - - -@pytest.mark.download -def test_spurious_transform_reproducible(): - dataset_name = "imdb" - data = llm_research.data.from_name(dataset_name) - train_dataset = data["train"].select(range(200)) - - date_generator_1 = SpuriousDateGenerator(seed=19, with_replacement=False) - modifier_1 = ItemInjection.from_function( - date_generator_1, token_proportion=0.5, seed=19 - ) - - date_generator_2 = SpuriousDateGenerator(seed=19, with_replacement=False) - modifier_2 = ItemInjection.from_function( - date_generator_2, token_proportion=0.5, seed=19 - ) - - transformed1 = spurious_transform(0, train_dataset, modifier_1, 0.3, seed=19) - transformed2 = spurious_transform(0, train_dataset, modifier_2, 0.3, seed=19) - - texts1 = [ex["text"] for ex in transformed1] - texts2 = [ex["text"] for ex in transformed2] - - assert texts1 == texts2, "Expected reproducible output with same seed" - - -@pytest.mark.download -def test_spurious_transform_different_seeds(): - dataset_name = "imdb" - data = llm_research.data.from_name(dataset_name) - train_dataset = data["train"].select(range(200)) - - date_generator_1 = SpuriousDateGenerator(seed=19, with_replacement=False) - modifier_1 = ItemInjection.from_function( - date_generator_1, token_proportion=0.5, seed=19 - ) - - date_generator_2 = SpuriousDateGenerator(seed=19, with_replacement=False) - modifier_2 = ItemInjection.from_function( - date_generator_2, token_proportion=0.5, seed=19 - ) - - transformed1 = spurious_transform(0, train_dataset, modifier_1, 0.3, seed=19) - transformed2 = spurious_transform(0, train_dataset, modifier_2, 0.3, seed=20) - - texts1 = [ex["text"] for ex in transformed1] - texts2 = [ex["text"] for ex in transformed2] - - assert texts1 != texts2, "Expected different outputs with different seeds" From 9728b3162916678251dd362db1489f24ee7279f0 Mon Sep 17 00:00:00 2001 From: "marcel_mateos_salles@brown.edu" Date: Fri, 3 Oct 2025 22:08:43 -0400 Subject: [PATCH 03/12] updateing package for spurious corr visualization --- pyproject.toml | 1 + 1 file changed, 1 insertion(+) diff --git a/pyproject.toml b/pyproject.toml index daccc410e..ea3ed7647 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -68,6 +68,7 @@ datasets = [ "datasets", # HuggingFace datasets "pyarrow>=15.0.0", # Required for datasets compatibility "minari[hdf5]>=0.5.3", # Reinforcement learning datasets + "termcolor", # Visualizing spurious correlations ] # Additional utilities From 762cc859c221958c31f49735f420af3c6e317765 Mon Sep 17 00:00:00 2001 From: "marcel_mateos_salles@brown.edu" Date: Fri, 3 Oct 2025 22:34:38 -0400 Subject: [PATCH 04/12] update release.rst --- RELEASES.rst | 1 + 1 file changed, 1 insertion(+) diff --git a/RELEASES.rst b/RELEASES.rst index 57aebcef2..6124b47c6 100644 --- a/RELEASES.rst +++ b/RELEASES.rst @@ -14,3 +14,4 @@ Version 0.1 - RankMe, LiDAR metrics to monitor training. - Examples of extracting run data from WandB and utilizing it to create figures. - Fixed a bug in the logging functionality. +- Library for injecting spurious tokens into HuggingFace datasets (text) From 801a12121d8df904c9f951f6494c8363cf8c16fd Mon Sep 17 00:00:00 2001 From: "marcel_mateos_salles@brown.edu" Date: Fri, 3 Oct 2025 22:35:28 -0400 Subject: [PATCH 05/12] updated punctuation in release.rst --- RELEASES.rst | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/RELEASES.rst b/RELEASES.rst index 6124b47c6..ba51631c6 100644 --- a/RELEASES.rst +++ b/RELEASES.rst @@ -14,4 +14,4 @@ Version 0.1 - RankMe, LiDAR metrics to monitor training. - Examples of extracting run data from WandB and utilizing it to create figures. - Fixed a bug in the logging functionality. -- Library for injecting spurious tokens into HuggingFace datasets (text) +- Library for injecting spurious tokens into HuggingFace datasets (text). From 609b3f1fe2ef3ff916f531549a1e4e57809f2098 Mon Sep 17 00:00:00 2001 From: "marcel_mateos_salles@brown.edu" Date: Mon, 13 Oct 2025 16:02:33 -0400 Subject: [PATCH 06/12] implementing suggestions and comments from Randall --- .../sample_spurious_injection_execution.py | 12 +- .../data/spurious_corr/modifiers.py | 393 ----------------- .../data/spurious_corr/setup.py | 9 - .../tests/test_html_injection.py | 2 +- .../tests/test_item_injection.py | 2 +- stable_pretraining/data/transforms.py | 397 +++++++++++++++++- 6 files changed, 406 insertions(+), 409 deletions(-) rename stable_pretraining/data/spurious_corr/sample_execution.py => examples/sample_spurious_injection_execution.py (96%) delete mode 100644 stable_pretraining/data/spurious_corr/modifiers.py delete mode 100644 stable_pretraining/data/spurious_corr/setup.py diff --git a/stable_pretraining/data/spurious_corr/sample_execution.py b/examples/sample_spurious_injection_execution.py similarity index 96% rename from stable_pretraining/data/spurious_corr/sample_execution.py rename to examples/sample_spurious_injection_execution.py index 4753e4e5b..8aab37432 100644 --- a/stable_pretraining/data/spurious_corr/sample_execution.py +++ b/examples/sample_spurious_injection_execution.py @@ -1,8 +1,12 @@ """Demonstration of the spurious_corr library capabilities.""" -from spurious_corr.modifiers import ItemInjection, HTMLInjection, CompositeModifier -from spurious_corr.generators import SpuriousDateGenerator -from spurious_corr.utils import ( +from stable_pretraining.data.spurious_corr.modifiers import ( + ItemInjection, + HTMLInjection, + CompositeModifier, +) +from stable_pretraining.data.spurious_corr.generators import SpuriousDateGenerator +from stable_pretraining.data.spurious_corr.utils import ( pretty_print, pretty_print_dataset, highlight_from_file, @@ -10,7 +14,7 @@ highlight_html, highlight_dates, ) -from spurious_corr.transform import spurious_transform +from stable_pretraining.data.spurious_corr.transform import spurious_transform from datasets import load_dataset diff --git a/stable_pretraining/data/spurious_corr/modifiers.py b/stable_pretraining/data/spurious_corr/modifiers.py deleted file mode 100644 index 1a1cd9037..000000000 --- a/stable_pretraining/data/spurious_corr/modifiers.py +++ /dev/null @@ -1,393 +0,0 @@ -"""modifiers.py. - -This module defines the base Modifier class, as well as subclasses for injecting items -(ItemInjection) and HTML tags (HTMLInjection) into text, as well as composing multiple -modifiers (CompositeModifier). -""" - -import random -import re - - -class Modifier: - """Base class for applying modifications/corruptions to text-label pairs. - - Subclasses must implement the __call__ method to define specific transformations. - - Example: - class MyModifier(Modifier): - def __call__(self, text: str, label: Any) -> tuple[str, Any]: - # custom transformation here - return transformed_text, transformed_label - """ - - def __call__(self, text: str, label): - """Apply the transformation to a single text-label pair. - - Args: - text (str): The input text to transform. - label: The associated label. - - Returns: - tuple: (transformed_text, transformed_label) - """ - raise NotImplementedError("Subclasses must implement __call__") - - -class CompositeModifier: - """CompositeModifier chains multiple Modifier instances together. - - Each modifier from the list is applied sequentially to the text. This enables - the combination of various transformations or injections into one composite operation. - """ - - def __init__(self, modifiers: list): - """Initialize a CompositeModifier instance. - - Args: - modifiers (list): A list of modifier instances (subclasses of Modifier) - to be applied sequentially. - """ - self.modifiers = modifiers - - def __call__(self, text: str, label): - """Apply all modifiers in sequence to the given (text, label). - - Args: - text (str): The input text. - label: The associated label. - - Returns: - tuple: The modified (text, label) pair after all transformations. - """ - for modifier in self.modifiers: - text, label = modifier(text, label) - return text, label - - -class ItemInjection(Modifier): - """A Modifier that injects items into text. - - This class supports creation via three different approaches: - - from_list: Using a predefined list of injection items. - - from_file: Reading injection items from a file. - - from_function: Using a custom function to generate injections. - """ - - def __init__( - self, - injection_source, - location: str = "random", - token_proportion: float = 0.1, - seed=None, - _rng=None, - ): - """Initialize an ItemInjection instance. - - Args: - injection_source (callable): A function that returns an injection token. - location (str): Where to inject the token ("beginning", "random", "end"). - token_proportion (float): Proportion of tokens in the text to be affected. - seed (int, optional): Seed for reproducibility. - """ - assert callable(injection_source), "injection_source must be callable" - self.injection_source = injection_source - self.location = location - self.token_proportion = token_proportion - self.rng = _rng or random.Random(seed) - - assert 0 <= token_proportion <= 1, "token_proportion must be between 0 and 1" - assert location in {"beginning", "random", "end"}, ( - "location must be 'beginning', 'random', or 'end'" - ) - - def __call__(self, text: str, label): - """Inject tokens into the text at specified locations. - - Args: - text (str): The input text to modify. - label: The original label (unchanged). - - Returns: - tuple: The modified text and the original label. - """ - words = text.split() - num_tokens = len(words) - - # Ensure at least one token is injected - num_to_inject = max(1, int(num_tokens * self.token_proportion)) - - injections = [self.injection_source() for _ in range(num_to_inject)] - - if self.location == "beginning": - words = injections + words - elif self.location == "end": - words = words + injections - elif self.location == "random": - for injection in injections: - pos = self.rng.randint(0, len(words)) - words.insert(pos, injection) - - return " ".join(words), label # return modified text and unchanged label - - @classmethod - def from_list( - cls, - items: list, - location: str = "random", - token_proportion: float = 0.1, - seed=None, - ): - """Create an ItemInjection instance using a predefined list of tokens. - - Args: - items (list): List of token strings to choose from. - location (str): Where to inject tokens ("beginning", "random", "end"). - token_proportion (float): Proportion of text tokens to be affected. - seed (int, optional): Seed for reproducibility. - - Returns: - ItemInjection: Configured instance. - """ - rng = random.Random(seed) - - def injection_source(): - return rng.choice(items) - - return cls( - injection_source, - location=location, - token_proportion=token_proportion, - seed=seed, - _rng=rng, - ) - - @classmethod - def from_file( - cls, - file_path: str, - location: str = "random", - token_proportion: float = 0.1, - seed=None, - ): - """Create an ItemInjection instance using tokens read from a file. - - Each non-empty line becomes a potential injection item. - - Args: - file_path (str): Path to the file with one token per line. - location (str): Where to inject tokens. - token_proportion (float): Proportion of tokens to inject. - seed (int, optional): Seed for reproducibility. - - Returns: - ItemInjection: Configured instance. - """ - with open(file_path, "r", encoding="utf-8") as file: - items = [line.strip() for line in file if line.strip()] - - rng = random.Random(seed) - - def injection_source(): - return rng.choice(items) - - return cls( - injection_source, - location=location, - token_proportion=token_proportion, - _rng=rng, - ) - - @classmethod - def from_function( - cls, - injection_func, - location: str = "random", - token_proportion: float = 0.1, - seed=None, - ): - """Create an ItemInjection instance using a custom function to generate injections. - - Args: - injection_func (callable): Function that returns a new injection token each time. - location (str): Where to inject tokens. - token_proportion (float): Proportion of text to inject into. - seed (int, optional): Seed for reproducibility (used only for insertion position). - - Returns: - ItemInjection: Configured instance. - """ - assert callable(injection_func), "injection_func must be callable" - return cls( - injection_func, - location=location, - token_proportion=token_proportion, - seed=seed, - ) - - -class HTMLInjection(Modifier): - """A Modifier that injects html into text. - - This class supports creation via two different approaches: - - from_list: Using a predefined list of injection items. - - from_file: Reading injection items from a file. - """ - - def __init__( - self, - file_path: str, - location: str = "random", - level: int = None, - token_proportion: float = None, - seed=None, - ): - with open(file_path, "r", encoding="utf-8") as f: - self.tags = [line.strip() for line in f if line.strip()] - self.location = location - self.level = level - self.token_proportion = token_proportion - self.rng = random.Random(seed) - - if token_proportion is not None: - assert 0 < token_proportion <= 1, "token_proportion must be between 0 and 1" - - @classmethod - def from_file( - cls, - file_path: str, - location: str = "random", - level: int = None, - token_proportion: float = None, - seed=None, - ): - return cls( - file_path, - location=location, - level=level, - token_proportion=token_proportion, - seed=seed, - ) - - @classmethod - def from_list( - cls, - tags: list, - location: str = "random", - level: int = None, - token_proportion: float = None, - seed=None, - ): - instance = cls.__new__(cls) - instance.tags = tags - instance.location = location - instance.level = level - instance.token_proportion = token_proportion - instance.rng = random.Random(seed) - - if token_proportion is not None: - assert 0 < token_proportion <= 1, "token_proportion must be between 0 and 1" - - return instance - - def _choose_tag(self): - """Randomly choose a tag from the loaded list. - - Returns: - tuple: (opening_tag, closing_tag or None) - """ - line = self.rng.choice(self.tags) - parts = line.split() - if len(parts) >= 2: - return parts[0], parts[1] - else: - return parts[0], None - - def _inject_into_tokens(self, tokens, location): - tokens = tokens[:] - n = len(tokens) - - if self.token_proportion is None: - opening, closing = self._choose_tag() - return self._inject_with_tags(tokens, opening, closing, location) - - # Otherwise, inject up to token_proportion of total tokens - num_insertions = max(1, int(n * self.token_proportion)) - for _ in range(num_insertions): - opening, closing = self._choose_tag() - tokens = self._inject_with_tags(tokens, opening, closing, location) - return tokens - - def _inject_with_tags(self, tokens, opening, closing, location): - if location == "beginning": - new_tokens = [opening] + tokens - if closing: - pos = self.rng.randint(1, len(new_tokens)) - new_tokens.insert(pos, closing) - return new_tokens - - elif location == "end": - new_tokens = tokens[:] - pos = self.rng.randint(0, len(new_tokens)) - new_tokens.insert(pos, opening) - if closing: - new_tokens.append(closing) - return new_tokens - - elif location == "random": - new_tokens = tokens[:] - pos_open = self.rng.randint(0, len(new_tokens)) - new_tokens.insert(pos_open, opening) - if closing: - pos_close = self.rng.randint(pos_open + 1, len(new_tokens)) - new_tokens.insert(pos_close, closing) - return new_tokens - - return tokens - - def _inject(self, text, location): - tokens = text.split() - new_tokens = self._inject_into_tokens(tokens, location) - return " ".join(new_tokens) - - def _find_level_span(self, text, level): - """Find the first span inside the desired HTML nesting level. - - Args: - text (str): Input HTML text. - level (int): Desired nesting level. - - Returns: - tuple or None: (start, end) of the content region, or None if not found. - """ - tag_regex = re.compile(r"]*>") - stack = [] - for match in tag_regex.finditer(text): - tag_str = match.group(0) - tag_name = match.group(1) - if not tag_str.startswith(" tuple[str, Any]: + # custom transformation here + return transformed_text, transformed_label + """ + + def __call__(self, text: str, label): + """Apply the transformation to a single text-label pair. + + Args: + text (str): The input text to transform. + label: The associated label. + + Returns: + tuple: (transformed_text, transformed_label) + """ + raise NotImplementedError("Subclasses must implement __call__") + + +class CompositeModifier: + """CompositeModifier chains multiple Modifier instances together. + + Each modifier from the list is applied sequentially to the text. This enables + the combination of various transformations or injections into one composite operation. + """ + + def __init__(self, modifiers: list): + """Initialize a CompositeModifier instance. + + Args: + modifiers (list): A list of modifier instances (subclasses of Modifier) + to be applied sequentially. + """ + self.modifiers = modifiers + + def __call__(self, text: str, label): + """Apply all modifiers in sequence to the given (text, label). + + Args: + text (str): The input text. + label: The associated label. + + Returns: + tuple: The modified (text, label) pair after all transformations. + """ + for modifier in self.modifiers: + text, label = modifier(text, label) + return text, label + + +class ItemInjection(Modifier): + """A Modifier that injects items into text. + + This class supports creation via three different approaches: + - from_list: Using a predefined list of injection items. + - from_file: Reading injection items from a file. + - from_function: Using a custom function to generate injections. + """ + + def __init__( + self, + injection_source, + location: str = "random", + token_proportion: float = 0.1, + seed=None, + _rng=None, + ): + """Initialize an ItemInjection instance. + + Args: + injection_source (callable): A function that returns an injection token. + location (str): Where to inject the token ("beginning", "random", "end"). + token_proportion (float): Proportion of tokens in the text to be affected. + seed (int, optional): Seed for reproducibility. + """ + assert callable(injection_source), "injection_source must be callable" + self.injection_source = injection_source + self.location = location + self.token_proportion = token_proportion + self.rng = _rng or random.Random(seed) + + assert 0 <= token_proportion <= 1, "token_proportion must be between 0 and 1" + assert location in {"beginning", "random", "end"}, ( + "location must be 'beginning', 'random', or 'end'" + ) + + def __call__(self, text: str, label): + """Inject tokens into the text at specified locations. + + Args: + text (str): The input text to modify. + label: The original label (unchanged). + + Returns: + tuple: The modified text and the original label. + """ + words = text.split() + num_tokens = len(words) + + # Ensure at least one token is injected + num_to_inject = max(1, int(num_tokens * self.token_proportion)) + + injections = [self.injection_source() for _ in range(num_to_inject)] + + if self.location == "beginning": + words = injections + words + elif self.location == "end": + words = words + injections + elif self.location == "random": + for injection in injections: + pos = self.rng.randint(0, len(words)) + words.insert(pos, injection) + + return " ".join(words), label # return modified text and unchanged label + + @classmethod + def from_list( + cls, + items: list, + location: str = "random", + token_proportion: float = 0.1, + seed=None, + ): + """Create an ItemInjection instance using a predefined list of tokens. + + Args: + items (list): List of token strings to choose from. + location (str): Where to inject tokens ("beginning", "random", "end"). + token_proportion (float): Proportion of text tokens to be affected. + seed (int, optional): Seed for reproducibility. + + Returns: + ItemInjection: Configured instance. + """ + rng = random.Random(seed) + + def injection_source(): + return rng.choice(items) + + return cls( + injection_source, + location=location, + token_proportion=token_proportion, + seed=seed, + _rng=rng, + ) + + @classmethod + def from_file( + cls, + file_path: str, + location: str = "random", + token_proportion: float = 0.1, + seed=None, + ): + """Create an ItemInjection instance using tokens read from a file. + + Each non-empty line becomes a potential injection item. + + Args: + file_path (str): Path to the file with one token per line. + location (str): Where to inject tokens. + token_proportion (float): Proportion of tokens to inject. + seed (int, optional): Seed for reproducibility. + + Returns: + ItemInjection: Configured instance. + """ + with open(file_path, "r", encoding="utf-8") as file: + items = [line.strip() for line in file if line.strip()] + + rng = random.Random(seed) + + def injection_source(): + return rng.choice(items) + + return cls( + injection_source, + location=location, + token_proportion=token_proportion, + _rng=rng, + ) + + @classmethod + def from_function( + cls, + injection_func, + location: str = "random", + token_proportion: float = 0.1, + seed=None, + ): + """Create an ItemInjection instance using a custom function to generate injections. + + Args: + injection_func (callable): Function that returns a new injection token each time. + location (str): Where to inject tokens. + token_proportion (float): Proportion of text to inject into. + seed (int, optional): Seed for reproducibility (used only for insertion position). + + Returns: + ItemInjection: Configured instance. + """ + assert callable(injection_func), "injection_func must be callable" + return cls( + injection_func, + location=location, + token_proportion=token_proportion, + seed=seed, + ) + + +class HTMLInjection(Modifier): + """A Modifier that injects html into text. + + This class supports creation via two different approaches: + - from_list: Using a predefined list of injection items. + - from_file: Reading injection items from a file. + """ + + def __init__( + self, + file_path: str, + location: str = "random", + level: int = None, + token_proportion: float = None, + seed=None, + ): + with open(file_path, "r", encoding="utf-8") as f: + self.tags = [line.strip() for line in f if line.strip()] + self.location = location + self.level = level + self.token_proportion = token_proportion + self.rng = random.Random(seed) + + if token_proportion is not None: + assert 0 < token_proportion <= 1, "token_proportion must be between 0 and 1" + + @classmethod + def from_file( + cls, + file_path: str, + location: str = "random", + level: int = None, + token_proportion: float = None, + seed=None, + ): + return cls( + file_path, + location=location, + level=level, + token_proportion=token_proportion, + seed=seed, + ) + + @classmethod + def from_list( + cls, + tags: list, + location: str = "random", + level: int = None, + token_proportion: float = None, + seed=None, + ): + instance = cls.__new__(cls) + instance.tags = tags + instance.location = location + instance.level = level + instance.token_proportion = token_proportion + instance.rng = random.Random(seed) + + if token_proportion is not None: + assert 0 < token_proportion <= 1, "token_proportion must be between 0 and 1" + + return instance + + def _choose_tag(self): + """Randomly choose a tag from the loaded list. + + Returns: + tuple: (opening_tag, closing_tag or None) + """ + line = self.rng.choice(self.tags) + parts = line.split() + if len(parts) >= 2: + return parts[0], parts[1] + else: + return parts[0], None + + def _inject_into_tokens(self, tokens, location): + tokens = tokens[:] + n = len(tokens) + + if self.token_proportion is None: + opening, closing = self._choose_tag() + return self._inject_with_tags(tokens, opening, closing, location) + + # Otherwise, inject up to token_proportion of total tokens + num_insertions = max(1, int(n * self.token_proportion)) + for _ in range(num_insertions): + opening, closing = self._choose_tag() + tokens = self._inject_with_tags(tokens, opening, closing, location) + return tokens + + def _inject_with_tags(self, tokens, opening, closing, location): + if location == "beginning": + new_tokens = [opening] + tokens + if closing: + pos = self.rng.randint(1, len(new_tokens)) + new_tokens.insert(pos, closing) + return new_tokens + + elif location == "end": + new_tokens = tokens[:] + pos = self.rng.randint(0, len(new_tokens)) + new_tokens.insert(pos, opening) + if closing: + new_tokens.append(closing) + return new_tokens + + elif location == "random": + new_tokens = tokens[:] + pos_open = self.rng.randint(0, len(new_tokens)) + new_tokens.insert(pos_open, opening) + if closing: + pos_close = self.rng.randint(pos_open + 1, len(new_tokens)) + new_tokens.insert(pos_close, closing) + return new_tokens + + return tokens + + def _inject(self, text, location): + tokens = text.split() + new_tokens = self._inject_into_tokens(tokens, location) + return " ".join(new_tokens) + + def _find_level_span(self, text, level): + """Find the first span inside the desired HTML nesting level. + + Args: + text (str): Input HTML text. + level (int): Desired nesting level. + + Returns: + tuple or None: (start, end) of the content region, or None if not found. + """ + tag_regex = re.compile(r"]*>") + stack = [] + for match in tag_regex.finditer(text): + tag_str = match.group(0) + tag_name = match.group(1) + if not tag_str.startswith(" Date: Mon, 13 Oct 2025 16:14:23 -0400 Subject: [PATCH 07/12] minor bug fixed from reformatting --- stable_pretraining/data/spurious_corr/__init__.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/stable_pretraining/data/spurious_corr/__init__.py b/stable_pretraining/data/spurious_corr/__init__.py index 9d47c78d0..ed8cdf61e 100644 --- a/stable_pretraining/data/spurious_corr/__init__.py +++ b/stable_pretraining/data/spurious_corr/__init__.py @@ -6,7 +6,7 @@ text, and utilities for printing and highlighting text. """ -from .modifiers import ( +from ..transforms import ( Modifier as Modifier, CompositeModifier as CompositeModifier, ItemInjection as ItemInjection, From 4a5bdd90f368744f984569a2c1d1cc74d813c35a Mon Sep 17 00:00:00 2001 From: "marcel_mateos_salles@brown.edu" Date: Mon, 13 Oct 2025 16:23:24 -0400 Subject: [PATCH 08/12] fixing import errors --- stable_pretraining/data/transforms.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/stable_pretraining/data/transforms.py b/stable_pretraining/data/transforms.py index 38fae63d0..8c98841b8 100644 --- a/stable_pretraining/data/transforms.py +++ b/stable_pretraining/data/transforms.py @@ -1,7 +1,8 @@ from contextlib import contextmanager from itertools import islice -from random import getstate, random, setstate +from random import getstate, setstate from random import seed as rseed +import random import re from typing import Any, Dict, List, Optional, Sequence, Tuple, Union From 65a203ed5b176f2ee77c32b38843abbec672fdb5 Mon Sep 17 00:00:00 2001 From: "marcel_mateos_salles@brown.edu" Date: Mon, 13 Oct 2025 16:43:27 -0400 Subject: [PATCH 09/12] removed the files for spurious text to huggingface --- .../data/spurious_corr/data/colors.txt | 52 ----- .../data/spurious_corr/data/countries.txt | 194 ----------------- .../spurious_corr/data/double_exclamation.txt | 1 - .../data/spurious_corr/data/exclamation.txt | 1 - .../data/spurious_corr/data/html_tags.txt | 106 ---------- .../data/spurious_corr/data/random.txt | 4 - .../spurious_corr/data/two_hundred_dates.txt | 200 ------------------ 7 files changed, 558 deletions(-) delete mode 100644 stable_pretraining/data/spurious_corr/data/colors.txt delete mode 100644 stable_pretraining/data/spurious_corr/data/countries.txt delete mode 100644 stable_pretraining/data/spurious_corr/data/double_exclamation.txt delete mode 100644 stable_pretraining/data/spurious_corr/data/exclamation.txt delete mode 100644 stable_pretraining/data/spurious_corr/data/html_tags.txt delete mode 100644 stable_pretraining/data/spurious_corr/data/random.txt delete mode 100644 stable_pretraining/data/spurious_corr/data/two_hundred_dates.txt diff --git a/stable_pretraining/data/spurious_corr/data/colors.txt b/stable_pretraining/data/spurious_corr/data/colors.txt deleted file mode 100644 index 82dfac34e..000000000 --- a/stable_pretraining/data/spurious_corr/data/colors.txt +++ /dev/null @@ -1,52 +0,0 @@ -Red -Blue -Green -Yellow -Orange -Purple -Pink -Brown -Black -White -Gray -Cyan -Magenta -Beige -Maroon -Olive -Navy -Teal -Lavender -Turquoise -Gold -Silver -Bronze -Ivory -Coral -Aqua -Crimson -Fuchsia -Amber -Chartreuse -Indigo -Emerald -Violet -Peach -Mint -Lilac -Ruby -Sapphire -Topaz -Periwinkle -Charcoal -Khaki -Plum -Scarlet -Azure -Tan -Cobalt -Mauve -Rust -Sand -Aquamarine -Burgundy diff --git a/stable_pretraining/data/spurious_corr/data/countries.txt b/stable_pretraining/data/spurious_corr/data/countries.txt deleted file mode 100644 index 7e4619afd..000000000 --- a/stable_pretraining/data/spurious_corr/data/countries.txt +++ /dev/null @@ -1,194 +0,0 @@ -Afghanistan -Albania -Algeria -Andorra -Angola -Antigua and Barbuda -Argentina -Armenia -Australia -Austria -Azerbaijan -Bahamas -Bahrain -Bangladesh -Barbados -Belarus -Belgium -Belize -Benin -Bhutan -Bolivia -Bosnia and Herzegovina -Botswana -Brazil -Brunei -Bulgaria -Burkina Faso -Burundi -Cabo Verde -Cambodia -Cameroon -Canada -Central African Republic -Chad -Chile -China -Colombia -Comoros -Congo (Congo-Brazzaville) -Costa Rica -Croatia -Cuba -Cyprus -Czechia (Czech Republic) -Denmark -Djibouti -Dominica -Dominican Republic -Ecuador -Egypt -El Salvador -Equatorial Guinea -Eritrea -Estonia -Eswatini (fmr. Swaziland) -Ethiopia -Fiji -Finland -France -Gabon -Gambia -Georgia -Germany -Ghana -Greece -Grenada -Guatemala -Guinea -Guinea-Bissau -Guyana -Haiti -Holy See -Honduras -Hungary -Iceland -India -Indonesia -Iran -Iraq -Ireland -Israel -Italy -Jamaica -Japan -Jordan -Kazakhstan -Kenya -Kiribati -Korea (North) -Korea (South) -Kosovo -Kuwait -Kyrgyzstan -Laos -Latvia -Lebanon -Lesotho -Liberia -Libya -Liechtenstein -Lithuania -Luxembourg -Madagascar -Malawi -Malaysia -Maldives -Mali -Malta -Marshall Islands -Mauritania -Mauritius -Mexico -Micronesia -Moldova -Monaco -Mongolia -Montenegro -Morocco -Mozambique -Myanmar -Namibia -Nauru -Nepal -Netherlands -New Zealand -Nicaragua -Niger -Nigeria -North Macedonia -Norway -Oman -Pakistan -Palau -Palestine State -Panama -Papua New Guinea -Paraguay -Peru -Philippines -Poland -Portugal -Qatar -Romania -Russia -Rwanda -Saint Kitts and Nevis -Saint Lucia -Saint Vincent and the Grenadines -Samoa -San Marino -Sao Tome and Principe -Saudi Arabia -Senegal -Serbia -Seychelles -Sierra Leone -Singapore -Slovakia -Slovenia -Solomon Islands -Somalia -South Africa -South Sudan -Spain -Sri Lanka -Sudan -Suriname -Sweden -Switzerland -Syria -Tajikistan -Tanzania -Thailand -Timor-Leste -Togo -Tonga -Trinidad and Tobago -Tunisia -Turkey -Turkmenistan -Tuvalu -Uganda -Ukraine -United Arab Emirates -United Kingdom -United States of America -Uruguay -Uzbekistan -Vanuatu -Venezuela -Vietnam -Yemen -Zambia -Zimbabwe diff --git a/stable_pretraining/data/spurious_corr/data/double_exclamation.txt b/stable_pretraining/data/spurious_corr/data/double_exclamation.txt deleted file mode 100644 index 79d895682..000000000 --- a/stable_pretraining/data/spurious_corr/data/double_exclamation.txt +++ /dev/null @@ -1 +0,0 @@ -!! diff --git a/stable_pretraining/data/spurious_corr/data/exclamation.txt b/stable_pretraining/data/spurious_corr/data/exclamation.txt deleted file mode 100644 index cdf4cb4fe..000000000 --- a/stable_pretraining/data/spurious_corr/data/exclamation.txt +++ /dev/null @@ -1 +0,0 @@ -! diff --git a/stable_pretraining/data/spurious_corr/data/html_tags.txt b/stable_pretraining/data/spurious_corr/data/html_tags.txt deleted file mode 100644 index 56b4a07ac..000000000 --- a/stable_pretraining/data/spurious_corr/data/html_tags.txt +++ /dev/null @@ -1,106 +0,0 @@ - - - - - - - - - -

    -

    -

    -

    -
    -
    -

    -
    -
    -
     
    - - - - - - - - - - - - - - - -
    - - - - - - - - -
    -
    -
  • -
    -
    -
    -
    - - - - -
    -
    - - - - - -
    - - - - - - - - - -
    - - - - - - - -
    - - - - - -
    - - - - - - - - - - - -
    -
    -
    -
    -
    - -
    - -
    diff --git a/stable_pretraining/data/spurious_corr/data/random.txt b/stable_pretraining/data/spurious_corr/data/random.txt deleted file mode 100644 index 8422d40f1..000000000 --- a/stable_pretraining/data/spurious_corr/data/random.txt +++ /dev/null @@ -1,4 +0,0 @@ -A -B -C -D diff --git a/stable_pretraining/data/spurious_corr/data/two_hundred_dates.txt b/stable_pretraining/data/spurious_corr/data/two_hundred_dates.txt deleted file mode 100644 index 853060bf0..000000000 --- a/stable_pretraining/data/spurious_corr/data/two_hundred_dates.txt +++ /dev/null @@ -1,200 +0,0 @@ -1975-02-20 -1975-04-19 -1976-11-02 -1976-11-21 -1976-12-10 -1976-12-30 -1977-05-16 -1977-07-21 -1977-10-17 -1977-10-27 -1977-10-31 -1977-12-04 -1978-05-23 -1979-04-20 -1979-07-29 -1979-08-30 -1979-10-09 -1979-10-25 -1979-11-21 -1980-04-08 -1980-05-11 -1980-06-30 -1980-09-26 -1981-02-17 -1981-03-12 -1981-03-18 -1981-05-09 -1981-08-01 -1982-03-12 -1982-03-13 -1982-04-13 -1982-09-27 -1982-11-05 -1982-11-21 -1982-12-07 -1983-01-26 -1983-06-03 -1983-06-07 -1983-09-14 -1983-09-21 -1983-10-26 -1983-11-06 -1984-01-23 -1984-06-07 -1984-08-19 -1984-10-25 -1984-11-21 -1984-11-30 -1985-02-20 -1985-07-26 -1985-10-23 -1986-01-18 -1986-04-01 -1986-08-07 -1986-11-08 -1986-11-16 -1986-12-24 -1987-02-27 -1987-10-16 -1988-01-21 -1988-05-03 -1989-03-11 -1989-08-12 -1989-08-27 -1989-09-27 -1990-02-09 -1990-08-14 -1990-12-24 -1991-01-08 -1991-02-05 -1991-10-11 -1991-11-29 -1992-02-11 -1992-02-18 -1992-06-30 -1992-08-07 -1992-09-28 -1992-11-24 -1993-06-16 -1994-03-21 -1994-06-13 -1994-06-27 -1994-09-26 -1994-10-22 -1995-02-11 -1995-06-12 -1995-06-21 -1995-07-02 -1995-07-17 -1995-10-18 -1995-10-27 -1996-07-10 -1996-07-29 -1998-01-07 -1998-02-18 -1998-03-06 -1998-06-24 -1998-08-06 -1998-09-15 -1998-12-21 -1999-03-17 -1999-05-30 -1999-08-01 -2000-01-07 -2000-03-13 -2000-04-30 -2000-06-15 -2000-07-29 -2000-09-17 -2000-12-13 -2000-12-22 -2000-12-30 -2001-01-29 -2001-03-04 -2001-08-04 -2002-04-19 -2002-06-07 -2002-08-24 -2002-09-25 -2003-01-11 -2003-05-02 -2004-01-11 -2004-05-02 -2004-05-31 -2004-11-11 -2004-12-31 -2005-02-03 -2005-02-20 -2005-04-10 -2005-07-21 -2005-10-06 -2006-05-25 -2006-07-22 -2006-09-21 -2006-12-29 -2007-04-06 -2007-04-25 -2007-08-26 -2007-09-03 -2008-01-08 -2008-06-01 -2008-06-30 -2008-10-17 -2009-02-28 -2009-10-10 -2010-02-01 -2010-03-26 -2010-06-18 -2011-01-16 -2011-02-24 -2011-03-15 -2011-04-06 -2011-07-27 -2011-10-20 -2011-12-20 -2012-09-10 -2012-10-04 -2013-04-04 -2013-07-15 -2013-11-24 -2014-03-12 -2014-03-19 -2014-11-19 -2015-08-05 -2016-01-26 -2016-01-29 -2016-03-05 -2016-06-05 -2016-12-26 -2017-04-18 -2017-05-21 -2017-09-01 -2017-09-04 -2018-02-24 -2018-03-13 -2018-04-21 -2018-07-20 -2018-10-13 -2019-06-05 -2019-07-14 -2019-08-22 -2019-10-30 -2020-05-30 -2020-08-23 -2020-09-06 -2020-11-27 -2021-06-10 -2021-07-04 -2021-09-15 -2021-10-16 -2021-11-04 -2022-06-28 -2022-08-09 -2022-08-16 -2023-08-29 -2024-03-23 -2024-07-03 -2024-08-06 -2024-12-28 -2025-11-14 From fc7e5151d8e19cad439c4fe76c0835d152a7a1e0 Mon Sep 17 00:00:00 2001 From: "marcel_mateos_salles@brown.edu" Date: Tue, 14 Oct 2025 15:27:56 -0400 Subject: [PATCH 10/12] beginning of visual spurious injection framework --- stable_pretraining/data/transforms.py | 2276 +++++++++++++------------ 1 file changed, 1217 insertions(+), 1059 deletions(-) diff --git a/stable_pretraining/data/transforms.py b/stable_pretraining/data/transforms.py index 8c98841b8..22c7cddd1 100644 --- a/stable_pretraining/data/transforms.py +++ b/stable_pretraining/data/transforms.py @@ -5,7 +5,6 @@ import random import re from typing import Any, Dict, List, Optional, Sequence, Tuple, Union - import numpy as np import PIL.Image import torch @@ -21,1341 +20,1500 @@ # ============================================================ -# ===================== TEXT MODIFIERS ======================= +# ===================== Images =============================== # ============================================================ -class Modifier: - """Base class for applying modifications/corruptions to text-label pairs. +class Transform(v2.Transform): + """Base transform class extending torchvision v2.Transform with nested data handling.""" - Subclasses must implement the __call__ method to define specific transformations. + def nested_get(self, v, name): + if name == "": + return v + i = name.split(".") + if i[0].isnumeric(): + i[0] = int(i[0]) + return self.nested_get(v[i[0]], ".".join(i[1:])) - Example: - class MyModifier(Modifier): - def __call__(self, text: str, label: Any) -> tuple[str, Any]: - # custom transformation here - return transformed_text, transformed_label - """ + def nested_set(self, original, value, name): + if "." not in name: + if name.isnumeric(): + name = int(name) + original[name] = value + else: + i = name.split(".") + if i[0].isnumeric(): + i[0] = int(i[0]) + self.nested_set(original[i[0]], value, ".".join(i[1:])) - def __call__(self, text: str, label): - """Apply the transformation to a single text-label pair. + def get_name(self, x): + base = self.name + assert "_" not in base + if base not in x: + return base + ctr = 0 + while f"{base}_{ctr}" in base: + ctr += 1 + return f"{base}_{ctr}" - Args: - text (str): The input text to transform. - label: The associated label. + @property + def name(self): + return self.__class__.__name__ - Returns: - tuple: (transformed_text, transformed_label) - """ - raise NotImplementedError("Subclasses must implement __call__") +@torch.jit.unused +def to_image( + input: Union[torch.Tensor, PIL.Image.Image, np.ndarray], +) -> tv_tensors.Image: + """See :class:`~torchvision.transforms.v2.ToImage` for details.""" + if isinstance(input, np.ndarray): + output = torch.from_numpy(np.atleast_3d(input)).transpose(-3, -1).contiguous() + elif isinstance(input, PIL.Image.Image): + output = torchvision.transforms.functional.pil_to_tensor(input) + elif isinstance(input, torch.Tensor): + output = input + else: + raise TypeError( + f"Input can either be a pure Tensor, a numpy array, or a PIL image, but got {type(input)} instead." + ) + return tv_tensors.Image(output) -class CompositeModifier: - """CompositeModifier chains multiple Modifier instances together. - Each modifier from the list is applied sequentially to the text. This enables - the combination of various transformations or injections into one composite operation. - """ +class ToImage(Transform): + """Convert input to image tensor with optional normalization.""" - def __init__(self, modifiers: list): - """Initialize a CompositeModifier instance. + def __init__( + self, + dtype=torch.float32, + scale=True, + mean=None, + std=None, + source: str = "image", + target: str = "image", + ): + super().__init__() + t = [to_image, v2.ToDtype(dtype, scale=scale)] + if mean is not None and std is not None: + t.append(v2.Normalize(mean=mean, std=std)) + self.t = v2.Compose(t) + self.source = source + self.target = target - Args: - modifiers (list): A list of modifier instances (subclasses of Modifier) - to be applied sequentially. - """ - self.modifiers = modifiers + def __call__(self, x): + self.nested_set(x, self.t(self.nested_get(x, self.source)), self.target) + return x - def __call__(self, text: str, label): - """Apply all modifiers in sequence to the given (text, label). - Args: - text (str): The input text. - label: The associated label. +class RandomGrayscale(Transform, v2.RandomGrayscale): + """Randomly convert image to grayscale with given probability.""" - Returns: - tuple: The modified (text, label) pair after all transformations. - """ - for modifier in self.modifiers: - text, label = modifier(text, label) - return text, label + def __init__(self, p=0.1, source: str = "image", target: str = "image"): + super().__init__(p) + self.source = source + self.target = target + def _get_params(self, inp: List[Any]) -> Dict[str, Any]: + num_input_channels, *_ = query_chw([inp]) + return dict(num_input_channels=num_input_channels) -class ItemInjection(Modifier): - """A Modifier that injects items into text. + def __call__(self, x) -> Any: + if self.p < 1 and torch.rand(1) >= self.p: + x[self.get_name(x)] = False + self.nested_set(x, self.nested_get(x, self.source), self.target) + return x + channels, *_ = query_chw([self.nested_get(x, self.source)]) + self.nested_set( + x, + F.rgb_to_grayscale( + self.nested_get(x, self.source), num_output_channels=channels + ), + self.target, + ) + x[self.get_name(x)] = True + return x - This class supports creation via three different approaches: - - from_list: Using a predefined list of injection items. - - from_file: Reading injection items from a file. - - from_function: Using a custom function to generate injections. - """ - def __init__( - self, - injection_source, - location: str = "random", - token_proportion: float = 0.1, - seed=None, - _rng=None, - ): - """Initialize an ItemInjection instance. +class Lambda(Transform): + """Applies a lambda callable to target key and store it in source.""" - Args: - injection_source (callable): A function that returns an injection token. - location (str): Where to inject the token ("beginning", "random", "end"). - token_proportion (float): Proportion of tokens in the text to be affected. - seed (int, optional): Seed for reproducibility. - """ - assert callable(injection_source), "injection_source must be callable" - self.injection_source = injection_source - self.location = location - self.token_proportion = token_proportion - self.rng = _rng or random.Random(seed) + def __init__(self, lambd, source: str = "image", target: str = "image"): + super().__init__() + self.source = source + self.target = target + self.lambd = lambd - assert 0 <= token_proportion <= 1, "token_proportion must be between 0 and 1" - assert location in {"beginning", "random", "end"}, ( - "location must be 'beginning', 'random', or 'end'" - ) + def __call__(self, x) -> Any: + self.nested_set(x, self.lambd(x), self.target) + return x - def __call__(self, text: str, label): - """Inject tokens into the text at specified locations. - Args: - text (str): The input text to modify. - label: The original label (unchanged). +class RoutingTransform(Transform): + """Applies a routing callable to conditionally apply a transform from many candidates.""" - Returns: - tuple: The modified text and the original label. - """ - words = text.split() - num_tokens = len(words) + def __init__(self, router: callable, transforms: Union[list, tuple, dict]): + self.router = router + self.transforms = transforms - # Ensure at least one token is injected - num_to_inject = max(1, int(num_tokens * self.token_proportion)) + def __call__(self, x) -> Any: + route = self.router(x) + return self.transforms[route](x) - injections = [self.injection_source() for _ in range(num_to_inject)] - if self.location == "beginning": - words = injections + words - elif self.location == "end": - words = words + injections - elif self.location == "random": - for injection in injections: - pos = self.rng.randint(0, len(words)) - words.insert(pos, injection) +class WrapTorchTransform(Transform, v2.Lambda): + """Applies a lambda callable to target key and store it in source.""" - return " ".join(words), label # return modified text and unchanged label + def __init__(self, transform, source: str = "image", target: str = "image"): + super().__init__(transform) + self.source = source + self.target = target - @classmethod - def from_list( - cls, - items: list, - location: str = "random", - token_proportion: float = 0.1, - seed=None, - ): - """Create an ItemInjection instance using a predefined list of tokens. + def __call__(self, x) -> Any: + self.nested_set( + x, super().__call__(self.nested_get(x, self.source)), self.target + ) + return x - Args: - items (list): List of token strings to choose from. - location (str): Where to inject tokens ("beginning", "random", "end"). - token_proportion (float): Proportion of text tokens to be affected. - seed (int, optional): Seed for reproducibility. - Returns: - ItemInjection: Configured instance. - """ - rng = random.Random(seed) +class RandomSolarize(Transform, v2.RandomSolarize): + """Randomly solarize image by inverting pixel values above threshold.""" - def injection_source(): - return rng.choice(items) + def __init__(self, threshold, p=0.5, source: str = "image", target: str = "image"): + super().__init__(threshold, p) + self.source = source + self.target = target - return cls( - injection_source, - location=location, - token_proportion=token_proportion, - seed=seed, - _rng=rng, + def __call__(self, x) -> Any: + if self.p < 1 and torch.rand(1) >= self.p: + x[self.get_name(x)] = False + return x + self.nested_set( + x, F.solarize(self.nested_get(x, self.source), self.threshold), self.target ) + x[self.get_name(x)] = True + return x - @classmethod - def from_file( - cls, - file_path: str, - location: str = "random", - token_proportion: float = 0.1, - seed=None, + +class GaussianBlur(Transform, v2.GaussianBlur): + """Apply Gaussian blur to image with random sigma values.""" + + _NAMES = ["sigma_x", "sigma_y"] + + def __init__( + self, + kernel_size, + sigma=(0.1, 2.0), + p=1, + source: str = "image", + target: str = "image", ): - """Create an ItemInjection instance using tokens read from a file. + super().__init__(kernel_size, sigma) + self.p = p + self.source = source + self.target = target - Each non-empty line becomes a potential injection item. + def __call__(self, x) -> Any: + if self.p < 1 and torch.rand(1) >= self.p: + x[self.get_name(x)] = torch.zeros((2,)) + return x + params = self.make_params([]) + self.nested_set( + x, self.transform(self.nested_get(x, self.source), params), self.target + ) + x[self.get_name(x)] = torch.Tensor(params["sigma"]) + return x - Args: - file_path (str): Path to the file with one token per line. - location (str): Where to inject tokens. - token_proportion (float): Proportion of tokens to inject. - seed (int, optional): Seed for reproducibility. - Returns: - ItemInjection: Configured instance. - """ - with open(file_path, "r", encoding="utf-8") as file: - items = [line.strip() for line in file if line.strip()] +class PILGaussianBlur(Transform): + """PIL-based Gaussian blur transform with random sigma sampling.""" - rng = random.Random(seed) + _NAMES = ["sigma_x", "sigma_y"] - def injection_source(): - return rng.choice(items) + def __init__(self, sigma=None, p=1, source: str = "image", target: str = "image"): + """Gaussian blur as a callable object. - return cls( - injection_source, - location=location, - token_proportion=token_proportion, - _rng=rng, - ) + Args: + sigma (Sequence[float]): range to sample the radius of the gaussian blur filter. + Defaults to [0.1, 2.0]. + p (float): probability of applying the transform. + source (str): source key in the data dictionary. + target (str): target key in the data dictionary. + """ + if sigma is None: + sigma = [0.1, 2.0] - @classmethod - def from_function( - cls, - injection_func, - location: str = "random", - token_proportion: float = 0.1, - seed=None, - ): - """Create an ItemInjection instance using a custom function to generate injections. + self.sigma = sigma + self.p = p + self.source = source + self.target = target + + def __call__(self, x): + """Applies gaussian blur to an input image. Args: - injection_func (callable): Function that returns a new injection token each time. - location (str): Where to inject tokens. - token_proportion (float): Proportion of text to inject into. - seed (int, optional): Seed for reproducibility (used only for insertion position). + x (dict): Data dictionary containing the image to transform. Returns: - ItemInjection: Configured instance. + dict: Data dictionary with blurred image. """ - assert callable(injection_func), "injection_func must be callable" - return cls( - injection_func, - location=location, - token_proportion=token_proportion, - seed=seed, + if self.p < 1 and torch.rand(1) >= self.p: + x[self.get_name(x)] = torch.zeros((1,)) + return x + sigma = torch.rand((1,)) * (self.sigma[1] - self.sigma[0]) + self.sigma[0] + x[self.get_name(x)] = sigma + self.nested_set( + x, + self.nested_get(x, self.source).filter( + ImageFilter.GaussianBlur(radius=sigma.item()) + ), + self.target, ) + return x -class HTMLInjection(Modifier): - """A Modifier that injects html into text. - - This class supports creation via two different approaches: - - from_list: Using a predefined list of injection items. - - from_file: Reading injection items from a file. - """ +class UniformTemporalSubsample(Transform): + """``nn.Module`` wrapper for ``pytorchvideo.transforms.functional.uniform_temporal_subsample``.""" def __init__( self, - file_path: str, - location: str = "random", - level: int = None, - token_proportion: float = None, - seed=None, + num_samples: int, + temporal_dim: int = -3, + source: str = "video", + target: str = "video", ): - with open(file_path, "r", encoding="utf-8") as f: - self.tags = [line.strip() for line in f if line.strip()] - self.location = location - self.level = level - self.token_proportion = token_proportion - self.rng = random.Random(seed) - - if token_proportion is not None: - assert 0 < token_proportion <= 1, "token_proportion must be between 0 and 1" + super().__init__(num_samples, temporal_dim) + self.source = source + self.target = target - @classmethod - def from_file( - cls, - file_path: str, - location: str = "random", - level: int = None, - token_proportion: float = None, - seed=None, - ): - return cls( - file_path, - location=location, - level=level, - token_proportion=token_proportion, - seed=seed, + def forward(self, x: dict) -> torch.Tensor: + self.nested_set( + x, super().forward(self, self.nested_get(x, self.source)), self.target ) + return x - @classmethod - def from_list( - cls, - tags: list, - location: str = "random", - level: int = None, - token_proportion: float = None, - seed=None, - ): - instance = cls.__new__(cls) - instance.tags = tags - instance.location = location - instance.level = level - instance.token_proportion = token_proportion - instance.rng = random.Random(seed) - if token_proportion is not None: - assert 0 < token_proportion <= 1, "token_proportion must be between 0 and 1" +class RandomContiguousTemporalSampler(Transform): + """Randomly sample contiguous frames from a video sequence.""" - return instance + def __init__(self, source, target, num_frames, frame_subsampling: int = 1): + self.source = source + self.target = target + self.num_frames = num_frames + self.frame_subsampling = frame_subsampling - def _choose_tag(self): - """Randomly choose a tag from the loaded list. + def __call__(self, x): + metadata = self.nested_get(x, self.source).get_metadata() + T = int(metadata["video"]["duration"][0] * metadata["video"]["fps"][0]) + covering = self.num_frames * self.frame_subsampling + start = torch.randint(low=0, high=T - covering, size=(1,)).item() + video_frames = [] # video frame buffer - Returns: - tuple: (opening_tag, closing_tag or None) - """ - line = self.rng.choice(self.tags) - parts = line.split() - if len(parts) >= 2: - return parts[0], parts[1] - else: - return parts[0], None + # Seek and return frames + count = 0 + for frame in islice( + self.nested_get(x, self.source).seek(start / metadata["video"]["fps"][0]), + covering, + ): + if count % self.frame_subsampling == 0: + video_frames.append(frame["data"]) + count += 1 + # Stack it into a tensor + self.nested_set(x, torch.stack(video_frames, 0), self.target) + x[self.get_name(x)] = start + return x - def _inject_into_tokens(self, tokens, location): - tokens = tokens[:] - n = len(tokens) - if self.token_proportion is None: - opening, closing = self._choose_tag() - return self._inject_with_tags(tokens, opening, closing, location) +class RGB(Transform, v2.RGB): + """Convert image to RGB format.""" - # Otherwise, inject up to token_proportion of total tokens - num_insertions = max(1, int(n * self.token_proportion)) - for _ in range(num_insertions): - opening, closing = self._choose_tag() - tokens = self._inject_with_tags(tokens, opening, closing, location) - return tokens + def __init__(self, source: str = "image", target: str = "image"): + super().__init__() + self.source = source + self.target = target - def _inject_with_tags(self, tokens, opening, closing, location): - if location == "beginning": - new_tokens = [opening] + tokens - if closing: - pos = self.rng.randint(1, len(new_tokens)) - new_tokens.insert(pos, closing) - return new_tokens + def __call__(self, x): + self.nested_set( + x, F.grayscale_to_rgb(self.nested_get(x, self.source)), self.target + ) + return x - elif location == "end": - new_tokens = tokens[:] - pos = self.rng.randint(0, len(new_tokens)) - new_tokens.insert(pos, opening) - if closing: - new_tokens.append(closing) - return new_tokens - elif location == "random": - new_tokens = tokens[:] - pos_open = self.rng.randint(0, len(new_tokens)) - new_tokens.insert(pos_open, opening) - if closing: - pos_close = self.rng.randint(pos_open + 1, len(new_tokens)) - new_tokens.insert(pos_close, closing) - return new_tokens - - return tokens - - def _inject(self, text, location): - tokens = text.split() - new_tokens = self._inject_into_tokens(tokens, location) - return " ".join(new_tokens) - - def _find_level_span(self, text, level): - """Find the first span inside the desired HTML nesting level. - - Args: - text (str): Input HTML text. - level (int): Desired nesting level. - - Returns: - tuple or None: (start, end) of the content region, or None if not found. - """ - tag_regex = re.compile(r"]*>") - stack = [] - for match in tag_regex.finditer(text): - tag_str = match.group(0) - tag_name = match.group(1) - if not tag_str.startswith(" None: + super().__init__(size, interpolation, max_size, antialias) + self.source = source + self.target = target -@torch.jit.unused -def to_image( - input: Union[torch.Tensor, PIL.Image.Image, np.ndarray], -) -> tv_tensors.Image: - """See :class:`~torchvision.transforms.v2.ToImage` for details.""" - if isinstance(input, np.ndarray): - output = torch.from_numpy(np.atleast_3d(input)).transpose(-3, -1).contiguous() - elif isinstance(input, PIL.Image.Image): - output = torchvision.transforms.functional.pil_to_tensor(input) - elif isinstance(input, torch.Tensor): - output = input - else: - raise TypeError( - f"Input can either be a pure Tensor, a numpy array, or a PIL image, but got {type(input)} instead." + def __call__(self, x): + self.nested_set( + x, self.transform(self.nested_get(x, self.source), []), self.target ) - return tv_tensors.Image(output) + return x -class ToImage(Transform): - """Convert input to image tensor with optional normalization.""" +class ColorJitter(Transform, v2.ColorJitter): + """Randomly change brightness, contrast, saturation, and hue of an image.""" def __init__( self, - dtype=torch.float32, - scale=True, - mean=None, - std=None, + brightness=None, + contrast=None, + saturation=None, + hue=None, + p=1, source: str = "image", target: str = "image", ): - super().__init__() - t = [to_image, v2.ToDtype(dtype, scale=scale)] - if mean is not None and std is not None: - t.append(v2.Normalize(mean=mean, std=std)) - self.t = v2.Compose(t) + super().__init__(brightness, contrast, saturation, hue) + self.p = p self.source = source self.target = target - def __call__(self, x): - self.nested_set(x, self.t(self.nested_get(x, self.source)), self.target) + def __call__(self, x) -> Any: + if self.p < 1 and torch.rand(1) > self.p: + self.nested_set(x, self.nested_get(x, self.source), self.target) + x[self.get_name(x)] = torch.zeros(8) + return x + params = self.make_params([]) + self.nested_set( + x, self.transform(self.nested_get(x, self.source), params), self.target + ) + brightness_factor = params["brightness_factor"] + contrast_factor = params["contrast_factor"] + saturation_factor = params["saturation_factor"] + hue_factor = params["hue_factor"] + perm = params["fn_idx"].tolist() + x[self.get_name(x)] = torch.Tensor( + [brightness_factor, contrast_factor, saturation_factor, hue_factor] + perm + ) return x -class RandomGrayscale(Transform, v2.RandomGrayscale): - """Randomly convert image to grayscale with given probability.""" +class RandomRotation(Transform, v2.RandomRotation): + """Rotate image by random angle within specified degrees range.""" - def __init__(self, p=0.1, source: str = "image", target: str = "image"): - super().__init__(p) + def __init__( + self, + degrees, + interpolation=InterpolationMode.NEAREST, + expand=False, + center=None, + fill=0, + source: str = "image", + target: str = "image", + ): + super().__init__(degrees, interpolation, expand, center, fill) self.source = source self.target = target - def _get_params(self, inp: List[Any]) -> Dict[str, Any]: - num_input_channels, *_ = query_chw([inp]) - return dict(num_input_channels=num_input_channels) - - def __call__(self, x) -> Any: - if self.p < 1 and torch.rand(1) >= self.p: - x[self.get_name(x)] = False - self.nested_set(x, self.nested_get(x, self.source), self.target) - return x - channels, *_ = query_chw([self.nested_get(x, self.source)]) + def __call__(self, x): + angle = self.make_params([]) self.nested_set( - x, - F.rgb_to_grayscale( - self.nested_get(x, self.source), num_output_channels=channels - ), - self.target, + x, self.transform(self.nested_get(x, self.source), angle), self.target ) - x[self.get_name(x)] = True + x[self.get_name(x)] = angle return x -class Lambda(Transform): - """Applies a lambda callable to target key and store it in source.""" +class RandomChannelPermutation(Transform, v2.RandomChannelPermutation): + """Randomly permute the channels of an image.""" - def __init__(self, lambd, source: str = "image", target: str = "image"): + def __init__(self, source: str = "image", target: str = "image"): super().__init__() self.source = source self.target = target - self.lambd = lambd def __call__(self, x) -> Any: - self.nested_set(x, self.lambd(x), self.target) + num_channels, *_ = query_chw([self.nested_get(x, self.source)]) + perm = torch.randperm(num_channels) + self.nested_set( + x, F.permute_channels(self.nested_get(x, self.source), perm), self.target + ) + x[self.get_name(x)] = perm return x -class RoutingTransform(Transform): - """Applies a routing callable to conditionally apply a transform from many candidates.""" - - def __init__(self, router: callable, transforms: Union[list, tuple, dict]): - self.router = router - self.transforms = transforms - - def __call__(self, x) -> Any: - route = self.router(x) - return self.transforms[route](x) - +class RandomCrop(Transform, v2.RandomCrop): + """Crop a random portion of image and resize it to given size.""" -class WrapTorchTransform(Transform, v2.Lambda): - """Applies a lambda callable to target key and store it in source.""" + _NAMES = ["needs_crop", "top", "left", "height", "width", "needs_pad", "padding"] - def __init__(self, transform, source: str = "image", target: str = "image"): - super().__init__(transform) + def __init__( + self, + size, + padding=None, + pad_if_needed=False, + fill=0, + padding_mode="constant", + source: str = "image", + target: str = "image", + ): + super().__init__(size, padding, pad_if_needed, fill, padding_mode) self.source = source self.target = target - def __call__(self, x) -> Any: + def __call__(self, x): + params = self.make_params([self.nested_get(x, self.source)]) self.nested_set( - x, super().__call__(self.nested_get(x, self.source)), self.target + x, self.transform(self.nested_get(x, self.source), params), self.target ) - return x - - -class RandomSolarize(Transform, v2.RandomSolarize): - """Randomly solarize image by inverting pixel values above threshold.""" + values = [] + values.append(params["needs_crop"]) + values.append(params["top"]) + values.append(params["left"]) + values.append(params["height"]) + values.append(params["width"]) + values.append(params["needs_pad"]) + values.extend(params["padding"]) + x[self.get_name(x)] = torch.Tensor(values) + return x - def __init__(self, threshold, p=0.5, source: str = "image", target: str = "image"): - super().__init__(threshold, p) + +class RandomHorizontalFlip(Transform, v2.RandomHorizontalFlip): + """Horizontally flip the given image randomly with a given probability.""" + + def __init__(self, p=0.5, source: str = "image", target: str = "image"): + super().__init__(p) self.source = source self.target = target def __call__(self, x) -> Any: - if self.p < 1 and torch.rand(1) >= self.p: + if self.p > 0 and torch.rand(1) < self.p: + self.nested_set( + x, F.horizontal_flip(self.nested_get(x, self.source)), self.target + ) + x[self.get_name(x)] = True + else: + self.nested_set(x, self.nested_get(x, self.source), self.target) x[self.get_name(x)] = False - return x - self.nested_set( - x, F.solarize(self.nested_get(x, self.source), self.threshold), self.target - ) - x[self.get_name(x)] = True return x -class GaussianBlur(Transform, v2.GaussianBlur): - """Apply Gaussian blur to image with random sigma values.""" +class RandomResizedCrop(Transform, v2.RandomResizedCrop): + """Crop a random portion of image and resize it to given size.""" - _NAMES = ["sigma_x", "sigma_y"] + _NAMES = ["top", "left", "height", "width"] def __init__( self, - kernel_size, - sigma=(0.1, 2.0), - p=1, + size: Union[int, Sequence[int]], + scale: Tuple[float, float] = (0.08, 1.0), + ratio: Tuple[float, float] = (3.0 / 4.0, 4.0 / 3.0), + interpolation: Union[InterpolationMode, int] = InterpolationMode.BILINEAR, + antialias: Optional[bool] = True, source: str = "image", target: str = "image", ): - super().__init__(kernel_size, sigma) - self.p = p + super().__init__(size, scale, ratio, interpolation, antialias) self.source = source self.target = target - def __call__(self, x) -> Any: - if self.p < 1 and torch.rand(1) >= self.p: - x[self.get_name(x)] = torch.zeros((2,)) - return x - params = self.make_params([]) + def __call__(self, x): + params = self.make_params([self.nested_get(x, self.source)]) self.nested_set( x, self.transform(self.nested_get(x, self.source), params), self.target ) - x[self.get_name(x)] = torch.Tensor(params["sigma"]) + values = [] + values.append(params["top"]) + values.append(params["left"]) + values.append(params["height"]) + values.append(params["width"]) + x[self.get_name(x)] = torch.Tensor(values) return x -class PILGaussianBlur(Transform): - """PIL-based Gaussian blur transform with random sigma sampling.""" - - _NAMES = ["sigma_x", "sigma_y"] - - def __init__(self, sigma=None, p=1, source: str = "image", target: str = "image"): - """Gaussian blur as a callable object. +class CenterCrop(Transform, v2.CenterCrop): + """Crop the center of an image to the given size.""" - Args: - sigma (Sequence[float]): range to sample the radius of the gaussian blur filter. - Defaults to [0.1, 2.0]. - p (float): probability of applying the transform. - source (str): source key in the data dictionary. - target (str): target key in the data dictionary. - """ - if sigma is None: - sigma = [0.1, 2.0] + _NAMES = [] - self.sigma = sigma - self.p = p + def __init__(self, size, source: str = "image", target: str = "image"): + super().__init__(size) self.source = source self.target = target def __call__(self, x): - """Applies gaussian blur to an input image. - - Args: - x (dict): Data dictionary containing the image to transform. - - Returns: - dict: Data dictionary with blurred image. - """ - if self.p < 1 and torch.rand(1) >= self.p: - x[self.get_name(x)] = torch.zeros((1,)) - return x - sigma = torch.rand((1,)) * (self.sigma[1] - self.sigma[0]) + self.sigma[0] - x[self.get_name(x)] = sigma self.nested_set( - x, - self.nested_get(x, self.source).filter( - ImageFilter.GaussianBlur(radius=sigma.item()) - ), - self.target, + x, self.transform(self.nested_get(x, self.source), []), self.target ) return x -class UniformTemporalSubsample(Transform): - """``nn.Module`` wrapper for ``pytorchvideo.transforms.functional.uniform_temporal_subsample``.""" +def set_seed(seeds): + if hasattr(seeds[0], "__len__"): + version, state, gauss = seeds[0] + setstate((version, tuple(state), gauss)) + else: + rseed(seeds[0]) + if hasattr(seeds[1], "__len__"): + np.random.set_state(seeds[1]) + else: + np.random.seed(seeds[1]) + if hasattr(seeds[2], "__len__"): + torch.set_rng_state(seeds[2]) + else: + torch.manual_seed(seeds[2]) + if len(seeds) == 4: + if hasattr(seeds[3], "__len__"): + torch.cuda.set_rng_state_all(seeds[3]) + else: + torch.cuda.manual_seed(seeds[3]) + + +@contextmanager +def random_seed(seed): + seeds = [getstate(), np.random.get_state(), torch.get_rng_state()] + if False: # torch.cuda.is_available(): + seeds.append(torch.cuda.get_rng_state_all()) + new_seeds = [int(seed)] * len(seeds) + set_seed(new_seeds) + yield + set_seed(seeds) + + +class ControlledTransform(Transform): + """Face Landmarks dataset.""" def __init__( - self, - num_samples: int, - temporal_dim: int = -3, - source: str = "video", - target: str = "video", + self, transform: callable, seed_offset: int = 0, key: Optional[str] = "idx" ): - super().__init__(num_samples, temporal_dim) - self.source = source - self.target = target + super().__init__() + self.seed_offset = seed_offset + self._transform = transform + self.key = key - def forward(self, x: dict) -> torch.Tensor: - self.nested_set( - x, super().forward(self, self.nested_get(x, self.source)), self.target - ) + def __call__(self, x): + with random_seed(x["idx"] + self.seed_offset): + x = self._transform(x) return x -class RandomContiguousTemporalSampler(Transform): - """Randomly sample contiguous frames from a video sequence.""" +class Conditional(Transform): + """Apply transform conditionally based on a data dictionary key.""" - def __init__(self, source, target, num_frames, frame_subsampling: int = 1): - self.source = source - self.target = target - self.num_frames = num_frames - self.frame_subsampling = frame_subsampling + def __init__(self, transform, condition_key, apply_on_true=True): + super().__init__() + self._transform = transform + self.condition_key = condition_key + self.apply_on_true = apply_on_true def __call__(self, x): - metadata = self.nested_get(x, self.source).get_metadata() - T = int(metadata["video"]["duration"][0] * metadata["video"]["fps"][0]) - covering = self.num_frames * self.frame_subsampling - start = torch.randint(low=0, high=T - covering, size=(1,)).item() - video_frames = [] # video frame buffer - - # Seek and return frames - count = 0 - for frame in islice( - self.nested_get(x, self.source).seek(start / metadata["video"]["fps"][0]), - covering, - ): - if count % self.frame_subsampling == 0: - video_frames.append(frame["data"]) - count += 1 - # Stack it into a tensor - self.nested_set(x, torch.stack(video_frames, 0), self.target) - x[self.get_name(x)] = start + if x[self.condition_key] and self.apply_on_true: + return self._transform(x) + elif not x[self.condition_key] and not self.apply_on_true: + return self._transform(x) + # if the transform is not applied we still inform the user + # otherwise collate_fn will complain + x[self._transform.get_name(x)] = self._transform.BYPASS_VALUE return x -class RGB(Transform, v2.RGB): - """Convert image to RGB format.""" +class AdditiveGaussian(Transform): + """Add Gaussian noise to input data.""" - def __init__(self, source: str = "image", target: str = "image"): + BYPASS_VALUE = False + + def __init__(self, sigma, p=1): super().__init__() - self.source = source - self.target = target + if not torch.is_tensor(sigma): + sigma = torch.Tensor([sigma])[0] + self.sigma = sigma + self.p = p def __call__(self, x): - self.nested_set( - x, F.grayscale_to_rgb(self.nested_get(x, self.source)), self.target - ) + if self.p == 0 or self.p < torch.rand(1): + x[self.get_name(x)] = self.BYPASS_VALUE + return x + x[self.get_name(x)] = True + out = torch.randn_like(x["image"]).mul_(self.sigma) + x["image"] = x["image"].add_(out) return x -class Resize(Transform, v2.Resize): - """Resize image to specified size.""" +class Compose(v2.Transform): + """Compose multiple transforms together in sequence.""" - def __init__( - self, - size, - interpolation=2, - max_size=None, - antialias=True, - source="image", - target="image", - ) -> None: - super().__init__(size, interpolation, max_size, antialias) - self.source = source - self.target = target + def __init__(self, *args): + super().__init__() + self.args = args - def __call__(self, x): - self.nested_set( - x, self.transform(self.nested_get(x, self.source), []), self.target - ) - return x + def __call__(self, sample): + for a in self.args: + sample = a(sample) + return sample -class ColorJitter(Transform, v2.ColorJitter): - """Randomly change brightness, contrast, saturation, and hue of an image.""" +class RoundRobinMultiViewTransform(v2.Transform): + """Round-robin multi-view transform that cycles through transforms using a counter. - def __init__( - self, - brightness=None, - contrast=None, - saturation=None, - hue=None, - p=1, - source: str = "image", - target: str = "image", - ): - super().__init__(brightness, contrast, saturation, hue) - self.p = p - self.source = source - self.target = target + IMPORTANT: This transform is designed to work with RepeatedRandomSampler, where + each image index appears multiple times consecutively in the batch. It uses an + internal counter to apply different augmentations to each repeated occurrence. - def __call__(self, x) -> Any: - if self.p < 1 and torch.rand(1) > self.p: - self.nested_set(x, self.nested_get(x, self.source), self.target) - x[self.get_name(x)] = torch.zeros(8) - return x - params = self.make_params([]) - self.nested_set( - x, self.transform(self.nested_get(x, self.source), params), self.target - ) - brightness_factor = params["brightness_factor"] - contrast_factor = params["contrast_factor"] - saturation_factor = params["saturation_factor"] - hue_factor = params["hue_factor"] - perm = params["fn_idx"].tolist() - x[self.get_name(x)] = torch.Tensor( - [brightness_factor, contrast_factor, saturation_factor, hue_factor] + perm - ) - return x + BATCH SIZE NOTE: When using this with RepeatedRandomSampler, the batch_size + parameter refers to the total number of augmented samples, NOT the number of + unique images. For example, with batch_size=256 and n_views=2, you get 128 + unique images, each appearing twice with different augmentations. + How it works: + 1. RepeatedRandomSampler produces indices like [0,0,1,1,2,2,...] (for n_views=2) + 2. DataLoader loads the same image multiple times + 3. This transform applies a different augmentation each time using round-robin -class RandomRotation(Transform, v2.RandomRotation): - """Rotate image by random angle within specified degrees range.""" + Args: + transforms: List of transforms, one for each view. The counter cycles + through these transforms in order. - def __init__( - self, - degrees, - interpolation=InterpolationMode.NEAREST, - expand=False, - center=None, - fill=0, - source: str = "image", - target: str = "image", - ): - super().__init__(degrees, interpolation, expand, center, fill) - self.source = source - self.target = target + Example: + # With RepeatedRandomSampler(dataset, n_views=2) + transform = RoundRobinMultiViewTransform([ + strong_augmentation, # Applied to 1st occurrence of each image + weak_augmentation, # Applied to 2nd occurrence of each image + ]) - def __call__(self, x): - angle = self.make_params([]) - self.nested_set( - x, self.transform(self.nested_get(x, self.source), angle), self.target - ) - x[self.get_name(x)] = angle - return x + Warning: The internal counter makes this transform stateful and not thread-safe. + """ + def __init__(self, transforms): + super().__init__() + self.transforms = transforms + self.n_transforms = len(transforms) + self.counter = 0 -class RandomChannelPermutation(Transform, v2.RandomChannelPermutation): - """Randomly permute the channels of an image.""" + def __call__(self, sample): + # Use round-robin to apply transforms + transform_idx = self.counter % self.n_transforms + self.counter += 1 + return self.transforms[transform_idx](sample) - def __init__(self, source: str = "image", target: str = "image"): - super().__init__() - self.source = source - self.target = target - def __call__(self, x) -> Any: - num_channels, *_ = query_chw([self.nested_get(x, self.source)]) - perm = torch.randperm(num_channels) - self.nested_set( - x, F.permute_channels(self.nested_get(x, self.source), perm), self.target - ) - x[self.get_name(x)] = perm - return x +class MultiViewTransform(v2.Transform): + """Creates multiple views from one sample by applying different transforms. + Takes a single sample and applies different transforms to create multiple + views, returning a list of complete sample dicts. Preserves all modifications + each transform makes (masks, augmentation params, metadata, etc.). -class RandomCrop(Transform, v2.RandomCrop): - """Crop a random portion of image and resize it to given size.""" + Implementation Note: + This transform uses shallow copy (dict.copy()) for the input sample before + applying each transform. This is efficient and safe because: + - The shallow copy shares references to the original tensors/objects + - Standard transforms create NEW tensors (e.g., through mul(), resize(), + crop()) rather than modifying inputs in-place + - The original sample remains unchanged - _NAMES = ["needs_crop", "top", "left", "height", "width", "needs_pad", "padding"] + Consequences of shallow copy: + - Memory efficient: Original tensors are not duplicated unnecessarily + - Safe with torchvision transforms: All torchvision transforms and our + custom transforms follow the pattern of creating new tensors + - Caution: If using custom transforms that modify tensors in-place (using + operations like mul_(), add_() with underscore), views may interfere with + each other. Always use non-in-place operations in custom transforms. - def __init__( - self, - size, - padding=None, - pad_if_needed=False, - fill=0, - padding_mode="constant", - source: str = "image", - target: str = "image", - ): - super().__init__(size, padding, pad_if_needed, fill, padding_mode) - self.source = source - self.target = target + Args: + transforms: Either a list or dict of transforms. + - List: Returns a list of views in the same order + - Dict: Returns a dict of views with the same keys - def __call__(self, x): - params = self.make_params([self.nested_get(x, self.source)]) - self.nested_set( - x, self.transform(self.nested_get(x, self.source), params), self.target - ) - values = [] - values.append(params["needs_crop"]) - values.append(params["top"]) - values.append(params["left"]) - values.append(params["height"]) - values.append(params["width"]) - values.append(params["needs_pad"]) - values.extend(params["padding"]) - x[self.get_name(x)] = torch.Tensor(values) - return x + Returns: + Union[List[dict], Dict[str, dict]]: + - If transforms is a list: Returns a list of transformed sample dicts + - If transforms is a dict: Returns a dict of transformed sample dicts with same keys + Each dict contains NEW tensors, not references to the original. + Example: + # List input - returns list of views + transform = MultiViewTransform([ + strong_augmentation, # Creates first view with strong aug + weak_augmentation, # Creates second view with weak aug + ]) + # Input: {"image": img, "label": 0} + # Output: [{"image": img_strong, "label": 0}, {"image": img_weak, "label": 0}] -class RandomHorizontalFlip(Transform, v2.RandomHorizontalFlip): - """Horizontally flip the given image randomly with a given probability.""" + # Dict input - returns dict of named views + transform = MultiViewTransform({ + "student": strong_augmentation, + "teacher": weak_augmentation, + }) + # Input: {"image": img, "label": 0} + # Output: {"student": {"image": img_strong, "label": 0}, + # "teacher": {"image": img_weak, "label": 0}} + """ - def __init__(self, p=0.5, source: str = "image", target: str = "image"): - super().__init__(p) - self.source = source - self.target = target + def __init__(self, transforms): + super().__init__() + self.transforms = transforms + self.return_dict = isinstance(transforms, dict) - def __call__(self, x) -> Any: - if self.p > 0 and torch.rand(1) < self.p: - self.nested_set( - x, F.horizontal_flip(self.nested_get(x, self.source)), self.target - ) - x[self.get_name(x)] = True + def __call__(self, sample): + """Create multiple views by applying different transforms to the sample.""" + if self.return_dict: + # Dict input - return dict of views + views = {} + for key, transform in self.transforms.items(): + # Copy to avoid transforms modifying the original + sample_copy = sample.copy() + # Apply transform to entire dict + transformed = transform(sample_copy) + views[key] = transformed else: - self.nested_set(x, self.nested_get(x, self.source), self.target) - x[self.get_name(x)] = False - return x + # List input - return list of views + views = [] + for transform in self.transforms: + # Copy to avoid transforms modifying the original + sample_copy = sample.copy() + # Apply transform to entire dict + transformed = transform(sample_copy) + views.append(transformed) + return views -class RandomResizedCrop(Transform, v2.RandomResizedCrop): - """Crop a random portion of image and resize it to given size.""" - _NAMES = ["top", "left", "height", "width"] +class ContextTargetsMultiBlockMask(Transform): + """Transform that adds multi-block masks to batch, with multiple target blocks and one disjoint context block. + + Args: + patch_size: Size of the patch in patches + num_blocks: Number of blocks to sample + context_scale: Scale of the context block + aspect_ratio: Aspect ratio of the blocks + min_keep: Minimum number of patches that must be in the block + + """ def __init__( self, - size: Union[int, Sequence[int]], - scale: Tuple[float, float] = (0.08, 1.0), - ratio: Tuple[float, float] = (3.0 / 4.0, 4.0 / 3.0), - interpolation: Union[InterpolationMode, int] = InterpolationMode.BILINEAR, - antialias: Optional[bool] = True, + patch_size=16, + context_scale=(0.85, 1.0), + context_aspect_ratio=(1.0, 1.0), + target_scales=((0.15, 0.2),) * 4, + target_aspect_ratios=((0.75, 1.5),) * 4, + min_keep=10, source: str = "image", - target: str = "image", + target_context: str = "mask_context", + target_targets: str = "masks_target", ): - super().__init__(size, scale, ratio, interpolation, antialias) + super().__init__() + self.patch_size = patch_size + self.context_scale = context_scale + self.context_aspect_ratio = context_aspect_ratio + self.target_scales = target_scales + self.target_aspect_ratios = target_aspect_ratios self.source = source - self.target = target - - def __call__(self, x): - params = self.make_params([self.nested_get(x, self.source)]) - self.nested_set( - x, self.transform(self.nested_get(x, self.source), params), self.target - ) - values = [] - values.append(params["top"]) - values.append(params["left"]) - values.append(params["height"]) - values.append(params["width"]) - x[self.get_name(x)] = torch.Tensor(values) - return x - - -class CenterCrop(Transform, v2.CenterCrop): - """Crop the center of an image to the given size.""" - - _NAMES = [] + self.target_context = target_context + self.target_targets = target_targets + if len(target_scales) != len(target_aspect_ratios): + raise ValueError( + "Each scale must have its associated aspect ratio and vice versa.", + "Received {len(target_scales)=} {len(target_aspect_ratios)=}", + ) - def __init__(self, size, source: str = "image", target: str = "image"): - super().__init__(size) - self.source = source - self.target = target + self.min_keep = min_keep def __call__(self, x): - self.nested_set( - x, self.transform(self.nested_get(x, self.source), []), self.target + source = self.nested_get(x, self.source) + if isinstance(source, PIL.Image.Image): + W, H = source.size # PIL is W,H + elif isinstance(source, torch.Tensor): + # assumes H W + H, W = source.shape[-2:] + else: + raise ValueError( + f"Source must be a PIL.Image.Image or a torch.Tensor, but got {type(source)} instead." + ) + + scales = [self.context_scale, *self.target_scales] + aspect_ratios = [self.context_aspect_ratio, *self.target_aspect_ratios] + context_mask, *target_masks = multi_block_mask( + H // self.patch_size, + W // self.patch_size, + block_scales=scales, + aspect_ratios=aspect_ratios, + min_keep=self.min_keep, ) + # makes targets disjoint with context + for mask in target_masks: + context_mask &= ~mask + + x[self.target_context] = torch.nonzero(context_mask.flatten()).squeeze() + x[self.target_targets] = [ + torch.nonzero(mask.flatten()).squeeze() for mask in target_masks + ] + x[self.get_name(x)] = torch.tensor([scales, aspect_ratios]) return x -def set_seed(seeds): - if hasattr(seeds[0], "__len__"): - version, state, gauss = seeds[0] - setstate((version, tuple(state), gauss)) - else: - rseed(seeds[0]) - if hasattr(seeds[1], "__len__"): - np.random.set_state(seeds[1]) - else: - np.random.seed(seeds[1]) - if hasattr(seeds[2], "__len__"): - torch.set_rng_state(seeds[2]) - else: - torch.manual_seed(seeds[2]) - if len(seeds) == 4: - if hasattr(seeds[3], "__len__"): - torch.cuda.set_rng_state_all(seeds[3]) - else: - torch.cuda.manual_seed(seeds[3]) +class RandomMask(Transform): + r"""Creates a random MAE-style mask for an image. + This transform generates a random permutation of all patch indices for an + input image. It then splits these indices into two disjoint sets: + 'visible' and 'masked', according to the specified `mask_ratio`. -@contextmanager -def random_seed(seed): - seeds = [getstate(), np.random.get_state(), torch.get_rng_state()] - if False: # torch.cuda.is_available(): - seeds.append(torch.cuda.get_rng_state_all()) - new_seeds = [int(seed)] * len(seeds) - set_seed(new_seeds) - yield - set_seed(seeds) + It also provides an `ids_restore` tensor, which can un-shuffle a sequence + of patches back to its original 2D grid order. All outputs are added as + new keys to the sample dictionary. + Example: + >>> # xdoctest: +SKIP + >>> transform = RandomMask(patch_size=16, mask_ratio=0.75) + >>> sample = {"image": torch.randn(3, 224, 224)} + >>> result = transform(sample) + >>> sorted(result.keys()) + ['image', 'ids_restore', 'len_keep', 'mask_masked', 'mask_visible'] + >>> result["len_keep"] + 49 + >>> result["mask_visible"].shape + torch.Size([49]) -class ControlledTransform(Transform): - """Face Landmarks dataset.""" + Args: + patch_size (int): The height and width of each square patch. + mask_ratio (float): The fraction of patches to be masked (e.g., 0.75). + source (str): The key in the sample dict for the source image tensor. + target_visible (str): The key to use when storing visible patch indices. + target_masked (str): The key to use when storing masked patch indices. + target_ids_restore (str): The key to use for the restoration indices. + target_len_keep (str): The key to use for the count of visible patches. + """ def __init__( - self, transform: callable, seed_offset: int = 0, key: Optional[str] = "idx" + self, + patch_size=16, + mask_ratio=0.75, + source: str = "image", + target_visible: str = "mask_visible", + target_masked: str = "mask_masked", + target_ids_restore: str = "ids_restore", + target_len_keep: str = "len_keep", ): super().__init__() - self.seed_offset = seed_offset - self._transform = transform - self.key = key + self.patch_size = patch_size + self.mask_ratio = mask_ratio + self.source = source + self.target_visible = target_visible + self.target_masked = target_masked + self.target_ids_restore = target_ids_restore + self.target_len_keep = target_len_keep def __call__(self, x): - with random_seed(x["idx"] + self.seed_offset): - x = self._transform(x) - return x + source = self.nested_get(x, self.source) + if isinstance(source, PIL.Image.Image): + W, H = source.size # PIL is W,H + elif isinstance(source, torch.Tensor): + # NOTE assumes _HW + H, W = source.shape[-2:] + else: + raise ValueError( + f"Source must be a PIL.Image.Image or a torch.Tensor, but got {type(source)} instead." + ) + num_patches = (H // self.patch_size) * (W // self.patch_size) + len_keep = int(num_patches * (1 - self.mask_ratio)) -class Conditional(Transform): - """Apply transform conditionally based on a data dictionary key.""" + # Generate random noise and shuffle indices (like MAE) + noise = torch.rand(num_patches) + ids_shuffle = torch.argsort(noise) + ids_restore = torch.argsort(ids_shuffle) # inverse permutation - def __init__(self, transform, condition_key, apply_on_true=True): - super().__init__() - self._transform = transform - self.condition_key = condition_key - self.apply_on_true = apply_on_true + # Split into visible and masked + mask_visible = ids_shuffle[:len_keep] # first len_keep are visible + mask_masked = ids_shuffle[len_keep:] # rest are masked + + # Add to sample + x[self.target_visible] = mask_visible + x[self.target_masked] = mask_masked + x[self.target_ids_restore] = ( + ids_restore # NEW: for reconstructing full sequence + ) + x[self.target_len_keep] = len_keep - def __call__(self, x): - if x[self.condition_key] and self.apply_on_true: - return self._transform(x) - elif not x[self.condition_key] and not self.apply_on_true: - return self._transform(x) - # if the transform is not applied we still inform the user - # otherwise collate_fn will complain - x[self._transform.get_name(x)] = self._transform.BYPASS_VALUE return x -class AdditiveGaussian(Transform): - """Add Gaussian noise to input data.""" +# class RandomClassSwitch(v2.Transform): +# def __init__( +# self, +# label_key: str, +# new_key: str, +# p: float, +# low: int = -2147483648, +# high: int = 0, +# ): +# super().__init__() +# self.p = p +# self.label_key = label_key +# self.new_key = new_key +# self.low = low +# self.high = high - BYPASS_VALUE = False +# def __call__(self, sample: dict): +# assert type(sample) is dict +# assert self.label_key in sample +# assert self.new_key not in sample +# if self.p > 0 and torch.rand(1) < self.p: +# if torch.is_tensor(sample[self.label_key]): +# sample[self.new_key] = torch.randint( +# low=self.low, high=self.high, size=() +# ) +# else: +# sample[self.new_key] = np.random.randint(low=self.low, high=self.high) +# else: +# sample[self.new_key] = sample[self.label_key] +# return sample - def __init__(self, sigma, p=1): - super().__init__() - if not torch.is_tensor(sigma): - sigma = torch.Tensor([sigma])[0] - self.sigma = sigma - self.p = p - def __call__(self, x): - if self.p == 0 or self.p < torch.rand(1): - x[self.get_name(x)] = self.BYPASS_VALUE - return x - x[self.get_name(x)] = True - out = torch.randn_like(x["image"]).mul_(self.sigma) - x["image"] = x["image"].add_(out) - return x +# ============================================================ +# ================ Spurious Correlations ===================== +# ============================================================ -class Compose(v2.Transform): - """Compose multiple transforms together in sequence.""" +# ============================================================ +# ===================== Image MODIFIERS ====================== +# ============================================================ - def __init__(self, *args): - super().__init__() - self.args = args - def __call__(self, sample): - for a in self.args: - sample = a(sample) - return sample +class AddSampleIdx(Transform): + """Add an "idx" key each sample to allow for deterministic injection.""" + def __init__(self): + super().__init__() + self._counter = 0 -class RoundRobinMultiViewTransform(v2.Transform): - """Round-robin multi-view transform that cycles through transforms using a counter. + def __call__(self, x: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: + if "idx" not in x: + x["idx"] = self._counter + self._counter += 1 - IMPORTANT: This transform is designed to work with RepeatedRandomSampler, where - each image index appears multiple times consecutively in the batch. It uses an - internal counter to apply different augmentations to each repeated occurrence. + return x - BATCH SIZE NOTE: When using this with RepeatedRandomSampler, the batch_size - parameter refers to the total number of augmented samples, NOT the number of - unique images. For example, with batch_size=256 and n_views=2, you get 128 - unique images, each appearing twice with different augmentations. - How it works: - 1. RepeatedRandomSampler produces indices like [0,0,1,1,2,2,...] (for n_views=2) - 2. DataLoader loads the same image multiple times - 3. This transform applies a different augmentation each time using round-robin +class AddPatch(Transform): + """Add a solid color patch to an image at a fixed position. Args: - transforms: List of transforms, one for each view. The counter cycles - through these transforms in order. + patch_size (float): Fraction of image width/height for the patch (0 < patch_size ≤ 1). + color (Tuple[float, float, float]): RGB values in [0, 1]. + position (str): Where to place the patch: 'top_left_corner', 'top_right_corner', + 'bottom_left_corner', 'bottom_right_corner'. + """ - Example: - # With RepeatedRandomSampler(dataset, n_views=2) - transform = RoundRobinMultiViewTransform([ - strong_augmentation, # Applied to 1st occurrence of each image - weak_augmentation, # Applied to 2nd occurrence of each image - ]) - - Warning: The internal counter makes this transform stateful and not thread-safe. - """ - - def __init__(self, transforms): + def __init__( + self, + patch_size: float = 0.1, + color: Tuple[float, float, float] = (1.0, 0.0, 0.0), + position: str = "bottom_right_corner", + ): super().__init__() - self.transforms = transforms - self.n_transforms = len(transforms) - self.counter = 0 - - def __call__(self, sample): - # Use round-robin to apply transforms - transform_idx = self.counter % self.n_transforms - self.counter += 1 - return self.transforms[transform_idx](sample) - - -class MultiViewTransform(v2.Transform): - """Creates multiple views from one sample by applying different transforms. - - Takes a single sample and applies different transforms to create multiple - views, returning a list of complete sample dicts. Preserves all modifications - each transform makes (masks, augmentation params, metadata, etc.). - - Implementation Note: - This transform uses shallow copy (dict.copy()) for the input sample before - applying each transform. This is efficient and safe because: - - The shallow copy shares references to the original tensors/objects - - Standard transforms create NEW tensors (e.g., through mul(), resize(), - crop()) rather than modifying inputs in-place - - The original sample remains unchanged - - Consequences of shallow copy: - - Memory efficient: Original tensors are not duplicated unnecessarily - - Safe with torchvision transforms: All torchvision transforms and our - custom transforms follow the pattern of creating new tensors - - Caution: If using custom transforms that modify tensors in-place (using - operations like mul_(), add_() with underscore), views may interfere with - each other. Always use non-in-place operations in custom transforms. - - Args: - transforms: Either a list or dict of transforms. - - List: Returns a list of views in the same order - - Dict: Returns a dict of views with the same keys - - Returns: - Union[List[dict], Dict[str, dict]]: - - If transforms is a list: Returns a list of transformed sample dicts - - If transforms is a dict: Returns a dict of transformed sample dicts with same keys - Each dict contains NEW tensors, not references to the original. - Example: - # List input - returns list of views - transform = MultiViewTransform([ - strong_augmentation, # Creates first view with strong aug - weak_augmentation, # Creates second view with weak aug - ]) - # Input: {"image": img, "label": 0} - # Output: [{"image": img_strong, "label": 0}, {"image": img_weak, "label": 0}] + # checking constraints + if patch_size <= 0 or patch_size > 1: + raise ValueError("patch_size must be between 0 and 1.") - # Dict input - returns dict of named views - transform = MultiViewTransform({ - "student": strong_augmentation, - "teacher": weak_augmentation, - }) - # Input: {"image": img, "label": 0} - # Output: {"student": {"image": img_strong, "label": 0}, - # "teacher": {"image": img_weak, "label": 0}} - """ + if len(color) != 3: + raise ValueError( + "color must be a tuple of size 3 in the form \ + Tuple[float, float, float]) with each representing RGB values in [0, 1]" + ) - def __init__(self, transforms): - super().__init__() - self.transforms = transforms - self.return_dict = isinstance(transforms, dict) + for value in color: + if value > 1 or value < 0: + raise ValueError("Each color value must be in [0, 1]") - def __call__(self, sample): - """Create multiple views by applying different transforms to the sample.""" - if self.return_dict: - # Dict input - return dict of views - views = {} - for key, transform in self.transforms.items(): - # Copy to avoid transforms modifying the original - sample_copy = sample.copy() - # Apply transform to entire dict - transformed = transform(sample_copy) - views[key] = transformed + self.patch_size = patch_size + self.color = color + self.position = position + + def __call__(self, x: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: + img = self.nested_get(x, "image") + _, H, W = img.shape + + patch_h = int(H * self.patch_size) + patch_w = int(W * self.patch_size) + + # Create a colored patch + patch = torch.zeros((3, patch_h, patch_w), device=img.device) + patch[0] = self.color[0] + patch[1] = self.color[1] + patch[2] = self.color[2] + + img = img.clone() + if self.position == "top_left_corner": + img[:, :patch_h, :patch_w] = patch + elif self.position == "top_right_corner": + img[:, :patch_h, -patch_w:] = patch + elif self.position == "bottom_left_corner": + img[:, -patch_h:, :patch_w] = patch + elif self.position == "bottom_right_corner": + img[:, -patch_h:, -patch_w:] = patch + elif self.position == "center": + center_y, center_x = H // 2, W // 2 + img[ + :, + center_y - patch_h // 2 : center_y + patch_h // 2, + center_x - patch_w // 2 : center_x + patch_w // 2, + ] = patch else: - # List input - return list of views - views = [] - for transform in self.transforms: - # Copy to avoid transforms modifying the original - sample_copy = sample.copy() - # Apply transform to entire dict - transformed = transform(sample_copy) - views.append(transformed) + raise ValueError( + f"Invalid position: {self.position}, valid positions are: \ + top_left_corner, top_right_corner, bottom_left_corner, bottom_right_corner, center" + ) - return views + self.nested_set(x, img, "image") + return x -class ContextTargetsMultiBlockMask(Transform): - """Transform that adds multi-block masks to batch, with multiple target blocks and one disjoint context block. +class ClassConditionalInjector(Transform): + """Applies transformations conditionally based on sample label. Args: - patch_size: Size of the patch in patches - num_blocks: Number of blocks to sample - context_scale: Scale of the context block - aspect_ratio: Aspect ratio of the blocks - min_keep: Minimum number of patches that must be in the block - + transformation (Transform): Transform to apply to the image. + label_key (str): Key for label in the sample dict. + target_labels (Union[int, list[int]]): Which labels to modify. + proportion (float): Fraction of samples with matching labels to modify (0-1). + total_samples (int, optional): Dataset size (for deterministic mask). + seed (int): Seed for randomization to determine which samples transformation is applied to """ def __init__( self, - patch_size=16, - context_scale=(0.85, 1.0), - context_aspect_ratio=(1.0, 1.0), - target_scales=((0.15, 0.2),) * 4, - target_aspect_ratios=((0.75, 1.5),) * 4, - min_keep=10, - source: str = "image", - target_context: str = "mask_context", - target_targets: str = "masks_target", + transformation: Transform, + label_key: str = "label", + target_labels: Union[int, list[int]] = 0, + proportion: float = 0.5, + total_samples: Optional[int] = None, + seed: int = 42, ): super().__init__() - self.patch_size = patch_size - self.context_scale = context_scale - self.context_aspect_ratio = context_aspect_ratio - self.target_scales = target_scales - self.target_aspect_ratios = target_aspect_ratios - self.source = source - self.target_context = target_context - self.target_targets = target_targets - if len(target_scales) != len(target_aspect_ratios): - raise ValueError( - "Each scale must have its associated aspect ratio and vice versa.", - "Received {len(target_scales)=} {len(target_aspect_ratios)=}", + self.transformation = transformation + self.label_key = label_key + self.target_labels = ( + [target_labels] if isinstance(target_labels, int) else target_labels + ) + self.proportion = proportion + self.total_samples = total_samples + self.seed = seed + + # Precompute deterministic mask if dataset size known + if total_samples is not None: + num_to_transform = int(total_samples * proportion) + rng = torch.Generator().manual_seed(seed) + self.indices_to_transform = set( + torch.randperm(total_samples, generator=rng)[:num_to_transform].tolist() ) + else: + self.indices_to_transform = None - self.min_keep = min_keep + def __call__(self, x: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: + label = self.nested_get(x, self.label_key) - def __call__(self, x): - source = self.nested_get(x, self.source) - if isinstance(source, PIL.Image.Image): - W, H = source.size # PIL is W,H - elif isinstance(source, torch.Tensor): - # assumes H W - H, W = source.shape[-2:] - else: - raise ValueError( - f"Source must be a PIL.Image.Image or a torch.Tensor, but got {type(source)} instead." - ) + # Determine if we apply the transformation + should_transform = False + idx = self.nested_get(x, "idx") + if label in self.target_labels: + if self.indices_to_transform is not None: + should_transform = idx in self.indices_to_transform + else: + should_transform = random.random() < self.proportion - scales = [self.context_scale, *self.target_scales] - aspect_ratios = [self.context_aspect_ratio, *self.target_aspect_ratios] - context_mask, *target_masks = multi_block_mask( - H // self.patch_size, - W // self.patch_size, - block_scales=scales, - aspect_ratios=aspect_ratios, - min_keep=self.min_keep, - ) - # makes targets disjoint with context - for mask in target_masks: - context_mask &= ~mask + if should_transform: + x = self.transformation(x) - x[self.target_context] = torch.nonzero(context_mask.flatten()).squeeze() - x[self.target_targets] = [ - torch.nonzero(mask.flatten()).squeeze() for mask in target_masks - ] - x[self.get_name(x)] = torch.tensor([scales, aspect_ratios]) return x -class RandomMask(Transform): - r"""Creates a random MAE-style mask for an image. +# ============================================================ +# ===================== TEXT MODIFIERS ======================= +# ============================================================ - This transform generates a random permutation of all patch indices for an - input image. It then splits these indices into two disjoint sets: - 'visible' and 'masked', according to the specified `mask_ratio`. - It also provides an `ids_restore` tensor, which can un-shuffle a sequence - of patches back to its original 2D grid order. All outputs are added as - new keys to the sample dictionary. +class Modifier: + """Base class for applying modifications/corruptions to text-label pairs. - Example: - >>> # xdoctest: +SKIP - >>> transform = RandomMask(patch_size=16, mask_ratio=0.75) - >>> sample = {"image": torch.randn(3, 224, 224)} - >>> result = transform(sample) - >>> sorted(result.keys()) - ['image', 'ids_restore', 'len_keep', 'mask_masked', 'mask_visible'] - >>> result["len_keep"] - 49 - >>> result["mask_visible"].shape - torch.Size([49]) + Subclasses must implement the __call__ method to define specific transformations. - Args: - patch_size (int): The height and width of each square patch. - mask_ratio (float): The fraction of patches to be masked (e.g., 0.75). - source (str): The key in the sample dict for the source image tensor. - target_visible (str): The key to use when storing visible patch indices. - target_masked (str): The key to use when storing masked patch indices. - target_ids_restore (str): The key to use for the restoration indices. - target_len_keep (str): The key to use for the count of visible patches. + Example: + class MyModifier(Modifier): + def __call__(self, text: str, label: Any) -> tuple[str, Any]: + # custom transformation here + return transformed_text, transformed_label """ - def __init__( - self, - patch_size=16, - mask_ratio=0.75, - source: str = "image", - target_visible: str = "mask_visible", - target_masked: str = "mask_masked", - target_ids_restore: str = "ids_restore", - target_len_keep: str = "len_keep", - ): - super().__init__() - self.patch_size = patch_size - self.mask_ratio = mask_ratio - self.source = source - self.target_visible = target_visible - self.target_masked = target_masked - self.target_ids_restore = target_ids_restore - self.target_len_keep = target_len_keep + def __call__(self, text: str, label): + """Apply the transformation to a single text-label pair. - def __call__(self, x): - source = self.nested_get(x, self.source) - if isinstance(source, PIL.Image.Image): - W, H = source.size # PIL is W,H - elif isinstance(source, torch.Tensor): - # NOTE assumes _HW - H, W = source.shape[-2:] - else: - raise ValueError( - f"Source must be a PIL.Image.Image or a torch.Tensor, but got {type(source)} instead." - ) + Args: + text (str): The input text to transform. + label: The associated label. - num_patches = (H // self.patch_size) * (W // self.patch_size) - len_keep = int(num_patches * (1 - self.mask_ratio)) + Returns: + tuple: (transformed_text, transformed_label) + """ + raise NotImplementedError("Subclasses must implement __call__") - # Generate random noise and shuffle indices (like MAE) - noise = torch.rand(num_patches) - ids_shuffle = torch.argsort(noise) - ids_restore = torch.argsort(ids_shuffle) # inverse permutation - # Split into visible and masked - mask_visible = ids_shuffle[:len_keep] # first len_keep are visible - mask_masked = ids_shuffle[len_keep:] # rest are masked +class CompositeModifier: + """CompositeModifier chains multiple Modifier instances together. - # Add to sample - x[self.target_visible] = mask_visible - x[self.target_masked] = mask_masked - x[self.target_ids_restore] = ( - ids_restore # NEW: for reconstructing full sequence + Each modifier from the list is applied sequentially to the text. This enables + the combination of various transformations or injections into one composite operation. + """ + + def __init__(self, modifiers: list): + """Initialize a CompositeModifier instance. + + Args: + modifiers (list): A list of modifier instances (subclasses of Modifier) + to be applied sequentially. + """ + self.modifiers = modifiers + + def __call__(self, text: str, label): + """Apply all modifiers in sequence to the given (text, label). + + Args: + text (str): The input text. + label: The associated label. + + Returns: + tuple: The modified (text, label) pair after all transformations. + """ + for modifier in self.modifiers: + text, label = modifier(text, label) + return text, label + + +class ItemInjection(Modifier): + """A Modifier that injects items into text. + + This class supports creation via three different approaches: + - from_list: Using a predefined list of injection items. + - from_file: Reading injection items from a file. + - from_function: Using a custom function to generate injections. + """ + + def __init__( + self, + injection_source, + location: str = "random", + token_proportion: float = 0.1, + seed=None, + _rng=None, + ): + """Initialize an ItemInjection instance. + + Args: + injection_source (callable): A function that returns an injection token. + location (str): Where to inject the token ("beginning", "random", "end"). + token_proportion (float): Proportion of tokens in the text to be affected. + seed (int, optional): Seed for reproducibility. + """ + assert callable(injection_source), "injection_source must be callable" + self.injection_source = injection_source + self.location = location + self.token_proportion = token_proportion + self.rng = _rng or random.Random(seed) + + assert 0 <= token_proportion <= 1, "token_proportion must be between 0 and 1" + assert location in {"beginning", "random", "end"}, ( + "location must be 'beginning', 'random', or 'end'" ) - x[self.target_len_keep] = len_keep - return x + def __call__(self, text: str, label): + """Inject tokens into the text at specified locations. + Args: + text (str): The input text to modify. + label: The original label (unchanged). -# class RandomClassSwitch(v2.Transform): -# def __init__( -# self, -# label_key: str, -# new_key: str, -# p: float, -# low: int = -2147483648, -# high: int = 0, -# ): -# super().__init__() -# self.p = p -# self.label_key = label_key -# self.new_key = new_key -# self.low = low -# self.high = high + Returns: + tuple: The modified text and the original label. + """ + words = text.split() + num_tokens = len(words) -# def __call__(self, sample: dict): -# assert type(sample) is dict -# assert self.label_key in sample -# assert self.new_key not in sample -# if self.p > 0 and torch.rand(1) < self.p: -# if torch.is_tensor(sample[self.label_key]): -# sample[self.new_key] = torch.randint( -# low=self.low, high=self.high, size=() -# ) -# else: -# sample[self.new_key] = np.random.randint(low=self.low, high=self.high) -# else: -# sample[self.new_key] = sample[self.label_key] -# return sample + # Ensure at least one token is injected + num_to_inject = max(1, int(num_tokens * self.token_proportion)) + + injections = [self.injection_source() for _ in range(num_to_inject)] + + if self.location == "beginning": + words = injections + words + elif self.location == "end": + words = words + injections + elif self.location == "random": + for injection in injections: + pos = self.rng.randint(0, len(words)) + words.insert(pos, injection) + + return " ".join(words), label # return modified text and unchanged label + + @classmethod + def from_list( + cls, + items: list, + location: str = "random", + token_proportion: float = 0.1, + seed=None, + ): + """Create an ItemInjection instance using a predefined list of tokens. + + Args: + items (list): List of token strings to choose from. + location (str): Where to inject tokens ("beginning", "random", "end"). + token_proportion (float): Proportion of text tokens to be affected. + seed (int, optional): Seed for reproducibility. + + Returns: + ItemInjection: Configured instance. + """ + rng = random.Random(seed) + + def injection_source(): + return rng.choice(items) + + return cls( + injection_source, + location=location, + token_proportion=token_proportion, + seed=seed, + _rng=rng, + ) + + @classmethod + def from_file( + cls, + file_path: str, + location: str = "random", + token_proportion: float = 0.1, + seed=None, + ): + """Create an ItemInjection instance using tokens read from a file. + + Each non-empty line becomes a potential injection item. + + Args: + file_path (str): Path to the file with one token per line. + location (str): Where to inject tokens. + token_proportion (float): Proportion of tokens to inject. + seed (int, optional): Seed for reproducibility. + + Returns: + ItemInjection: Configured instance. + """ + with open(file_path, "r", encoding="utf-8") as file: + items = [line.strip() for line in file if line.strip()] + + rng = random.Random(seed) + + def injection_source(): + return rng.choice(items) + + return cls( + injection_source, + location=location, + token_proportion=token_proportion, + _rng=rng, + ) + + @classmethod + def from_function( + cls, + injection_func, + location: str = "random", + token_proportion: float = 0.1, + seed=None, + ): + """Create an ItemInjection instance using a custom function to generate injections. + + Args: + injection_func (callable): Function that returns a new injection token each time. + location (str): Where to inject tokens. + token_proportion (float): Proportion of text to inject into. + seed (int, optional): Seed for reproducibility (used only for insertion position). + + Returns: + ItemInjection: Configured instance. + """ + assert callable(injection_func), "injection_func must be callable" + return cls( + injection_func, + location=location, + token_proportion=token_proportion, + seed=seed, + ) + + +class HTMLInjection(Modifier): + """A Modifier that injects html into text. + + This class supports creation via two different approaches: + - from_list: Using a predefined list of injection items. + - from_file: Reading injection items from a file. + """ + + def __init__( + self, + file_path: str, + location: str = "random", + level: int = None, + token_proportion: float = None, + seed=None, + ): + with open(file_path, "r", encoding="utf-8") as f: + self.tags = [line.strip() for line in f if line.strip()] + self.location = location + self.level = level + self.token_proportion = token_proportion + self.rng = random.Random(seed) + + if token_proportion is not None: + assert 0 < token_proportion <= 1, "token_proportion must be between 0 and 1" + + @classmethod + def from_file( + cls, + file_path: str, + location: str = "random", + level: int = None, + token_proportion: float = None, + seed=None, + ): + return cls( + file_path, + location=location, + level=level, + token_proportion=token_proportion, + seed=seed, + ) + + @classmethod + def from_list( + cls, + tags: list, + location: str = "random", + level: int = None, + token_proportion: float = None, + seed=None, + ): + instance = cls.__new__(cls) + instance.tags = tags + instance.location = location + instance.level = level + instance.token_proportion = token_proportion + instance.rng = random.Random(seed) + + if token_proportion is not None: + assert 0 < token_proportion <= 1, "token_proportion must be between 0 and 1" + + return instance + + def _choose_tag(self): + """Randomly choose a tag from the loaded list. + + Returns: + tuple: (opening_tag, closing_tag or None) + """ + line = self.rng.choice(self.tags) + parts = line.split() + if len(parts) >= 2: + return parts[0], parts[1] + else: + return parts[0], None + + def _inject_into_tokens(self, tokens, location): + tokens = tokens[:] + n = len(tokens) + + if self.token_proportion is None: + opening, closing = self._choose_tag() + return self._inject_with_tags(tokens, opening, closing, location) + + # Otherwise, inject up to token_proportion of total tokens + num_insertions = max(1, int(n * self.token_proportion)) + for _ in range(num_insertions): + opening, closing = self._choose_tag() + tokens = self._inject_with_tags(tokens, opening, closing, location) + return tokens + + def _inject_with_tags(self, tokens, opening, closing, location): + if location == "beginning": + new_tokens = [opening] + tokens + if closing: + pos = self.rng.randint(1, len(new_tokens)) + new_tokens.insert(pos, closing) + return new_tokens + + elif location == "end": + new_tokens = tokens[:] + pos = self.rng.randint(0, len(new_tokens)) + new_tokens.insert(pos, opening) + if closing: + new_tokens.append(closing) + return new_tokens + + elif location == "random": + new_tokens = tokens[:] + pos_open = self.rng.randint(0, len(new_tokens)) + new_tokens.insert(pos_open, opening) + if closing: + pos_close = self.rng.randint(pos_open + 1, len(new_tokens)) + new_tokens.insert(pos_close, closing) + return new_tokens + + return tokens + + def _inject(self, text, location): + tokens = text.split() + new_tokens = self._inject_into_tokens(tokens, location) + return " ".join(new_tokens) + + def _find_level_span(self, text, level): + """Find the first span inside the desired HTML nesting level. + + Args: + text (str): Input HTML text. + level (int): Desired nesting level. + + Returns: + tuple or None: (start, end) of the content region, or None if not found. + """ + tag_regex = re.compile(r"]*>") + stack = [] + for match in tag_regex.finditer(text): + tag_str = match.group(0) + tag_name = match.group(1) + if not tag_str.startswith(" Date: Thu, 16 Oct 2025 22:44:07 -0400 Subject: [PATCH 11/12] expanding the visual injections available --- stable_pretraining/data/transforms.py | 109 +++++++++++++++++- .../tests/unit/test_transforms.py | 19 +++ 2 files changed, 127 insertions(+), 1 deletion(-) diff --git a/stable_pretraining/data/transforms.py b/stable_pretraining/data/transforms.py index 22c7cddd1..ef8c44b20 100644 --- a/stable_pretraining/data/transforms.py +++ b/stable_pretraining/data/transforms.py @@ -15,6 +15,8 @@ from torchvision.transforms.functional import InterpolationMode from torchvision.transforms.v2 import functional as F from torchvision.transforms.v2._utils import query_chw +from torchvision.io import read_image +from torchvision.transforms.functional import resize from stable_pretraining.data.masking import multi_block_mask @@ -1003,7 +1005,7 @@ class AddPatch(Transform): patch_size (float): Fraction of image width/height for the patch (0 < patch_size ≤ 1). color (Tuple[float, float, float]): RGB values in [0, 1]. position (str): Where to place the patch: 'top_left_corner', 'top_right_corner', - 'bottom_left_corner', 'bottom_right_corner'. + 'bottom_left_corner', 'bottom_right_corner', 'center'. """ def __init__( @@ -1071,6 +1073,111 @@ def __call__(self, x: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: return x +class AddColorTint(Transform): + """Adds a color tint to the overall image (additive tint). + + Args: + tint (Tuple[float, float, float]): RGB representation of the tint that will be applied to the overall image + alpha (Float): mixing ratio for how much to blend the new color with the existing image + """ + + def __init__( + self, tint: Tuple[float, float, float] = (1.0, 0.8, 0.8), alpha: float = 0.3 + ): + super().__init__() + self.tint = torch.tensor(tint).view(3, 1, 1) + self.alpha = alpha + + def __call__(self, x): + img = self.nested_get(x, "image") + img = torch.clamp(img * (1 - self.alpha) + self.tint * self.alpha, 0, 1) + self.nested_set(x, img, "image") + return x + + +class AddBorder(Transform): + """Adds a border around an image. + + Args: + thickness (Float): how thick the border around the image will be + color (Tuple[float, float, float]): RGB representation of the color of the border + """ + + def __init__( + self, thickness: float = 0.05, color: Tuple[float, float, float] = (0, 1, 0) + ): + super().__init__() + self.thickness = thickness + self.color = color + + def __call__(self, x): + img = self.nested_get(x, "image").clone() + _, H, W = img.shape + + # scale to match image size + t = int(min(H, W) * self.thickness) + color_tensor = torch.tensor(self.color, device=img.device).view(3, 1, 1) + + img[:, :t, :] = color_tensor + img[:, -t:, :] = color_tensor + img[:, :, :t] = color_tensor + img[:, :, -t:] = color_tensor + self.nested_set(x, img, "image") + + return x + + +class AddWatermark(Transform): + """Overlay another image (logo, emoji, etc.) onto the base image. + + Args: + watermark_path (str): Path to the watermark image (e.g. 'smile.png'). + size (float): Fraction of base image size to scale watermark. + position (str): One of ['top_left', 'top_right', 'bottom_left', 'bottom_right', 'center']. + alpha (float): Opacity of watermark (0-1). + """ + + def __init__(self, watermark_path, size=0.2, position="bottom_right", alpha=0.8): + super().__init__() + # [C,H,W] tensor in [0,1] + self.watermark = read_image(watermark_path).float() / 255.0 + self.size = size + self.position = position + self.alpha = alpha + + def __call__(self, x): + img = self.nested_get(x, "image").clone() + _, H, W = img.shape + + # Resize watermark + w_h, w_w = self.watermark.shape[1:] + target_h = int(H * self.size) + target_w = int(w_w / w_h * target_h) + wm = resize(self.watermark, [target_h, target_w]) + + # Compute position + if self.position == "top_left": + y0, x0 = 0, 0 + elif self.position == "top_right": + y0, x0 = 0, W - target_w + elif self.position == "bottom_left": + y0, x0 = H - target_h, 0 + elif self.position == "bottom_right": + y0, x0 = H - target_h, W - target_w + elif self.position == "center": + y0, x0 = (H - target_h) // 2, (W - target_w) // 2 + else: + raise ValueError(f"Unknown position: {self.position}") + + background_region = img[:, y0 : y0 + target_h, x0 : x0 + target_w] + img[:, y0 : y0 + target_h, x0 : x0 + target_w] = ( + background_region * (1 - self.alpha) + wm * self.alpha + ) + + self.nested_set(x, img, "image") + return x + + class ClassConditionalInjector(Transform): """Applies transformations conditionally based on sample label. diff --git a/stable_pretraining/tests/unit/test_transforms.py b/stable_pretraining/tests/unit/test_transforms.py index 90d277f51..f968ffed1 100644 --- a/stable_pretraining/tests/unit/test_transforms.py +++ b/stable_pretraining/tests/unit/test_transforms.py @@ -75,6 +75,18 @@ def test_normalize_transform(self): # Check that normalization was applied assert not torch.allclose(result["image"], image) + def test_add_sample_idx_transform(self): + """Test adding the sample idex for the injections.""" + pass + + def test_add_patch_transform(self): + """Test adding the spurious patch into images.""" + pass + + def test_class_conditional_injector(self): + """Test conditional injections.""" + pass + def test_transform_params_initialization(self): """Test that transforms can be initialized with various parameters.""" # Test each transform can be created @@ -87,6 +99,13 @@ def test_transform_params_initialization(self): transforms.RandomResizedCrop(size=(32, 32)), transforms.RandomSolarize(threshold=0.5, p=0.2), transforms.RandomRotation(degrees=90), + transforms.AddSampleIdx(), + transforms.ClassConditionalInjector( + transformation=transforms.AddPatch( + patch_size=0.1, color=(1.0, 0.0, 0.0), position="center" + ), + total_samples=10000, + ), ] for t in transforms_to_test: From 77a9e77047c62df35d2d7b2209b0b83f9047c749 Mon Sep 17 00:00:00 2001 From: "marcel_mateos_salles@brown.edu" Date: Thu, 16 Oct 2025 23:00:18 -0400 Subject: [PATCH 12/12] updated tests --- .../tests/unit/test_transforms.py | 105 ++++++++++++++++-- 1 file changed, 93 insertions(+), 12 deletions(-) diff --git a/stable_pretraining/tests/unit/test_transforms.py b/stable_pretraining/tests/unit/test_transforms.py index f968ffed1..1fa7ca5c8 100644 --- a/stable_pretraining/tests/unit/test_transforms.py +++ b/stable_pretraining/tests/unit/test_transforms.py @@ -75,18 +75,6 @@ def test_normalize_transform(self): # Check that normalization was applied assert not torch.allclose(result["image"], image) - def test_add_sample_idx_transform(self): - """Test adding the sample idex for the injections.""" - pass - - def test_add_patch_transform(self): - """Test adding the spurious patch into images.""" - pass - - def test_class_conditional_injector(self): - """Test conditional injections.""" - pass - def test_transform_params_initialization(self): """Test that transforms can be initialized with various parameters.""" # Test each transform can be created @@ -100,6 +88,8 @@ def test_transform_params_initialization(self): transforms.RandomSolarize(threshold=0.5, p=0.2), transforms.RandomRotation(degrees=90), transforms.AddSampleIdx(), + transforms.AddColorTint(), + transforms.AddBorder(), transforms.ClassConditionalInjector( transformation=transforms.AddPatch( patch_size=0.1, color=(1.0, 0.0, 0.0), position="center" @@ -110,3 +100,94 @@ def test_transform_params_initialization(self): for t in transforms_to_test: assert t is not None + + # --------------------------- + # Spurious correlation tests + # --------------------------- + + def test_add_sample_idx_transform(self): + """Test that AddSampleIdx correctly increments indices.""" + transform = transforms.AddSampleIdx() + x1 = {"image": torch.zeros(3, 32, 32)} + x2 = {"image": torch.zeros(3, 32, 32)} + out1 = transform(x1) + out2 = transform(x2) + assert out1["idx"] == 0 + assert out2["idx"] == 1 + + def test_add_patch_transform(self): + """Test that AddPatch overlays a colored patch.""" + img = torch.zeros(3, 32, 32) + data = {"image": img.clone()} + transform = transforms.AddPatch( + patch_size=0.25, color=(1.0, 0.0, 0.0), position="top_left_corner" + ) + result = transform(data) + # Top-left corner should now contain red pixels + patch_area = result["image"][:, :8, :8] + assert torch.allclose(patch_area[0], torch.ones_like(patch_area[0]), atol=1e-3) + assert torch.allclose( + patch_area[1:], torch.zeros_like(patch_area[1:]), atol=1e-3 + ) + + def test_add_color_tint_transform(self): + """Test AddColorTint applies an additive tint.""" + img = torch.zeros(3, 16, 16) + data = {"image": img} + transform = transforms.AddColorTint(tint=(1.0, 0.5, 0.5), alpha=0.5) + result = transform(data) + # Image should not be all zeros anymore + assert torch.any(result["image"] > 0) + + def test_add_border_transform(self): + """Test AddBorder draws a colored border.""" + img = torch.zeros(3, 20, 20) + data = {"image": img} + transform = transforms.AddBorder(thickness=0.1, color=(0, 1, 0)) + result = transform(data) + # Corners should have green (0,1,0) + assert torch.allclose(result["image"][1, 0, 0], torch.tensor(1.0), atol=1e-3) + assert torch.allclose(result["image"][0, 0, 0], torch.tensor(0.0), atol=1e-3) + + def test_add_watermark_transform(self, tmp_path): + """Test AddWatermark overlays another image.""" + # Create a dummy watermark (white square) + wm_path = tmp_path / "wm.png" + from torchvision.utils import save_image + + save_image(torch.ones(3, 8, 8), wm_path) + data = {"image": torch.zeros(3, 32, 32)} + transform = transforms.AddWatermark( + str(wm_path), size=0.25, position="center", alpha=1.0 + ) + result = transform(data) + # There should be a bright region in the center + center = result["image"][:, 12:20, 12:20] + assert torch.mean(center) > 0.5 + + def test_class_conditional_injector(self): + """Test ClassConditionalInjector applies transform to correct labels only.""" + base_transform = transforms.AddPatch(color=(0, 1, 0)) + injector = transforms.ClassConditionalInjector( + transformation=base_transform, + target_labels=[1], + proportion=1.0, + total_samples=5, + seed=42, + ) + + # Prepare samples with idx + label + samples = [ + {"image": torch.zeros(3, 16, 16), "label": torch.tensor(label), "idx": idx} + for idx, label in enumerate([0, 1, 1, 0, 1]) + ] + + outputs = [injector(s) for s in samples] + + # Check that only samples with label of 1 were modified + for s_in, s_out in zip(samples, outputs): + mean_pixel = s_out["image"].mean().item() + if s_in["label"] == 1: + assert mean_pixel > 0 # patch added + else: + assert mean_pixel == 0 # unchanged