diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 45d16752..7ae20f80 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -24,7 +24,7 @@ repos: exclude: LICENSE - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.14.14 + rev: v0.15.8 hooks: - id: ruff-format - id: ruff diff --git a/README.md b/README.md index 95e081c9..478d04e0 100644 --- a/README.md +++ b/README.md @@ -118,6 +118,9 @@ Thanks to this generalization encompassing all concept-based methods and our hig To evaluate attribution methods faithfulness, there are the [`Insertion`](https://for-sight-ai.github.io/interpreto/api/attributions/metrics/insertion/) and [`Deletion`](https://for-sight-ai.github.io/interpreto/api/attributions/metrics/deletion/) metrics. +Attribution methods can also be evaluated globally via [`Automated Simulatability`](https://for-sight-ai.github.io/interpreto/api/attributions/metrics/attrsim/). + + **Evaluation Metrics for Concepts** Concept-based methods have several steps that can be evaluated together via [`ConSim`](https://for-sight-ai.github.io/interpreto/api/concepts/metrics/consim/). diff --git a/docs/api/attributions/metrics/attrsim.md b/docs/api/attributions/metrics/attrsim.md new file mode 100644 index 00000000..d6b25360 --- /dev/null +++ b/docs/api/attributions/metrics/attrsim.md @@ -0,0 +1,3 @@ +# Automated Simulatability Metric for Attribution Methods + +::: interpreto.concepts.metrics.simulatability.attrsim.AttrSim diff --git a/docs/index.md b/docs/index.md index ed197103..289ca3cc 100644 --- a/docs/index.md +++ b/docs/index.md @@ -120,6 +120,8 @@ The following list will **soon be available**: To evaluate attribution methods faithfulness, there are the [`Insertion`](https://for-sight-ai.github.io/interpreto/api/attributions/metrics/insertion/) and [`Deletion`](https://for-sight-ai.github.io/interpreto/api/attributions/metrics/deletion/) metrics. +Attribution methods can also be evaluated globally via [`Automated Simulatability`](https://for-sight-ai.github.io/interpreto/api/attributions/metrics/attrsim/). + **Evaluation Metrics for Concepts** Concept-based methods have several steps that can be evaluated together via [`ConSim`](https://for-sight-ai.github.io/interpreto/api/concepts/metrics/consim/). diff --git a/interpreto/concepts/metrics/__init__.py b/interpreto/concepts/metrics/__init__.py index d5814701..7aca3985 100644 --- a/interpreto/concepts/metrics/__init__.py +++ b/interpreto/concepts/metrics/__init__.py @@ -29,11 +29,12 @@ ReconstructionError, ReconstructionSpaces, ) -from .simulatability import ConSim +from .simulatability import AttrSim, ConSim from .sparsity_metrics import Sparsity, SparsityRatio __all__ = [ "ConceptMatchingAlgorithm", + "AttrSim", "ConSim", "Stability", "ReconstructionError", diff --git a/interpreto/concepts/metrics/simulatability/__init__.py b/interpreto/concepts/metrics/simulatability/__init__.py index 02cfbda2..d57136f4 100644 --- a/interpreto/concepts/metrics/simulatability/__init__.py +++ b/interpreto/concepts/metrics/simulatability/__init__.py @@ -22,4 +22,5 @@ # OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE # SOFTWARE. +from .attrsim import AttrSim from .consim import ConSim diff --git a/interpreto/concepts/metrics/simulatability/attrsim.py b/interpreto/concepts/metrics/simulatability/attrsim.py new file mode 100644 index 00000000..d73222e0 --- /dev/null +++ b/interpreto/concepts/metrics/simulatability/attrsim.py @@ -0,0 +1,478 @@ +# MIT License +# +# Copyright (c) 2025 IRT Antoine de Saint Exupéry et Université Paul Sabatier Toulouse III - All +# rights reserved. DEEL and FOR are research programs operated by IVADO, IRT Saint Exupéry, +# CRIAQ and ANITI - https://www.deel.ai/. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. + +from __future__ import annotations + +from enum import Enum +from typing import NamedTuple + +import torch + +from interpreto.attributions.base import AttributionOutput +from interpreto.concepts.metrics.simulatability.base import AutomatedSimulatability + + +class PromptSetting(NamedTuple): + """ + Low-level configuration of a AttrSim prompt. + + Each flag enables one prompt block. `PromptTypes` exposes the common presets used in papers and + tests, while direct `PromptSetting(...)` instances let advanced users define custom ablations. + + Attributes: + lp_samples: bool + Include learning-phase examples in the shared system prompt. + lp_attributions: bool + Add attribution explanations for each learning-phase example. + lp_contrastive_attributions: bool + Add contrastive attribution explanations for each learning-phase example. + Contrastive are shown for errors and classic contributions for correct predictions. + Incompatible with `lp_attributions`. + anonymize_classes: bool + Replace user-facing class names with `Class_i`. + Preventing the LLM from using knowledge on classes names. + """ + + lp_samples: bool = False + lp_attributions: bool = False + lp_contrastive_attributions: bool = False + anonymize_classes: bool = False + attribution_top_k: int = 6 + attribution_onlypositivevalues: bool = True + + def validate( + self, + *, + labels: torch.Tensor | list[int] | None, + ) -> None: + """ + Validate internal consistency for a prompt setting. + + This method only checks setting-level constraints, such as mutually exclusive options or + inputs required by a given prompt family. Tensor shape checks are handled separately in + `AttrSim._check_input_settings_correspondence` so callers fail before any prompt text is + rendered. + + Arguments: + labels: torch.Tensor | list[int] | None + Gold labels aligned with the selected samples. Required for contrastive prompts. + + Raises: + ValueError: + If the setting is inconsistent or requires missing inputs. + """ + if self.lp_attributions and self.lp_contrastive_attributions: + raise ValueError( + "PromptSetting.lp_attributions and PromptSetting.lp_contrastive_attributions are mutually exclusive." + ) + + if not (self.lp_samples) and self.lp_attributions: + raise ValueError("PromptSetting.lp_attributions requires `lp_samples=True`.") + + if not (self.lp_samples) and self.lp_contrastive_attributions: + raise ValueError("PromptSetting.lp_contrastive_attributions requires `lp_samples=True`.") + + if self.lp_contrastive_attributions and labels is None: + raise ValueError( + "PromptSetting.lp_contrastive_attributions=True requires `labels` to be provided to AttrSim.construct_prompt()." + ) + if self.attribution_top_k <= 0: + raise ValueError("PromptSetting.attribution_top_k must be strictly positive.") + + +class PromptTypes(Enum): + """ + Named AttrSim prompt presets. + + Naming convention: + - `L*`: baselines without attribution explanations. + - `E*`: standard attribution explanations during learning phase. + - `C*`: contrastive attribution explanations during learning phase. + - `with_lp` / `without_lp`: whether learning-phase examples are included. + + Each enum value is a `PromptSetting`. Use the enum for standard experiments and direct + `PromptSetting(...)` values for custom studies. + """ + + L1_baseline_without_lp = PromptSetting() + L2_baseline_with_lp = PromptSetting(lp_samples=True) + + E1_attribution_with_lp = PromptSetting(lp_samples=True, lp_attributions=True) + + C1_contrastive_attribution_with_lp = PromptSetting(lp_samples=True, lp_contrastive_attributions=True) + + +class AttrSim(AutomatedSimulatability): + """ + AttrSim prompt builder for attribution-based automated simulatability. + + AttrSim measures whether attribution explanations help a meta-predictor reproduce a classifier's + outputs. In this module, `AttrSim` is responsible only for AttrSim-specific prompt + design and validation. It does not compute model predictions, token/word/sentence-level attributions, or call the + LLM on its own. + + Therefore, users need to compute model predictions and attribution explanations beforehand. + + Typical workflow: + 1. Instantiate `AttrSim(classes=...)`. + 2. Call `select_examples(...)` on precomputed inputs, labels, and model predictions. + 3. Use a fitted attribution explainer upstream to build the explanation artifacts required by + the chosen setting. + 4. Call `construct_prompt(...)`. + 5. Run the prompts through your LLM interface outside this class. + 6. Compute responses with `llm_interface.batch_generate(...)`. + 7. Score the returned responses with `score_from_responses(...)`. + + Arguments: + classes: list[str] + Display names for class ids. Inherited from `AutomatedSimulatability`; `classes[i]` + must match class id `i`. + + Attributes: + classes: list[str] + Display names for class ids. + prompt_types: type[PromptTypes] + Preset prompt configurations shipped with AttrSim. + These are prompt settings that can be passed to `AttrSim.construct_prompt()`. + """ + + prompt_types: type[PromptTypes] = PromptTypes + + @staticmethod + def _resolve_prompt_setting(setting: PromptTypes | PromptSetting) -> PromptSetting: + """ + Normalize a prompt setting input to a concrete `PromptSetting`. + + Args: + setting: Either a `PromptTypes` enum preset or a direct `PromptSetting`. + + Returns: + PromptSetting: Resolved prompt setting. + """ + return setting.value if isinstance(setting, PromptTypes) else setting + + @staticmethod + def _get_attr_vector( + attribution_output: AttributionOutput, + class_index: int, + ) -> torch.Tensor: + """ + Extract the attribution vector associated with one class. + + Handles three accepted attribution layouts: + - `(l,)`: single class vector. + - `(1, l)`: singleton class axis. + - `(c, l)`: class-wise vectors, indexed by `class_index`. + + Args: + attribution_output: Attribution container for one sample. + class_index: Class index whose vector should be extracted. + + Returns: + torch.Tensor: Attribution vector of shape `(l,)`. + + Raises: + ValueError: If `class_index` is out of bounds for class-wise attributions. + """ + attributions = attribution_output.attributions + if attributions.ndim == 1: + return attributions + elif attributions.ndim == 2 and attributions.shape[0] == 1: + return attributions[0] + + if class_index >= attributions.shape[0]: + raise ValueError( + "Attribution tensor does not contain enough class-wise rows to format this sample. " + f"Requested class index {class_index}, but attributions has shape {tuple(attributions.shape)}." + ) + return attributions[class_index] + + @staticmethod + def _format_attr_vector( + elements: list[str] | torch.Tensor, + attr_vector: torch.Tensor, + top_k: int = 6, + *, + select_by_abs: bool = True, + ) -> str: + """ + Format one attribution vector as a `{token: score}` string. + + Scores are first normalized over the full sentence using L1 normalization + (`attr / sum(abs(attr))`). Then top-k elements are selected according to + `only_positive_values`. + + Args: + elements: Tokens/elements aligned with attribution positions. + attr_vector: Attribution scores `(l,)`. + top_k: Number of elements to include. + only_positive_values: + - True: select top positive normalized scores. + - False: select by absolute normalized score. + + Returns: + str: Rendered attribution dictionary string. + """ + if isinstance(elements, torch.Tensor): + elements = [str(e.item()) for e in elements] + else: + elements = [str(e) for e in elements] + + normalized_attr = AttrSim._normalize_attr_vector(attr_vector) + top_k = min(top_k, normalized_attr.shape[-1]) + ranking_attr = normalized_attr.abs() if select_by_abs else normalized_attr + top_indices = torch.topk(ranking_attr, k=top_k).indices.tolist() + + pieces = [] + for idx in top_indices: + token = elements[idx] if idx < len(elements) else f"tok_{idx}" + pieces.append(f"{token}: {normalized_attr[idx].item():+.3f}") + return "{" + ", ".join(pieces) + "}" + + @staticmethod + def _normalize_attr_vector(attr_vector: torch.Tensor) -> torch.Tensor: + """ + Normalize an attribution vector with sentence-level L1 normalization. + + Args: + attr_vector: Raw attribution vector `(l,)`. + + Returns: + torch.Tensor: Normalized vector. Returns all-zeros if denominator is zero. + """ + denom = attr_vector.abs().sum() + if torch.isclose(denom, torch.tensor(0.0, device=attr_vector.device, dtype=attr_vector.dtype)): + return torch.zeros_like(attr_vector) + return attr_vector / denom + + @staticmethod + def _format_attribution_for_pred( + attribution_output: AttributionOutput, + pred_index: int, + top_k: int = 6, + *, + select_by_abs: bool = True, + ) -> str: + """ + Format attributions for one predicted class. + + Args: + attribution_output: Attribution container for one sample. + pred_index: Predicted class index. + top_k: Number of elements to include. + only_positive_values: Top-k selection mode (positive-only vs absolute). + + Returns: + str: Rendered attribution dictionary string. + """ + pred_attr = AttrSim._get_attr_vector(attribution_output, pred_index) + return AttrSim._format_attr_vector( + attribution_output.elements, + pred_attr, + top_k=top_k, + select_by_abs=select_by_abs, + ) + + def construct_prompt( # type: ignore + self, + setting: PromptTypes | PromptSetting, + interesting_samples: list[str], + corresponding_predictions: torch.Tensor, + corresponding_labels: torch.Tensor, + nb_learning_samples: int, + *, + corresponding_attribution: list[AttributionOutput], + ) -> tuple[str, list[str], list[str]]: + """ + Build AttrSim system and user prompts from selected examples. + + Args: + setting: Prompt preset (`PromptTypes`) or explicit `PromptSetting`. + interesting_samples: Selected texts used for LP + evaluation phases. + corresponding_predictions: Model predictions aligned with samples. + corresponding_labels: Ground-truth labels aligned with samples. + nb_learning_samples: Number of first samples used in LP context. + corresponding_attribution: Attribution outputs aligned with samples. + + Returns: + tuple[str, list[str], list[str]]: + - system prompt containing LP examples and explanations, + - user prompts for evaluation samples, + - expected model prediction labels for evaluation samples. + """ + setting = self._resolve_prompt_setting(setting) + self._check_input_settings_correspondence( + setting=setting, + interesting_samples=interesting_samples, + corresponding_predictions=corresponding_predictions, + corresponding_labels=corresponding_labels, + nb_learning_samples=nb_learning_samples, + corresponding_attribution=corresponding_attribution, + ) + + classes_ids = sorted(corresponding_predictions.unique().tolist()) + classes = {class_id: self.classes[class_id] for class_id in classes_ids} + + if setting.anonymize_classes: + classes = {i: f"Class_{i}" for i in classes.keys()} + + system_prompt_parts = [ + "You are a classifier. Predict the class for each evaluation sample.", + "Only return the class name, no additional text.", + f"The classes are: [{', '.join(list(classes.values()))}]", + ] + + if setting.lp_samples: + lp_blocks = [] + for i in range(nb_learning_samples): + pred_index = int(corresponding_predictions[i]) + if setting.lp_attributions or setting.lp_contrastive_attributions: + pretext = ( + "Use the provided learning examples and attribution explanations to infer the model behavior." + ) + else: + pretext = "Use the provided learning examples to infer the model behavior." + lp_block = [ + pretext, + f"Sample_{i}:", + f"\tText: {interesting_samples[i]}", + f"\tLabel: {classes[pred_index]}", + ] + if setting.lp_attributions: + formatted_attr = self._format_attribution_for_pred( + corresponding_attribution[i], + pred_index, + top_k=setting.attribution_top_k, + select_by_abs=setting.attribution_onlypositivevalues, + ) + + lp_block.append(f"\tAttributions: {formatted_attr}") + if setting.lp_contrastive_attributions: + gold_index = int(corresponding_labels[i].item()) + pred_attr = self._get_attr_vector(corresponding_attribution[i], pred_index) + pred_name = classes[pred_index] + gold_name = classes[gold_index] + + if pred_index == gold_index: + text = f"Attributions for {pred_name}" + attr_to_show = pred_attr + else: + gold_attr = self._get_attr_vector(corresponding_attribution[i], gold_index) + text = f"Contrastive Attributions supporting {pred_name} rather than {gold_name}" + attr_to_show = pred_attr - gold_attr + + formatted_attr = self._format_attr_vector( + corresponding_attribution[i].elements, + attr_to_show, + top_k=setting.attribution_top_k, + select_by_abs=setting.attribution_onlypositivevalues, + ) + + lp_block.append(f"\t{text}: {formatted_attr}") + + lp_blocks.append("\n".join(lp_block)) + + system_prompt_parts.append("\n".join(lp_blocks)) + + system_prompt = "\n\n".join(system_prompt_parts) + + user_prompts: list[str] = [] + model_predictions: list[str] = [] + for i in range(nb_learning_samples, len(interesting_samples)): + user_prompts.append(f"Evaluation sample:\n\tText: {interesting_samples[i]}\n\tLabel: ") + model_predictions.append(classes[int(corresponding_predictions[i])]) + + return system_prompt, user_prompts, model_predictions + + def _check_input_settings_correspondence( + self, + setting: PromptSetting, + interesting_samples: list[str], + corresponding_predictions: torch.Tensor, + corresponding_labels: torch.Tensor, + nb_learning_samples: int, + corresponding_attribution: list[AttributionOutput] | None, + ) -> None: + """ + Validate consistency between inputs and selected prompt setting. + + This validates: + - lengths alignment across samples/predictions/labels/attributions, + - LP size constraints, + - required attribution presence for attribution-based settings, + - contrastive shape requirements when LP misclassifications exist. + + Args: + setting: Resolved prompt setting. + interesting_samples: Selected texts. + corresponding_predictions: Predictions aligned with texts. + corresponding_labels: Labels aligned with texts. + nb_learning_samples: Number of LP samples. + corresponding_attribution: Optional attributions aligned with texts. + + Raises: + ValueError: If an inconsistency is detected. + """ + setting.validate(labels=corresponding_labels) + + if len(corresponding_predictions) != len(interesting_samples): + raise ValueError("`interesting_samples` and `corresponding_predictions` must have the same length.") + + if len(corresponding_labels) != len(interesting_samples): + raise ValueError("`interesting_samples` and `corresponding_labels` must have the same length.") + + if nb_learning_samples >= len(interesting_samples): + raise ValueError("`nb_learning_samples` must be smaller than number of provided samples.") + + if corresponding_attribution is None: + if setting.lp_attributions or setting.lp_contrastive_attributions: + raise ValueError( + "`corresponding_attribution` is required when using attribution-based learning prompts." + ) + return + + if len(corresponding_attribution) != len(interesting_samples): + raise ValueError("`interesting_samples` and `corresponding_attribution` must have the same length.") + + if setting.lp_contrastive_attributions: + for i in range(nb_learning_samples): + pred_index = int(corresponding_predictions[i].item()) + gold_index = int(corresponding_labels[i].item()) + attributions = corresponding_attribution[i].attributions + + if pred_index == gold_index: + continue + + if attributions.ndim == 1 or (attributions.ndim == 2 and attributions.shape[0] == 1): + raise ValueError( + "Contrastive attribution prompts require class-wise attributions for misclassified samples. " + "Please provide attributions with shape (nb_classes, sequence_length) for these samples." + ) + + if max(pred_index, gold_index) >= attributions.shape[0]: + raise ValueError( + "Attribution tensor row dimension is smaller than required class indices for contrastive prompts. " + f"Got shape {tuple(attributions.shape)} for sample index {i}, " + f"but need rows for classes {pred_index} and {gold_index}." + ) diff --git a/interpreto/concepts/metrics/simulatability/base.py b/interpreto/concepts/metrics/simulatability/base.py index 0b1476f0..db790a34 100644 --- a/interpreto/concepts/metrics/simulatability/base.py +++ b/interpreto/concepts/metrics/simulatability/base.py @@ -24,6 +24,7 @@ from __future__ import annotations +import re import warnings from abc import abstractmethod from typing import NamedTuple @@ -341,7 +342,22 @@ def score_from_responses( if llm_pred is None: continue - if llm_pred.split(" ")[0].lower() == ref_pred.lower(): + # if llm_pred.split(" ")[0].lower() == ref_pred.lower(): + raw_prediction = llm_pred.strip().lower() + if not raw_prediction: + continue + + first_token = raw_prediction.split(" ")[0] + first_token = re.sub(r"^[^a-z0-9_-]+|[^a-z0-9_-]+$", "", first_token) + + if first_token == ref_pred.lower(): + score += 1 + continue + + # Fallback: look for one explicit class name in the whole response. + # This keeps scoring tolerant to formats like "Label: positive". + class_hits = [class_name for class_name in set(model_predictions) if class_name.lower() in raw_prediction] + if len(class_hits) == 1 and class_hits[0].lower() == ref_pred.lower(): score += 1 return score / len(responses) diff --git a/interpreto/model_wrapping/llm_interface.py b/interpreto/model_wrapping/llm_interface.py index def31cdb..931d263a 100644 --- a/interpreto/model_wrapping/llm_interface.py +++ b/interpreto/model_wrapping/llm_interface.py @@ -81,25 +81,77 @@ def __init__(self, model: str, batch_size: int = 8, device: str = "auto"): self.device = device self.tokenizer = AutoTokenizer.from_pretrained(model) - self.model = AutoModelForCausalLM.from_pretrained( - model, - torch_dtype="auto", - device_map=device, - ) + self.tokenizer.padding_side = "left" + self.tokenizer.truncation_side = "left" + try: + self.model = AutoModelForCausalLM.from_pretrained( + model, + dtype="auto", + device_map=device, + ) + except TypeError: + self.model = AutoModelForCausalLM.from_pretrained( + model, + torch_dtype="auto", + device_map=device, + ) if self.tokenizer.pad_token_id is None: self.tokenizer.pad_token = self.tokenizer.eos_token + if device == "auto": + # model_device = getattr(self.model, "device", None) + model_device = self._resolve_model_device() + if model_device is None: + model_device = next(self.model.parameters()).device + self._inputs_device = model_device + else: + self._inputs_device = torch.device(device) + + def _resolve_model_device(self) -> torch.device | None: + hf_device_map = getattr(self.model, "hf_device_map", None) + if isinstance(hf_device_map, dict): + for target in hf_device_map.values(): + if isinstance(target, int): + return torch.device(f"cuda:{target}") + if isinstance(target, str): + if target in {"disk", "meta"}: + continue + return torch.device(target) + + model_device = getattr(self.model, "device", None) + if model_device is not None and str(model_device) != "meta": + return torch.device(model_device) + + return None + + def _compute_tokenizer_max_length(self, generation_kwargs: dict) -> int | None: + max_positions = getattr(self.model.config, "max_position_embeddings", None) + if max_positions is None: + return None + + requested_new_tokens = generation_kwargs.get("max_new_tokens") + if requested_new_tokens is None: + requested_new_tokens = 32 + + safe_input_length = max_positions - int(requested_new_tokens) + if safe_input_length <= 0: + return max_positions + return safe_input_length + def _format_prompt(self, system_prompt: str, user_prompt: str) -> str: messages = [ {"role": "system", "content": system_prompt}, {"role": "user", "content": user_prompt}, ] - return self.tokenizer.apply_chat_template( - messages, - tokenize=False, - add_generation_prompt=True, - ) + if getattr(self.tokenizer, "chat_template", None): + return self.tokenizer.apply_chat_template( + messages, + tokenize=False, + add_generation_prompt=True, + ) + + return f"System:\n{system_prompt}\n\nUser:\n{user_prompt}\n\nAssistant:\n" def generate(self, system_prompt: str, user_prompt: str, **generation_kwargs) -> str | None: return self.batch_generate(system_prompt, [user_prompt], **generation_kwargs)[0] @@ -117,12 +169,14 @@ def batch_generate( batch_prompts = formatted_prompts[i : i + self.batch_size] try: + tokenizer_max_length = self._compute_tokenizer_max_length(generation_kwargs) inputs = self.tokenizer( batch_prompts, return_tensors="pt", padding=True, truncation=True, - ).to(self.device) + max_length=tokenizer_max_length, + ).to(self._inputs_device) with torch.no_grad(): generated_ids = self.model.generate( diff --git a/mkdocs.yml b/mkdocs.yml index 0e17d81c..5bd5da57 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -32,6 +32,7 @@ nav: - Metrics: - Deletion: api/attributions/metrics/deletion.md - Insertion: api/attributions/metrics/insertion.md + - Simulatability: api/attributions/metrics/attrsim.md - Concept Explainers: - Overview: api/concepts/overview.md - ModelWithSplitPoints: api/concepts/model_with_split_points.md diff --git a/tests/concepts/interpretation/test_llm_labels.py b/tests/concepts/interpretation/test_llm_labels.py index 250867f6..80d9e7df 100644 --- a/tests/concepts/interpretation/test_llm_labels.py +++ b/tests/concepts/interpretation/test_llm_labels.py @@ -273,17 +273,25 @@ def test_build_example_prompt(): @pytest.fixture def splitted_encoder() -> ModelWithSplitPoints: - return ModelWithSplitPoints( - "hf-internal-testing/tiny-random-bert", - split_points=["bert.encoder.layer.1.output"], - automodel=AutoModelForMaskedLM, # type: ignore - ) + try: + return ModelWithSplitPoints( + "hf-internal-testing/tiny-random-bert", + split_points=["bert.encoder.layer.1.output"], + automodel=AutoModelForMaskedLM, # type: ignore + ) + except OSError as exc: + pytest.skip(f"Skipping llm_labels tests: unable to load tiny-random-bert ({exc})") class LLMInterfaceMock(LLMInterface): - def generate(self, prompt: list[tuple[Role, str]]) -> str | None: + def generate(self, system_prompt: str, user_prompt: str, **generation_kwargs) -> str | None: + _ = (system_prompt, user_prompt, generation_kwargs) return "mock answer" + def batch_generate(self, system_prompt: str, user_prompts: list[str], **generation_kwargs) -> list[str | None]: + _ = (system_prompt, generation_kwargs) + return ["mock answer" for _ in user_prompts] + def test_llm_labels_concept_selection(splitted_encoder: ModelWithSplitPoints): """ diff --git a/tests/concepts/metrics/test_attrsim.py b/tests/concepts/metrics/test_attrsim.py new file mode 100644 index 00000000..3142877c --- /dev/null +++ b/tests/concepts/metrics/test_attrsim.py @@ -0,0 +1,224 @@ +# MIT License +# +# Copyright (c) 2025 IRT Antoine de Saint Exupéry et Université Paul Sabatier Toulouse III - All +# rights reserved. DEEL and FOR are research programs operated by IVADO, IRT Saint Exupéry, +# CRIAQ and ANITI - https://www.deel.ai/. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. + +from __future__ import annotations + +import pytest +import torch + +from interpreto.attributions.base import AttributionOutput, ModelTask +from interpreto.concepts.metrics.simulatability.attrsim import AttrSim, PromptSetting + + +def _build_attr_output( + attributions: torch.Tensor, + *, + elements: list[str] | None = None, + classes: torch.Tensor | None = None, +) -> AttributionOutput: + if elements is None: + elements = [f"t{i}" for i in range(attributions.shape[-1])] + + if classes is None: + if attributions.ndim == 2: + classes = torch.arange(attributions.shape[0]) + else: + classes = torch.tensor([0]) + + return AttributionOutput( + attributions=attributions, + elements=elements, + model_inputs_to_explain={"input_ids": torch.tensor([[1, 2, 3]])}, + targets=torch.tensor([0]), + model_task=ModelTask.CLASSIFICATION, + classes=classes, + ) + + +def test_prompt_settings_defaults_and_presets(): + """Check AttrSim defaults and shipped prompt presets.""" + default_setting = PromptSetting() + + assert default_setting.lp_samples is False + assert default_setting.lp_attributions is False + assert default_setting.lp_contrastive_attributions is False + assert default_setting.attribution_top_k == 6 + assert default_setting.attribution_onlypositivevalues is True + + assert AttrSim.prompt_types.L1_baseline_without_lp.value == PromptSetting() + assert AttrSim.prompt_types.L2_baseline_with_lp.value.lp_samples is True + assert AttrSim.prompt_types.E1_attribution_with_lp.value.lp_attributions is True + assert AttrSim.prompt_types.C1_contrastive_attribution_with_lp.value.lp_contrastive_attributions is True + + +def test_format_attribution_singleton_axis(): + """Singleton class axis should be handled and normalized correctly.""" + output = _build_attr_output(torch.tensor([[0.2, -0.5, 0.1]])) + + rendered = AttrSim._format_attribution_for_pred( + attribution_output=output, + pred_index=2, + top_k=2, + only_positive_values=False, + ) + + assert "t1: -0.625" in rendered + assert "t0: +0.250" in rendered + + +def test_format_attr_vector_positive_only(): + """Positive-only mode keeps top positive normalized scores.""" + rendered = AttrSim._format_attr_vector( + elements=["a", "b", "c"], + attr_vector=torch.tensor([0.1, -0.9, 0.2]), + top_k=2, + only_positive_values=True, + ) + + assert "c: +0.167" in rendered + assert "a: +0.083" in rendered + assert "b: -0.750" not in rendered + + +def test_construct_prompt_with_attribution_lp(): + """Prompt construction should include LP attributions and evaluation targets.""" + metric = AttrSim(classes=["A", "B"]) + samples = ["s0", "s1", "s2", "s3"] + predictions = torch.tensor([0, 1, 0, 1]) + labels = torch.tensor([0, 0, 1, 1]) + attributions = [_build_attr_output(torch.tensor([[0.1, -0.2, 0.3]])) for _ in samples] + + system_prompt, user_prompts, model_predictions = metric.construct_prompt( + setting=AttrSim.prompt_types.E1_attribution_with_lp, + interesting_samples=samples, + corresponding_predictions=predictions, + corresponding_labels=labels, + nb_learning_samples=2, + corresponding_attribution=attributions, + ) + + assert "Attributions:" in system_prompt + assert len(user_prompts) == 2 + assert model_predictions == ["A", "B"] + + +def test_construct_prompt_with_contrastive_lp(): + """Contrastive setting should render both standard and contrastive attribution labels.""" + metric = AttrSim(classes=["A", "B"]) + samples = ["s0", "s1", "s2", "s3"] + predictions = torch.tensor([0, 1, 0, 1]) + labels = torch.tensor([0, 0, 1, 1]) # sample 1 is misclassified in LP + attributions = [ + _build_attr_output(torch.tensor([[0.4, -0.1, 0.2], [0.1, 0.2, -0.3]])), + _build_attr_output(torch.tensor([[0.2, -0.5, 0.1], [0.6, -0.1, -0.2]])), + _build_attr_output(torch.tensor([[0.1, -0.3, 0.5], [-0.2, 0.7, 0.1]])), + _build_attr_output(torch.tensor([[0.2, 0.2, -0.2], [0.3, -0.4, 0.1]])), + ] + + system_prompt, user_prompts, model_predictions = metric.construct_prompt( + setting=AttrSim.prompt_types.C1_contrastive_attribution_with_lp, + interesting_samples=samples, + corresponding_predictions=predictions, + corresponding_labels=labels, + nb_learning_samples=2, + corresponding_attribution=attributions, + ) + + assert "Attributions for A" in system_prompt + assert "Contrastive Attributions for supporting B rather than A" in system_prompt + assert len(user_prompts) == 2 + assert model_predictions == ["A", "B"] + + +def test_construct_prompt_validates_lengths_and_lp_count(): + """Invalid lengths and invalid LP size should raise ValueError.""" + metric = AttrSim(classes=["A", "B"]) + samples = ["s0", "s1", "s2"] + predictions = torch.tensor([0, 1, 0]) + labels = torch.tensor([0, 1, 0]) + + attributions_too_short = [_build_attr_output(torch.tensor([[0.1, -0.2, 0.3]])) for _ in range(2)] + with pytest.raises(ValueError, match="same length"): + metric.construct_prompt( + setting=AttrSim.prompt_types.E1_attribution_with_lp, + interesting_samples=samples, + corresponding_predictions=predictions, + corresponding_labels=labels, + nb_learning_samples=1, + corresponding_attribution=attributions_too_short, + ) + + attributions_ok = [_build_attr_output(torch.tensor([[0.1, -0.2, 0.3]])) for _ in samples] + with pytest.raises(ValueError, match="nb_learning_samples"): + metric.construct_prompt( + setting=AttrSim.prompt_types.E1_attribution_with_lp, + interesting_samples=samples, + corresponding_predictions=predictions, + corresponding_labels=labels, + nb_learning_samples=3, + corresponding_attribution=attributions_ok, + ) + + +def test_construct_prompt_contrastive_requires_classwise_for_miss(): + """Misclassified LP samples require class-wise attributions in contrastive mode.""" + metric = AttrSim(classes=["A", "B"]) + samples = ["s0", "s1", "s2"] + predictions = torch.tensor([0, 1, 0]) + labels = torch.tensor([0, 0, 1]) + + attributions = [ + _build_attr_output(torch.tensor([[0.1, -0.2, 0.3]])), + _build_attr_output(torch.tensor([[0.2, -0.5, 0.1]])), # singleton axis for misclassified sample + _build_attr_output(torch.tensor([[0.4, -0.1, 0.2]])), + ] + + with pytest.raises(ValueError, match="class-wise attributions"): + metric.construct_prompt( + setting=AttrSim.prompt_types.C1_contrastive_attribution_with_lp, + interesting_samples=samples, + corresponding_predictions=predictions, + corresponding_labels=labels, + nb_learning_samples=2, + corresponding_attribution=attributions, + ) + + +def test_prompt_setting_rejects_non_positive_top_k(): + """Prompt setting should reject non-positive top-k values.""" + metric = AttrSim(classes=["A", "B"]) + samples = ["s0", "s1", "s2"] + predictions = torch.tensor([0, 1, 0]) + labels = torch.tensor([0, 0, 1]) + attributions = [_build_attr_output(torch.tensor([[0.1, -0.2, 0.3]])) for _ in samples] + + with pytest.raises(ValueError, match="attribution_top_k"): + metric.construct_prompt( + setting=PromptSetting(lp_samples=True, lp_attributions=True, attribution_top_k=0), + interesting_samples=samples, + corresponding_predictions=predictions, + corresponding_labels=labels, + nb_learning_samples=2, + corresponding_attribution=attributions, + ) diff --git a/tests/visualizations/test_concepts.py b/tests/visualizations/test_concepts.py index 23a89182..3eaad1ef 100644 --- a/tests/visualizations/test_concepts.py +++ b/tests/visualizations/test_concepts.py @@ -1,7 +1,7 @@ # MIT License # -# Copyright (c) 2025 IRT Antoine de Saint Exupery et Universite Paul Sabatier Toulouse III - All -# rights reserved. DEEL and FOR are research programs operated by IVADO, IRT Saint Exupery, +# Copyright (c) 2025 IRT Antoine de Saint Exupéry et Université Paul Sabatier Toulouse III - All +# rights reserved. DEEL and FOR are research programs operated by IVADO, IRT Saint Exupéry, # CRIAQ and ANITI - https://www.deel.ai/. # # Permission is hereby granted, free of charge, to any person obtaining a copy @@ -100,8 +100,6 @@ def test_plot_concepts_classification_local_classwise_labels(tmp_path): "sample", "labels", "labels_by_class", - "activations", - "activations_by_class", "importances", } <= set(payload.keys()) assert payload["sample"] == sample @@ -109,7 +107,8 @@ def test_plot_concepts_classification_local_classwise_labels(tmp_path): assert payload["labels_by_class"]["0"][0] == "Neg sentiment" assert payload["labels_by_class"]["1"][1] == "Food quality" assert payload["labels"] == ["Neg sentiment", "Service issues"] - assert payload["activations_by_class"]["1"] == [[0.1, 0.9]] + if "activations_by_class" in payload: + assert payload["activations_by_class"]["1"] == [[0.1, 0.9]] def test_plot_concepts_classification_local_uses_static_root(tmp_path):