From 3504b2831bf36861833077aae559c05c0c23e857 Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Mon, 13 Apr 2026 09:01:29 +0000 Subject: [PATCH 01/17] try huggingfaceLLM --- docs/notebooks/test.ipynb | 2692 +++++++++++++++++++++++++++++++++++++ 1 file changed, 2692 insertions(+) create mode 100644 docs/notebooks/test.ipynb diff --git a/docs/notebooks/test.ipynb b/docs/notebooks/test.ipynb new file mode 100644 index 00000000..8378f218 --- /dev/null +++ b/docs/notebooks/test.ipynb @@ -0,0 +1,2692 @@ +{ + "cells": [ + { + "cell_type": "code", + "execution_count": null, + "id": "193bf4f9", + "metadata": {}, + "outputs": [ + { + "data": { + "application/vnd.jupyter.widget-view+json": { + "model_id": "b030d96f6f54448eb5589fb47d8c6ed8", + "version_major": 2, + "version_minor": 0 + }, + "text/plain": [ + "Loading weights: 0%| | 0/148 [00:00

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + } + ], + "source": [ + "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", + "\n", + "model_gen.to(device)\n", + "\n", + "text = \"Hi there, how are you?\"\n", + "tokenized_inputs = tokenizer_gen(text, return_tensors=\"pt\").to(device)\n", + "print(tokenized_inputs)\n", + "\n", + "target = model_gen.generate(**tokenized_inputs, max_length=16)\n", + "\n", + "print(target)\n", + "\n", + "\n", + "explainer = SmoothGrad(\n", + " model_gen,\n", + " tokenizer_gen,\n", + " granularity=Granularity.WORD,\n", + " granularity_aggregation_strategy=GranularityAggregationStrategy.MAX,\n", + ")\n", + "\n", + "\n", + "attributions = explainer(tokenized_inputs, target)\n", + "plot_attributions(attributions[0])" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": ".venv (3.12.3)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} From 26dd81a4cf66842776d2674de2ca479af6ce8626 Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Mon, 13 Apr 2026 15:55:22 +0000 Subject: [PATCH 02/17] try huggingfaceLLM --- .pre-commit-config.yaml | 2 +- docs/notebooks/test.ipynb | 2735 +---------------- .../concepts/metrics/simulatability/base.py | 18 +- .../concepts/metrics/simulatability/consim.py | 8 + interpreto/model_wrapping/llm_interface.py | 86 +- .../interpretation/test_llm_labels.py | 26 +- tests/visualizations/test_concepts.py | 29 +- 7 files changed, 255 insertions(+), 2649 deletions(-) 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/docs/notebooks/test.ipynb b/docs/notebooks/test.ipynb index 8378f218..babdc986 100644 --- a/docs/notebooks/test.ipynb +++ b/docs/notebooks/test.ipynb @@ -2,2669 +2,150 @@ "cells": [ { "cell_type": "code", - "execution_count": null, - "id": "193bf4f9", + "execution_count": 1, + "id": "274dccbe", "metadata": {}, - "outputs": [ - { - "data": { - "application/vnd.jupyter.widget-view+json": { - "model_id": "b030d96f6f54448eb5589fb47d8c6ed8", - "version_major": 2, - "version_minor": 0 - }, - "text/plain": [ - "Loading weights: 0%| | 0/148 [00:00

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" } ], "source": [ - "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", - "\n", - "model_gen.to(device)\n", - "\n", - "text = \"Hi there, how are you?\"\n", - "tokenized_inputs = tokenizer_gen(text, return_tensors=\"pt\").to(device)\n", - "print(tokenized_inputs)\n", - "\n", - "target = model_gen.generate(**tokenized_inputs, max_length=16)\n", + "# HuggingFaceLLM now supports both chat-template and plain CausalLM tokenizers.\n", + "# Choose any local/cached CausalLM checkpoint.\n", "\n", - "print(target)\n", + "llm = HuggingFaceLLM(model=\"Qwen/Qwen3-0.6B\", batch_size=2, device=\"auto\")\n", "\n", - "\n", - "explainer = SmoothGrad(\n", - " model_gen,\n", - " tokenizer_gen,\n", - " granularity=Granularity.WORD,\n", - " granularity_aggregation_strategy=GranularityAggregationStrategy.MAX,\n", + "responses = llm.batch_generate(\n", + " system_prompt,\n", + " user_prompts,\n", + " max_new_tokens=40,\n", + " do_sample=False,\n", ")\n", "\n", - "\n", - "attributions = explainer(tokenized_inputs, target)\n", - "plot_attributions(attributions[0])" + "print(\"Raw responses:\")\n", + "for response in responses:\n", + " print(\"---\")\n", + " print(response)" + ] + }, + { + "cell_type": "code", + "execution_count": 4, + "id": "d950fc53", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "ConSim score: 0.0\n" + ] + } + ], + "source": [ + "score = metric.score_from_responses(responses, model_predictions)\n", + "print(\"ConSim score:\", score)" ] } ], 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/concepts/metrics/simulatability/consim.py b/interpreto/concepts/metrics/simulatability/consim.py index 5387d2ee..d42215ab 100644 --- a/interpreto/concepts/metrics/simulatability/consim.py +++ b/interpreto/concepts/metrics/simulatability/consim.py @@ -517,6 +517,14 @@ def _setting_to_prompt( # type: ignore[override] # noqa: PLR0912 # ignore too "The most important concepts and their importance for each class are:\n" + "\n".join( [ + # f"\t{class_name}: { + # ConSim._concepts_to_string( + # global_importances[class_index], + # concepts_interpretation, + # top_k=top_k, + # threshold=importance_threshold, + # ) + # }" f"\t{class_name}: { ConSim._concepts_to_string( global_importances[class_index], diff --git a/interpreto/model_wrapping/llm_interface.py b/interpreto/model_wrapping/llm_interface.py index def31cdb..b9e471ad 100644 --- a/interpreto/model_wrapping/llm_interface.py +++ b/interpreto/model_wrapping/llm_interface.py @@ -81,25 +81,87 @@ 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" + # self.model = AutoModelForCausalLM.from_pretrained( + # model, + # torch_dtype="auto", + # device_map=device, + # ) + 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, - ) + # 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 +179,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/tests/concepts/interpretation/test_llm_labels.py b/tests/concepts/interpretation/test_llm_labels.py index 250867f6..daffabf3 100644 --- a/tests/concepts/interpretation/test_llm_labels.py +++ b/tests/concepts/interpretation/test_llm_labels.py @@ -273,17 +273,31 @@ 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 - ) + # 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, 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/visualizations/test_concepts.py b/tests/visualizations/test_concepts.py index 23a89182..f3f798dc 100644 --- a/tests/visualizations/test_concepts.py +++ b/tests/visualizations/test_concepts.py @@ -1,3 +1,27 @@ +# 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. + # MIT License # # Copyright (c) 2025 IRT Antoine de Saint Exupery et Universite Paul Sabatier Toulouse III - All @@ -100,8 +124,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 +131,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): From 3d8fd478b4a0133aadea95c8972820b0445d71d3 Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Tue, 14 Apr 2026 15:06:12 +0000 Subject: [PATCH 03/17] add attrsim --- docs/notebooks/test_sentences.ipynb | 41548 ++++++++++++++++ interpreto/concepts/metrics/__init__.py | 3 +- .../metrics/simulatability/__init__.py | 1 + .../metrics/simulatability/attrsim.py | 143 + tests/concepts/metrics/test_attrsim.py | 24 + 5 files changed, 41718 insertions(+), 1 deletion(-) create mode 100644 docs/notebooks/test_sentences.ipynb create mode 100644 interpreto/concepts/metrics/simulatability/attrsim.py create mode 100644 tests/concepts/metrics/test_attrsim.py diff --git a/docs/notebooks/test_sentences.ipynb b/docs/notebooks/test_sentences.ipynb new file mode 100644 index 00000000..abe03e5e --- /dev/null +++ b/docs/notebooks/test_sentences.ipynb @@ -0,0 +1,41548 @@ +{ + "cells": [ + { + "cell_type": "code", + "execution_count": 8, + "id": "7eb999b1", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "The autoreload extension is already loaded. To reload it, use:\n", + " %reload_ext autoreload\n" + ] + } + ], + "source": [ + "%load_ext autoreload\n", + "%autoreload 2\n", + "\n", + "import sys\n", + "\n", + "sys.path.append(\"../..\")\n", + "\n", + "from transformers import AutoModelForCausalLM, AutoTokenizer\n", + "\n", + "from interpreto import (\n", + " Granularity,\n", + " plot_attributions,\n", + ")\n", + "from interpreto.commons import GranularityAggregationStrategy" + ] + }, + { + "cell_type": "markdown", + "id": "ae0b758d", + "metadata": {}, + "source": [ + "Modèles testés qui marchent:\n", + "- gpt2\n", + "- \n", + "\n", + "\n", + "Modèles testés qui ne marchent pas:\n", + "- Qwen/Qwen3.5-0.8B\n", + "- mistralai/Mistral-7B-v0.1" + ] + }, + { + "cell_type": "code", + "execution_count": 9, + "id": "ff3d10c5", + "metadata": {}, + "outputs": [ + { + "data": { + "application/vnd.jupyter.widget-view+json": { + "model_id": "13d73f90100e428fbe60f339b08a0dd3", + "version_major": 2, + "version_minor": 0 + }, + "text/plain": [ + "Loading weights: 0%| | 0/64 [00:00

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 2\n", + "Attributions shape: torch.Size([5, 9])\n", + "Number of elements: 9\n", + "Attribution: tensor([[-9.2138e-03, -9.2428e-03, 7.6194e-03, -2.8685e-02, nan,\n", + " nan, nan, nan, nan],\n", + " [ 1.6982e-02, 1.1031e-02, -6.1722e-03, -4.1739e-02, -1.0123e-01,\n", + " nan, nan, nan, nan],\n", + " [-1.6126e-03, -4.8020e-04, 4.5713e-04, 6.0391e-05, 8.5584e-03,\n", + " -1.7940e-02, nan, nan, nan],\n", + " [ 3.1949e-03, -2.8566e-04, 7.6152e-03, -1.8436e-02, -1.2982e-02,\n", + " 9.0746e-03, 2.5908e-01, nan, nan],\n", + " [-1.9200e-02, 1.2491e-02, -3.2375e-03, -8.9267e-03, -2.3039e-03,\n", + " 1.1168e-03, -3.0267e-03, 6.4192e-02, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 3\n", + "Attributions shape: torch.Size([1, 4])\n", + "Number of elements: 4\n", + "Attribution: tensor([[8.9990e-05, 1.2292e-02, 1.7112e-02, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 4\n", + "Attributions shape: torch.Size([3, 5])\n", + "Number of elements: 5\n", + "Attribution: tensor([[-0.0237, -0.0976, nan, nan, nan],\n", + " [-0.0186, 0.0137, -0.0373, nan, nan],\n", + " [-0.0087, -0.0029, 0.0004, -0.0142, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 5\n", + "Attributions shape: torch.Size([3, 11])\n", + "Number of elements: 11\n", + "Attribution: tensor([[ 0.0047, 0.0133, 0.0140, 0.0023, -0.0094, -0.0006, 0.0075, 0.0023,\n", + " nan, nan, nan],\n", + " [-0.0004, -0.0060, 0.0020, -0.0049, -0.0020, 0.0011, 0.0049, -0.0042,\n", + " 0.0354, nan, nan],\n", + " [-0.0094, -0.0002, 0.0027, -0.0060, 0.0017, -0.0005, 0.0008, -0.0014,\n", + " -0.0002, -0.1018, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 6\n", + "Attributions shape: torch.Size([1, 4])\n", + "Number of elements: 4\n", + "Attribution: tensor([[-0.0026, 0.0024, 0.0007, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 7\n", + "Attributions shape: torch.Size([5, 9])\n", + "Number of elements: 9\n", + "Attribution: tensor([[-9.2138e-03, -9.2428e-03, 7.6194e-03, -2.8685e-02, nan,\n", + " nan, nan, nan, nan],\n", + " [ 1.6982e-02, 1.1031e-02, -6.1722e-03, -4.1739e-02, -1.0123e-01,\n", + " nan, nan, nan, nan],\n", + " [-1.6126e-03, -4.8020e-04, 4.5713e-04, 6.0391e-05, 8.5584e-03,\n", + " -1.7940e-02, nan, nan, nan],\n", + " [ 3.1949e-03, -2.8566e-04, 7.6152e-03, -1.8436e-02, -1.2982e-02,\n", + " 9.0746e-03, 2.5908e-01, nan, nan],\n", + " [-1.9200e-02, 1.2491e-02, -3.2375e-03, -8.9267e-03, -2.3039e-03,\n", + " 1.1168e-03, -3.0267e-03, 6.4192e-02, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 8\n", + "Attributions shape: torch.Size([1, 5])\n", + "Number of elements: 5\n", + "Attribution: tensor([[-0.0005, 0.0114, 0.0193, -0.0177, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 9\n", + "Attributions shape: torch.Size([3, 5])\n", + "Number of elements: 5\n", + "Attribution: tensor([[-0.0237, -0.0976, nan, nan, nan],\n", + " [-0.0186, 0.0137, -0.0373, nan, nan],\n", + " [-0.0087, -0.0029, 0.0004, -0.0142, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 10\n", + "Attributions shape: torch.Size([3, 11])\n", + "Number of elements: 11\n", + "Attribution: tensor([[ 0.0047, 0.0133, 0.0140, 0.0023, -0.0094, -0.0006, 0.0075, 0.0023,\n", + " nan, nan, nan],\n", + " [-0.0004, -0.0060, 0.0020, -0.0049, -0.0020, 0.0011, 0.0049, -0.0042,\n", + " 0.0354, nan, nan],\n", + " [-0.0094, -0.0002, 0.0027, -0.0060, 0.0017, -0.0005, 0.0008, -0.0014,\n", + " -0.0002, -0.1018, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 11\n", + "Attributions shape: torch.Size([1, 4])\n", + "Number of elements: 4\n", + "Attribution: tensor([[-0.0026, 0.0024, 0.0007, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 12\n", + "Attributions shape: torch.Size([5, 9])\n", + "Number of elements: 9\n", + "Attribution: tensor([[-9.2138e-03, -9.2428e-03, 7.6194e-03, -2.8685e-02, nan,\n", + " nan, nan, nan, nan],\n", + " [ 1.6982e-02, 1.1031e-02, -6.1722e-03, -4.1739e-02, -1.0123e-01,\n", + " nan, nan, nan, nan],\n", + " [-1.6126e-03, -4.8020e-04, 4.5713e-04, 6.0391e-05, 8.5584e-03,\n", + " -1.7940e-02, nan, nan, nan],\n", + " [ 3.1949e-03, -2.8566e-04, 7.6152e-03, -1.8436e-02, -1.2982e-02,\n", + " 9.0746e-03, 2.5908e-01, nan, nan],\n", + " [-1.9200e-02, 1.2491e-02, -3.2375e-03, -8.9267e-03, -2.3039e-03,\n", + " 1.1168e-03, -3.0267e-03, 6.4192e-02, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 13\n", + "Attributions shape: torch.Size([1, 4])\n", + "Number of elements: 4\n", + "Attribution: tensor([[8.9990e-05, 1.2292e-02, 1.7112e-02, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 14\n", + "Attributions shape: torch.Size([3, 5])\n", + "Number of elements: 5\n", + "Attribution: tensor([[-0.0237, -0.0976, nan, nan, nan],\n", + " [-0.0186, 0.0137, -0.0373, nan, nan],\n", + " [-0.0087, -0.0029, 0.0004, -0.0142, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Example 15\n", + "Attributions shape: torch.Size([3, 11])\n", + "Number of elements: 11\n", + "Attribution: tensor([[ 0.0047, 0.0133, 0.0140, 0.0023, -0.0094, -0.0006, 0.0075, 0.0023,\n", + " nan, nan, nan],\n", + " [-0.0004, -0.0060, 0.0020, -0.0049, -0.0020, 0.0011, 0.0049, -0.0042,\n", + " 0.0354, nan, nan],\n", + " [-0.0094, -0.0002, 0.0027, -0.0060, 0.0017, -0.0005, 0.0008, -0.0014,\n", + " -0.0002, -0.1018, nan]])\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + } + ], + "source": [ + "import torch\n", + "\n", + "from interpreto import Occlusion\n", + "\n", + "DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", + "\n", + "list_texts = [\n", + " \"I like this\",\n", + " \"Oh it's cool\",\n", + " [\"My dog is \", \"this is very\"],\n", + " \"Interpreto is\",\n", + " \"This is two sentences. The goal is\",\n", + "]\n", + "list_targets = [\"video\", \"and I like it.\", [\"nice\", \"good\"], \"a great library\", \"to test.\"]\n", + "\n", + "list_tokenized_texts = [\n", + " tokenizer_gen(text, return_tensors=\"pt\", padding=True, truncation=True, return_offsets_mapping=True)\n", + " for text in list_texts\n", + "]\n", + "\n", + "list_tokenized_targets = [\n", + " tokenizer_gen(target, return_tensors=\"pt\", padding=True, truncation=True, return_offsets_mapping=True)\n", + " for target in list_targets\n", + "]\n", + "list_texts_complete = list_texts + list_tokenized_texts + list_texts\n", + "list_targets_complete = list_targets + list_targets + list_tokenized_targets\n", + "\n", + "explainer = Occlusion(\n", + " model_gen,\n", + " tokenizer_gen,\n", + " granularity=Granularity.WORD,\n", + " granularity_aggregation_strategy=GranularityAggregationStrategy.MEAN,\n", + ")\n", + "\n", + "i = 0\n", + "\n", + "for text, target in zip(list_texts_complete, list_targets_complete):\n", + " i += 1\n", + " print(f\"Example {i}\")\n", + " attributions = explainer.explain(text, targets=target)\n", + " print(f\"Attributions shape: {attributions[0].attributions.shape}\")\n", + " print(f\"Number of elements: {len(attributions[0].elements)}\")\n", + " print(f\"Attribution: {attributions[0].attributions}\")\n", + " # print(\"Elements:\", attributions[0].elements)\n", + "\n", + " # there is a third visualization class for generation attributions\n", + " plot_attributions(attributions[0])" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "234a4c3c", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "AttributionOutput(attributions=tensor([[-0.0407, 0.0271, -0.0071, nan, nan, nan, nan],\n", + " [-0.0039, -0.0123, -0.0051, -0.0076, nan, nan, nan],\n", + " [-0.0046, 0.0243, -0.0079, 0.0180, -0.0141, nan, nan],\n", + " [ 0.0095, 0.0375, 0.0051, 0.0175, 0.0005, -0.0044, nan]]), elements=['I', ' like', ' you', ' with', ' you', ' my', ' friend'], model_inputs_to_explain={'input_ids': tensor([[ 41, 250, 783, 407, 235, 296, 407, 235, 230, 89, 214, 337, 483]]), 'attention_mask': tensor([[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]]), 'offset_mapping': tensor([[[ 0, 1],\n", + " [ 1, 3],\n", + " [ 3, 6],\n", + " [ 6, 8],\n", + " [ 8, 10],\n", + " [10, 15],\n", + " [15, 17],\n", + " [17, 19],\n", + " [19, 21],\n", + " [21, 22],\n", + " [22, 24],\n", + " [24, 26],\n", + " [26, 29]]])}, targets=tensor([296, 407, 235, 230, 89, 214, 337, 483]), model_task=, classes=None, granularity=, granularity_aggregation_strategy=, inference_mode=)\n" + ] + }, + { + "data": { + "text/html": [ + "

Inputs

\n", + "

Outputs

\n", + "\n", + " \n", + " \n", + " " + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + } + ], + "source": [ + "from interpreto import Occlusion\n", + "\n", + "explainer = Occlusion(\n", + " model_gen,\n", + " tokenizer_gen,\n", + " granularity=Granularity.WORD,\n", + " granularity_aggregation_strategy=GranularityAggregationStrategy.MAX,\n", + ")\n", + "\n", + "text = \"I like you\"\n", + "target = \"with you my friend\"\n", + "\n", + "tokenized_texts = tokenizer_gen(text, return_tensors=\"pt\", padding=True, truncation=True, return_offsets_mapping=True)\n", + "\n", + "tokenized_targets = tokenizer_gen(\n", + " target, return_tensors=\"pt\", padding=True, truncation=True, return_offsets_mapping=True\n", + ")[\"input_ids\"]\n", + "attributions = explainer.explain(tokenized_texts, targets=tokenized_targets)\n", + "\n", + "\n", + "print(attributions[0])\n", + "plot_attributions(attributions[0])" + ] + }, + { + "cell_type": "code", + "execution_count": 13, + "id": "8cb4d137", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "torch.Size([4, 7])\n", + "7\n" + ] + } + ], + "source": [ + "print(attributions[0].attributions.shape)\n", + "print(len(attributions[0].elements))" + ] + }, + { + "cell_type": "code", + "execution_count": 14, + "id": "9d29d0cf", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "7\n", + "['I', ' like', ' you', ' with', ' you', ' my', ' friend']\n" + ] + } + ], + "source": [ + "print(len(attributions[0].elements))\n", + "print(attributions[0].elements)" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": ".venv (3.12.3)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} 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..c4796a19 --- /dev/null +++ b/interpreto/concepts/metrics/simulatability/attrsim.py @@ -0,0 +1,143 @@ +# 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 AttrSim prompt configuration.""" + + lp_samples: bool = True + lp_attributions: bool = True + anonymize_classes: bool = False + + +class PromptTypes(Enum): + """Named AttrSim prompt presets.""" + + A1_attributions_with_lp = PromptSetting(lp_samples=True, lp_attributions=True) + + +class AttrSim(AutomatedSimulatability): + """Attribution-based simulatability prompt builder. + + AttrSim mirrors ConSim but uses token-level attribution explanations instead of concept explanations. + """ + + prompt_types: type[PromptTypes] = PromptTypes + + @staticmethod + def _format_attribution_for_pred( + attribution_output: AttributionOutput, + pred_index: int, + top_k: int = 6, + ) -> str: + elements = attribution_output.elements + if isinstance(elements, torch.Tensor): + elements = [str(e.item()) for e in elements] + else: + elements = [str(e) for e in elements] + + attributions = attribution_output.attributions + if attributions.ndim == 1: + pred_attr = attributions + else: + pred_attr = attributions[pred_index] + + top_k = min(top_k, pred_attr.shape[-1]) + top_indices = torch.topk(pred_attr.abs(), 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}: {pred_attr[idx].item():+.3f}") + return "{" + ", ".join(pieces) + "}" + + def construct_prompt( + 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]]: + if isinstance(setting, PromptTypes): + setting = setting.value + + if len(interesting_samples) != len(corresponding_predictions): + raise ValueError("`interesting_samples` and `corresponding_predictions` must have the same length.") + if len(interesting_samples) != len(corresponding_labels): + raise ValueError("`interesting_samples` and `corresponding_labels` must have the same length.") + if len(interesting_samples) != len(corresponding_attribution): + raise ValueError("`interesting_samples` and `corresponding_attribution` 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.") + + classes = {i: c for i, c in enumerate(self.classes)} + 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.", + "Use the provided learning examples and attribution explanations to infer the model behavior.", + "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]) + lp_block = [ + f"Sample_{i}:", + f"\tText: {interesting_samples[i]}", + f"\tLabel: {classes[pred_index]}", + ] + if setting.lp_attributions: + lp_block.append( + f"\tAttributions: {self._format_attribution_for_pred(corresponding_attribution[i], pred_index)}" + ) + 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 diff --git a/tests/concepts/metrics/test_attrsim.py b/tests/concepts/metrics/test_attrsim.py new file mode 100644 index 00000000..17a725ea --- /dev/null +++ b/tests/concepts/metrics/test_attrsim.py @@ -0,0 +1,24 @@ +# 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 f22eea0c9ff58b8eb9ae294e040e98339810f078 Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Fri, 17 Apr 2026 09:38:09 +0000 Subject: [PATCH 04/17] add attrsim --- docs/notebooks/test_attrsim.ipynb | 206 ++++++++++++++++++++++++++++++ 1 file changed, 206 insertions(+) create mode 100644 docs/notebooks/test_attrsim.ipynb diff --git a/docs/notebooks/test_attrsim.ipynb b/docs/notebooks/test_attrsim.ipynb new file mode 100644 index 00000000..24f964b4 --- /dev/null +++ b/docs/notebooks/test_attrsim.ipynb @@ -0,0 +1,206 @@ +{ + "cells": [ + { + "cell_type": "code", + "execution_count": null, + "id": "bc397185", + "metadata": {}, + "outputs": [], + "source": [ + "import sys\n", + "\n", + "sys.path.append(\"../..\")\n", + "import torch\n", + "from datasets import load_dataset\n", + "from transformers import AutoModelForSequenceClassification, AutoTokenizer\n", + "\n", + "from interpreto.attributions import Saliency\n", + "from interpreto.concepts.metrics import AttrSim\n", + "from interpreto.model_wrapping.llm_interface import HuggingFaceLLM" + ] + }, + { + "cell_type": "code", + "execution_count": 3, + "id": "0ab1ae9f", + "metadata": {}, + "outputs": [], + "source": [ + "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", + "\n", + "model_name = \"textattack/distilbert-base-uncased-ag-news\"\n", + "tokenizer = AutoTokenizer.from_pretrained(model_name)\n", + "\n", + "model = AutoModelForSequenceClassification.from_pretrained(model_name).to(device)\n", + "dataset = load_dataset(\"fancyzhx/ag_news\")\n", + "\n", + "n_train = 700\n", + "train_inputs = dataset[\"train\"][\"text\"][:n_train]\n", + "classes_names = dataset[\"train\"].features[\"label\"].names" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "281281a0", + "metadata": {}, + "outputs": [], + "source": [ + "def predict_in_batches(model, tokenizer, texts, device, batch_size=32):\n", + " model.eval()\n", + " predictions = []\n", + "\n", + " with torch.no_grad():\n", + " for i in range(0, len(texts), batch_size):\n", + " batch_texts = texts[i : i + batch_size]\n", + "\n", + " encoded = tokenizer(batch_texts, return_tensors=\"pt\", padding=True, truncation=True)\n", + " encoded = {k: v.to(device) for k, v in encoded.items()}\n", + "\n", + " outputs = model(**encoded)\n", + " preds = outputs.logits.argmax(dim=-1).cpu()\n", + " predictions.append(preds)\n", + "\n", + " del encoded, outputs, preds\n", + " torch.cuda.empty_cache()\n", + "\n", + " return torch.cat(predictions, dim=0)\n", + "\n", + "\n", + "train_predictions = predict_in_batches(model, tokenizer, train_inputs, device, batch_size=32)" + ] + }, + { + "cell_type": "code", + "execution_count": 8, + "id": "a3d39b0d", + "metadata": {}, + "outputs": [], + "source": [ + "attrsim = AttrSim(classes=classes_names)\n", + "train_labels = torch.tensor(dataset[\"train\"][\"label\"][:n_train])\n", + "\n", + "# Select a balanced pool (good/miss) then force 10 good + 10 miss\n", + "# for the ConSim learning phase.\n", + "indices, samples, selected_labels, selected_predictions = attrsim.select_examples(\n", + " inputs=train_inputs,\n", + " labels=train_labels,\n", + " predictions=train_predictions,\n", + " nb_samples=30,\n", + " seed=0,\n", + ")\n", + "\n", + "good_mask = selected_labels == selected_predictions\n", + "good_idx = torch.where(good_mask)[0]\n", + "miss_idx = torch.where(~good_mask)[0]\n", + "\n", + "lp_good = good_idx[:10]\n", + "lp_miss = miss_idx[:10]\n", + "lp_idx = torch.cat([lp_good, lp_miss])\n", + "\n", + "all_idx = torch.arange(len(samples))\n", + "ep_idx = all_idx[~torch.isin(all_idx, lp_idx)]\n", + "ordered_idx = torch.cat([lp_idx, ep_idx])\n", + "\n", + "samples = [samples[i] for i in ordered_idx.tolist()]\n", + "selected_labels = selected_labels[ordered_idx]\n", + "selected_predictions = selected_predictions[ordered_idx]\n", + "indices = indices[ordered_idx]\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "1f0b5f0b", + "metadata": {}, + "outputs": [ + { + "ename": "IndexError", + "evalue": "index 2 is out of bounds for dimension 0 with size 1", + "output_type": "error", + "traceback": [ + "\u001b[31m---------------------------------------------------------------------------\u001b[39m", + "\u001b[31mIndexError\u001b[39m Traceback (most recent call last)", + "\u001b[36mCell\u001b[39m\u001b[36m \u001b[39m\u001b[32mIn[10]\u001b[39m\u001b[32m, line 16\u001b[39m\n\u001b[32m 10\u001b[39m \u001b[38;5;66;03m# Build attributions on the selected samples for each sample predicted class.\u001b[39;00m\n\u001b[32m 11\u001b[39m attr_outputs = saliency.explain(\n\u001b[32m 12\u001b[39m samples,\n\u001b[32m 13\u001b[39m targets=selected_predictions,\n\u001b[32m 14\u001b[39m )\n\u001b[32m---> \u001b[39m\u001b[32m16\u001b[39m attr_system_prompt, attr_user_prompts, attr_model_predictions = \u001b[43mattrsim\u001b[49m\u001b[43m.\u001b[49m\u001b[43mconstruct_prompt\u001b[49m\u001b[43m(\u001b[49m\n\u001b[32m 17\u001b[39m \u001b[43m \u001b[49m\u001b[43msetting\u001b[49m\u001b[43m=\u001b[49m\u001b[43mAttrSim\u001b[49m\u001b[43m.\u001b[49m\u001b[43mprompt_types\u001b[49m\u001b[43m.\u001b[49m\u001b[43mA1_attributions_with_lp\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 18\u001b[39m \u001b[43m \u001b[49m\u001b[43minteresting_samples\u001b[49m\u001b[43m=\u001b[49m\u001b[43msamples\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 19\u001b[39m \u001b[43m \u001b[49m\u001b[43mcorresponding_predictions\u001b[49m\u001b[43m=\u001b[49m\u001b[43mselected_predictions\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 20\u001b[39m \u001b[43m \u001b[49m\u001b[43mcorresponding_labels\u001b[49m\u001b[43m=\u001b[49m\u001b[43mselected_labels\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 21\u001b[39m \u001b[43m \u001b[49m\u001b[43mnb_learning_samples\u001b[49m\u001b[43m=\u001b[49m\u001b[32;43m20\u001b[39;49m\u001b[43m,\u001b[49m\n\u001b[32m 22\u001b[39m \u001b[43m \u001b[49m\u001b[43mcorresponding_attribution\u001b[49m\u001b[43m=\u001b[49m\u001b[43mattr_outputs\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 23\u001b[39m \u001b[43m)\u001b[49m\n\u001b[32m 25\u001b[39m attr_responses = llm.batch_generate(\n\u001b[32m 26\u001b[39m attr_system_prompt,\n\u001b[32m 27\u001b[39m attr_user_prompts,\n\u001b[32m 28\u001b[39m max_new_tokens=\u001b[32m16\u001b[39m,\n\u001b[32m 29\u001b[39m do_sample=\u001b[38;5;28;01mFalse\u001b[39;00m,\n\u001b[32m 30\u001b[39m )\n\u001b[32m 32\u001b[39m attr_score = attrsim.score_from_responses(attr_responses, attr_model_predictions)\n", + "\u001b[36mFile \u001b[39m\u001b[32m~/dev/interpreto/docs/notebooks/../../interpreto/concepts/metrics/simulatability/attrsim.py:129\u001b[39m, in \u001b[36mAttrSim.construct_prompt\u001b[39m\u001b[34m(self, setting, interesting_samples, corresponding_predictions, corresponding_labels, nb_learning_samples, corresponding_attribution)\u001b[39m\n\u001b[32m 122\u001b[39m lp_block = [\n\u001b[32m 123\u001b[39m \u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[33mSample_\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mi\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m:\u001b[39m\u001b[33m\"\u001b[39m,\n\u001b[32m 124\u001b[39m \u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[38;5;130;01m\\t\u001b[39;00m\u001b[33mText: \u001b[39m\u001b[38;5;132;01m{\u001b[39;00minteresting_samples[i]\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m\"\u001b[39m,\n\u001b[32m 125\u001b[39m \u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[38;5;130;01m\\t\u001b[39;00m\u001b[33mLabel: \u001b[39m\u001b[38;5;132;01m{\u001b[39;00mclasses[pred_index]\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m\"\u001b[39m,\n\u001b[32m 126\u001b[39m ]\n\u001b[32m 127\u001b[39m \u001b[38;5;28;01mif\u001b[39;00m setting.lp_attributions:\n\u001b[32m 128\u001b[39m lp_block.append(\n\u001b[32m--> \u001b[39m\u001b[32m129\u001b[39m \u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[38;5;130;01m\\t\u001b[39;00m\u001b[33mAttributions: \u001b[39m\u001b[38;5;132;01m{\u001b[39;00m\u001b[38;5;28;43mself\u001b[39;49m\u001b[43m.\u001b[49m\u001b[43m_format_attribution_for_pred\u001b[49m\u001b[43m(\u001b[49m\u001b[43mcorresponding_attribution\u001b[49m\u001b[43m[\u001b[49m\u001b[43mi\u001b[49m\u001b[43m]\u001b[49m\u001b[43m,\u001b[49m\u001b[38;5;250;43m \u001b[39;49m\u001b[43mpred_index\u001b[49m\u001b[43m)\u001b[49m\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m\"\u001b[39m\n\u001b[32m 130\u001b[39m )\n\u001b[32m 131\u001b[39m lp_blocks.append(\u001b[33m\"\u001b[39m\u001b[38;5;130;01m\\n\u001b[39;00m\u001b[33m\"\u001b[39m.join(lp_block))\n\u001b[32m 133\u001b[39m system_prompt_parts.append(\u001b[33m\"\u001b[39m\u001b[38;5;130;01m\\n\u001b[39;00m\u001b[33m\"\u001b[39m.join(lp_blocks))\n", + "\u001b[36mFile \u001b[39m\u001b[32m~/dev/interpreto/docs/notebooks/../../interpreto/concepts/metrics/simulatability/attrsim.py:74\u001b[39m, in \u001b[36mAttrSim._format_attribution_for_pred\u001b[39m\u001b[34m(attribution_output, pred_index, top_k)\u001b[39m\n\u001b[32m 72\u001b[39m pred_attr = attributions\n\u001b[32m 73\u001b[39m \u001b[38;5;28;01melse\u001b[39;00m:\n\u001b[32m---> \u001b[39m\u001b[32m74\u001b[39m pred_attr = \u001b[43mattributions\u001b[49m\u001b[43m[\u001b[49m\u001b[43mpred_index\u001b[49m\u001b[43m]\u001b[49m\n\u001b[32m 76\u001b[39m top_k = \u001b[38;5;28mmin\u001b[39m(top_k, pred_attr.shape[-\u001b[32m1\u001b[39m])\n\u001b[32m 77\u001b[39m top_indices = torch.topk(pred_attr.abs(), k=top_k).indices.tolist()\n", + "\u001b[31mIndexError\u001b[39m: index 2 is out of bounds for dimension 0 with size 1" + ] + } + ], + "source": [ + "# AttrSim evaluation (same selected samples, but with token attributions as explanations)\n", + "\n", + "saliency = Saliency(\n", + " model=model,\n", + " tokenizer=tokenizer,\n", + " batch_size=8,\n", + " device=device,\n", + ")\n", + "\n", + "# Build attributions on the selected samples for each sample predicted class.\n", + "attr_outputs = saliency.explain(\n", + " samples,\n", + " targets=selected_predictions,\n", + ")\n", + "\n", + "attr_system_prompt, attr_user_prompts, attr_model_predictions = attrsim.construct_prompt(\n", + " setting=AttrSim.prompt_types.A1_attributions_with_lp,\n", + " interesting_samples=samples,\n", + " corresponding_predictions=selected_predictions,\n", + " corresponding_labels=selected_labels,\n", + " nb_learning_samples=20,\n", + " corresponding_attribution=attr_outputs,\n", + ")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "14d0a0fe", + "metadata": {}, + "outputs": [], + "source": [ + "llm = HuggingFaceLLM(\n", + " model=\"HuggingFaceTB/SmolLM2-360M-Instruct\",\n", + " batch_size=2,\n", + " device=device,\n", + ")\n", + "\n", + "\n", + "attr_responses = llm.batch_generate(\n", + " attr_system_prompt,\n", + " attr_user_prompts,\n", + " max_new_tokens=16,\n", + " do_sample=False,\n", + ")\n", + "\n", + "attr_score = attrsim.score_from_responses(attr_responses, attr_model_predictions)\n", + "\n", + "print(\"AttrSim score:\", attr_score)\n", + "print(\"AttrSim responses preview:\", attr_responses[:2])" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": ".venv (3.12.3)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} From 99da6190aa717092c2b197c1d950c7ea7e946f82 Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Fri, 17 Apr 2026 09:38:22 +0000 Subject: [PATCH 05/17] add attrsim --- docs/notebooks/test.ipynb | 242 ++++++++++++++++++++++++-------------- 1 file changed, 156 insertions(+), 86 deletions(-) diff --git a/docs/notebooks/test.ipynb b/docs/notebooks/test.ipynb index babdc986..4b962831 100644 --- a/docs/notebooks/test.ipynb +++ b/docs/notebooks/test.ipynb @@ -2,7 +2,7 @@ "cells": [ { "cell_type": "code", - "execution_count": 1, + "execution_count": null, "id": "274dccbe", "metadata": {}, "outputs": [], @@ -11,7 +11,12 @@ "\n", "sys.path.append(\"../..\")\n", "import torch\n", + "from datasets import load_dataset\n", + "from transformers import AutoModelForSequenceClassification\n", "\n", + "from interpreto import ModelWithSplitPoints\n", + "from interpreto.concepts import ICAConcepts\n", + "from interpreto.concepts.interpretations import TopKInputs\n", "from interpreto.concepts.metrics import ConSim\n", "from interpreto.model_wrapping.llm_interface import HuggingFaceLLM" ] @@ -19,133 +24,198 @@ { "cell_type": "code", "execution_count": null, - "id": "35f89943", + "id": "44dfe233", + "metadata": {}, + "outputs": [], + "source": [ + "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", + "granularity = ModelWithSplitPoints.activation_granularities.CLS_TOKEN\n", + "\n", + "model_with_split_points = ModelWithSplitPoints(\n", + " model_or_repo_id=\"textattack/distilbert-base-uncased-ag-news\",\n", + " automodel=AutoModelForSequenceClassification,\n", + " split_points=[5],\n", + " device_map=device,\n", + " batch_size=64,\n", + ")\n", + "\n", + "dataset = load_dataset(\"fancyzhx/ag_news\")\n", + "\n", + "n_train = 700\n", + "train_inputs = dataset[\"train\"][\"text\"][:n_train]\n", + "classes_names = dataset[\"train\"].features[\"label\"].names\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6616b5f7", + "metadata": {}, + "outputs": [], + "source": [ + "# Compute activations\n", + "train_activations = model_with_split_points.get_activations(\n", + " inputs=train_inputs,\n", + " activation_granularity=granularity,\n", + " include_predicted_classes=True,\n", + ")\n", + "\n", + "train_predictions = train_activations.pop(\"predictions\").cpu()\n", + "\n", + "# Fit ICA concepts\n", + "concept_explainer = ICAConcepts(\n", + " model_with_split_points,\n", + " nb_concepts=16,\n", + " device=device,\n", + ")\n", + "concept_explainer.fit(train_activations)\n", + "\n", + "# Extract top-k words per concept\n", + "topk_inputs = TopKInputs(\n", + " concept_explainer=concept_explainer,\n", + " k=5,\n", + " activation_granularity=granularity,\n", + " use_unique_words=True,\n", + ")\n", + "\n", + "topk_words = topk_inputs.interpret(\n", + " inputs=train_inputs,\n", + " concepts_indices=\"all\",\n", + ")\n", + "\n", + "# Build concept interpretations\n", + "concepts_interpretation = {}\n", + "\n", + "for concept_id, words_dict in topk_words.items():\n", + " if words_dict is None:\n", + " concepts_interpretation[concept_id] = \"No strong activation pattern\"\n", + " else:\n", + " concepts_interpretation[concept_id] = \", \".join(list(words_dict.keys())[:5])\n", + "\n", + "# Compute global importances\n", + "global_gradients = concept_explainer.concept_output_gradient(\n", + " inputs=train_inputs,\n", + " targets=None,\n", + " activation_granularity=granularity,\n", + " concepts_x_gradients=True,\n", + " batch_size=32,\n", + ")\n", + "\n", + "global_importances = torch.stack(global_gradients).abs().squeeze().mean(0).cpu()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "7005d16f", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ - "Built 2 query prompt(s).\n" + "Selected 25 samples total.\n", + "Learning phase: 10 good / 10 miss\n", + "Built 5 ConSim prompt(s).\n" ] } ], "source": [ - "# Define classes\n", - "classes = [\"negative\", \"positive\"]\n", - "metric = ConSim(classes=classes)\n", - "\n", - "# Sample inputs\n", - "samples = [\n", - " \"I loved the movie.\",\n", - " \"It was boring and slow.\",\n", - " \"Excellent acting and story.\",\n", - " \"I would not recommend it.\",\n", - "]\n", - "\n", - "# Ground truth labels and model predictions\n", - "labels = torch.tensor([1, 0, 1, 0])\n", - "predictions = torch.tensor([1, 0, 1, 0])\n", - "\n", - "# Concept interpretations\n", - "concepts_interpretation = {\n", - " 0: \"positive sentiment words\",\n", - " 1: \"negative sentiment words\",\n", - "}\n", - "\n", - "# Global importances\n", - "global_importances = torch.tensor(\n", - " [\n", - " [-0.8, 0.7],\n", - " [0.8, -0.7],\n", - " ]\n", + "consim = ConSim(classes=classes_names)\n", + "train_labels = torch.tensor(dataset[\"train\"][\"label\"][:n_train])\n", + "\n", + "# Select a balanced pool (good/miss) then force 10 good + 10 miss\n", + "# for the ConSim learning phase.\n", + "indices, samples, selected_labels, selected_predictions = consim.select_examples(\n", + " inputs=train_inputs,\n", + " labels=train_labels,\n", + " predictions=train_predictions,\n", + " nb_samples=30,\n", + " seed=0,\n", + ")\n", + "\n", + "good_mask = selected_labels == selected_predictions\n", + "good_idx = torch.where(good_mask)[0]\n", + "miss_idx = torch.where(~good_mask)[0]\n", + "\n", + "lp_good = good_idx[:10]\n", + "lp_miss = miss_idx[:10]\n", + "lp_idx = torch.cat([lp_good, lp_miss])\n", + "\n", + "all_idx = torch.arange(len(samples))\n", + "ep_idx = all_idx[~torch.isin(all_idx, lp_idx)]\n", + "ordered_idx = torch.cat([lp_idx, ep_idx])\n", + "\n", + "samples = [samples[i] for i in ordered_idx.tolist()]\n", + "selected_labels = selected_labels[ordered_idx]\n", + "selected_predictions = selected_predictions[ordered_idx]\n", + "indices = indices[ordered_idx]\n", + "\n", + "local_gradients = concept_explainer.concept_output_gradient(\n", + " inputs=samples,\n", + " targets=None,\n", + " activation_granularity=granularity,\n", + " concepts_x_gradients=True,\n", + " batch_size=16,\n", ")\n", "\n", - "# Local importances (per sample)\n", - "local_importances = [\n", - " torch.tensor([[-0.1, 0.7], [0.1, -0.7]]),\n", - " torch.tensor([[0.7, -0.1], [-0.7, 0.1]]),\n", - " torch.tensor([[-0.2, 0.8], [0.2, -0.8]]),\n", - " torch.tensor([[0.8, -0.2], [-0.8, 0.2]]),\n", - "]\n", + "local_importances = [g.squeeze(1).cpu() if g.ndim == 3 else g.cpu() for g in local_gradients]\n", "\n", - "# Build prompts\n", - "system_prompt, user_prompts, model_predictions = metric.construct_prompt(\n", + "system_prompt, user_prompts, model_predictions = consim.construct_prompt(\n", " setting=ConSim.prompt_types.E3_global_and_local_concepts_with_lp,\n", " interesting_samples=samples,\n", - " corresponding_predictions=predictions,\n", - " corresponding_labels=labels,\n", - " nb_learning_samples=2,\n", + " corresponding_predictions=selected_predictions,\n", + " corresponding_labels=selected_labels,\n", + " nb_learning_samples=20,\n", " concepts_interpretation=concepts_interpretation,\n", " global_importances=global_importances,\n", " local_importances=local_importances,\n", ")\n", "\n", - "print(f\"Built {len(user_prompts)} query prompt(s).\")" + "print(f\"Selected {len(samples)} samples total.\")\n", + "print(\n", + " f\"Learning phase: \"\n", + " f\"{int((selected_labels[:20] == selected_predictions[:20]).sum())} good / \"\n", + " f\"{int((selected_labels[:20] != selected_predictions[:20]).sum())} miss\"\n", + ")\n", + "print(f\"Built {len(user_prompts)} ConSim prompt(s).\")" ] }, { "cell_type": "code", "execution_count": 5, - "id": "434c5cc8", + "id": "8f022054", "metadata": {}, "outputs": [ - { - "name": "stderr", - "output_type": "stream", - "text": [ - "The following generation flags are not valid and may be ignored: ['temperature', 'top_p', 'top_k']. Set `TRANSFORMERS_VERBOSITY=info` for more details.\n", - "A decoder-only architecture is being used, but right-padding was detected! For correct generation results, please set `padding_side='left'` when initializing the tokenizer.\n" - ] - }, { "name": "stdout", "output_type": "stream", "text": [ - "Raw responses:\n", - "---\n", - ".public!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!\n", - "---\n", - ".public!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!+x!\n" + "ConSim score: 0.4\n", + "Responses preview: ['This evaluation sample is from the Nikkei, a Japanese stock exchange. The', 'of 3 minutes 13.17 seconds.\\n\\tLabel: \\nassistant\\nThe evaluation sample provided is a sports event report. The text mentions that No Gold']\n" ] } ], "source": [ - "# HuggingFaceLLM now supports both chat-template and plain CausalLM tokenizers.\n", - "# Choose any local/cached CausalLM checkpoint.\n", - "\n", - "llm = HuggingFaceLLM(model=\"Qwen/Qwen3-0.6B\", batch_size=2, device=\"auto\")\n", + "# Prefer an instruction-tuned model for better ConSim performance\n", + "llm = HuggingFaceLLM(\n", + " model=\"HuggingFaceTB/SmolLM2-360M-Instruct\",\n", + " batch_size=2,\n", + " device=device,\n", + ")\n", "\n", "responses = llm.batch_generate(\n", " system_prompt,\n", " user_prompts,\n", - " max_new_tokens=40,\n", + " max_new_tokens=16,\n", " do_sample=False,\n", ")\n", "\n", - "print(\"Raw responses:\")\n", - "for response in responses:\n", - " print(\"---\")\n", - " print(response)" - ] - }, - { - "cell_type": "code", - "execution_count": 4, - "id": "d950fc53", - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "ConSim score: 0.0\n" - ] - } - ], - "source": [ - "score = metric.score_from_responses(responses, model_predictions)\n", - "print(\"ConSim score:\", score)" + "# Compute score\n", + "score = consim.score_from_responses(responses, model_predictions)\n", + "\n", + "print(\"ConSim score:\", score)\n", + "print(\"Responses preview:\", responses[:2])" ] } ], From 66da2483619f608bb4292afa6abe6be6042ff24c Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Fri, 17 Apr 2026 14:33:56 +0000 Subject: [PATCH 06/17] add attrsim --- .../metrics/simulatability/attrsim.py | 114 +++++++++++++++++- 1 file changed, 108 insertions(+), 6 deletions(-) diff --git a/interpreto/concepts/metrics/simulatability/attrsim.py b/interpreto/concepts/metrics/simulatability/attrsim.py index c4796a19..c8b820bd 100644 --- a/interpreto/concepts/metrics/simulatability/attrsim.py +++ b/interpreto/concepts/metrics/simulatability/attrsim.py @@ -34,23 +34,123 @@ class PromptSetting(NamedTuple): - """Low-level AttrSim prompt configuration.""" + """ + 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 = True lp_attributions: bool = True + lp_contrastive_attributions: bool = False anonymize_classes: bool = False + 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()." + ) + class PromptTypes(Enum): - """Named AttrSim prompt presets.""" + """ + Named AttrSim prompt presets. - A1_attributions_with_lp = PromptSetting(lp_samples=True, lp_attributions=True) + 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. + """ -class AttrSim(AutomatedSimulatability): - """Attribution-based simulatability prompt builder. + L1_baseline_without_lp = PromptSetting() + L2_baseline_with_lp = PromptSetting(lp_samples=True) + + E1_attribution_with_lp = PromptSetting(lp_samples=True, lp_attributions=True) - AttrSim mirrors ConSim but uses token-level attribution explanations instead of concept explanations. + 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 @@ -128,6 +228,8 @@ def construct_prompt( lp_block.append( f"\tAttributions: {self._format_attribution_for_pred(corresponding_attribution[i], pred_index)}" ) + if setting.lp_contrastive_attributions: + raise NotImplementedError("Contrastive attribution formatting is not implemented yet.") lp_blocks.append("\n".join(lp_block)) system_prompt_parts.append("\n".join(lp_blocks)) From cf04e47d75f51392b987584ef215068dcc917d69 Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Fri, 17 Apr 2026 14:37:40 +0000 Subject: [PATCH 07/17] add attrsim --- docs/notebooks/test_attrsim.ipynb | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/notebooks/test_attrsim.ipynb b/docs/notebooks/test_attrsim.ipynb index 24f964b4..7fbcf0e7 100644 --- a/docs/notebooks/test_attrsim.ipynb +++ b/docs/notebooks/test_attrsim.ipynb @@ -145,7 +145,7 @@ ")\n", "\n", "attr_system_prompt, attr_user_prompts, attr_model_predictions = attrsim.construct_prompt(\n", - " setting=AttrSim.prompt_types.A1_attributions_with_lp,\n", + " setting=AttrSim.prompt_types.E1_attribution_with_lp,\n", " interesting_samples=samples,\n", " corresponding_predictions=selected_predictions,\n", " corresponding_labels=selected_labels,\n", From b2baa46fdedac66979632cd60dafccb8bd14830a Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Fri, 17 Apr 2026 16:14:13 +0000 Subject: [PATCH 08/17] add attrsim --- docs/notebooks/test_attrsim.ipynb | 62 ++++++++++++------- .../metrics/simulatability/attrsim.py | 60 ++++++++++++++---- 2 files changed, 87 insertions(+), 35 deletions(-) diff --git a/docs/notebooks/test_attrsim.ipynb b/docs/notebooks/test_attrsim.ipynb index 7fbcf0e7..b6959531 100644 --- a/docs/notebooks/test_attrsim.ipynb +++ b/docs/notebooks/test_attrsim.ipynb @@ -2,7 +2,7 @@ "cells": [ { "cell_type": "code", - "execution_count": null, + "execution_count": 7, "id": "bc397185", "metadata": {}, "outputs": [], @@ -21,7 +21,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": 8, "id": "0ab1ae9f", "metadata": {}, "outputs": [], @@ -41,7 +41,7 @@ }, { "cell_type": "code", - "execution_count": null, + "execution_count": 9, "id": "281281a0", "metadata": {}, "outputs": [], @@ -72,7 +72,7 @@ }, { "cell_type": "code", - "execution_count": 8, + "execution_count": 10, "id": "a3d39b0d", "metadata": {}, "outputs": [], @@ -105,29 +105,15 @@ "samples = [samples[i] for i in ordered_idx.tolist()]\n", "selected_labels = selected_labels[ordered_idx]\n", "selected_predictions = selected_predictions[ordered_idx]\n", - "indices = indices[ordered_idx]\n" + "indices = indices[ordered_idx]" ] }, { "cell_type": "code", - "execution_count": null, + "execution_count": 11, "id": "1f0b5f0b", "metadata": {}, - "outputs": [ - { - "ename": "IndexError", - "evalue": "index 2 is out of bounds for dimension 0 with size 1", - "output_type": "error", - "traceback": [ - "\u001b[31m---------------------------------------------------------------------------\u001b[39m", - "\u001b[31mIndexError\u001b[39m Traceback (most recent call last)", - "\u001b[36mCell\u001b[39m\u001b[36m \u001b[39m\u001b[32mIn[10]\u001b[39m\u001b[32m, line 16\u001b[39m\n\u001b[32m 10\u001b[39m \u001b[38;5;66;03m# Build attributions on the selected samples for each sample predicted class.\u001b[39;00m\n\u001b[32m 11\u001b[39m attr_outputs = saliency.explain(\n\u001b[32m 12\u001b[39m samples,\n\u001b[32m 13\u001b[39m targets=selected_predictions,\n\u001b[32m 14\u001b[39m )\n\u001b[32m---> \u001b[39m\u001b[32m16\u001b[39m attr_system_prompt, attr_user_prompts, attr_model_predictions = \u001b[43mattrsim\u001b[49m\u001b[43m.\u001b[49m\u001b[43mconstruct_prompt\u001b[49m\u001b[43m(\u001b[49m\n\u001b[32m 17\u001b[39m \u001b[43m \u001b[49m\u001b[43msetting\u001b[49m\u001b[43m=\u001b[49m\u001b[43mAttrSim\u001b[49m\u001b[43m.\u001b[49m\u001b[43mprompt_types\u001b[49m\u001b[43m.\u001b[49m\u001b[43mA1_attributions_with_lp\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 18\u001b[39m \u001b[43m \u001b[49m\u001b[43minteresting_samples\u001b[49m\u001b[43m=\u001b[49m\u001b[43msamples\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 19\u001b[39m \u001b[43m \u001b[49m\u001b[43mcorresponding_predictions\u001b[49m\u001b[43m=\u001b[49m\u001b[43mselected_predictions\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 20\u001b[39m \u001b[43m \u001b[49m\u001b[43mcorresponding_labels\u001b[49m\u001b[43m=\u001b[49m\u001b[43mselected_labels\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 21\u001b[39m \u001b[43m \u001b[49m\u001b[43mnb_learning_samples\u001b[49m\u001b[43m=\u001b[49m\u001b[32;43m20\u001b[39;49m\u001b[43m,\u001b[49m\n\u001b[32m 22\u001b[39m \u001b[43m \u001b[49m\u001b[43mcorresponding_attribution\u001b[49m\u001b[43m=\u001b[49m\u001b[43mattr_outputs\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 23\u001b[39m \u001b[43m)\u001b[49m\n\u001b[32m 25\u001b[39m attr_responses = llm.batch_generate(\n\u001b[32m 26\u001b[39m attr_system_prompt,\n\u001b[32m 27\u001b[39m attr_user_prompts,\n\u001b[32m 28\u001b[39m max_new_tokens=\u001b[32m16\u001b[39m,\n\u001b[32m 29\u001b[39m do_sample=\u001b[38;5;28;01mFalse\u001b[39;00m,\n\u001b[32m 30\u001b[39m )\n\u001b[32m 32\u001b[39m attr_score = attrsim.score_from_responses(attr_responses, attr_model_predictions)\n", - "\u001b[36mFile \u001b[39m\u001b[32m~/dev/interpreto/docs/notebooks/../../interpreto/concepts/metrics/simulatability/attrsim.py:129\u001b[39m, in \u001b[36mAttrSim.construct_prompt\u001b[39m\u001b[34m(self, setting, interesting_samples, corresponding_predictions, corresponding_labels, nb_learning_samples, corresponding_attribution)\u001b[39m\n\u001b[32m 122\u001b[39m lp_block = [\n\u001b[32m 123\u001b[39m \u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[33mSample_\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mi\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m:\u001b[39m\u001b[33m\"\u001b[39m,\n\u001b[32m 124\u001b[39m \u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[38;5;130;01m\\t\u001b[39;00m\u001b[33mText: \u001b[39m\u001b[38;5;132;01m{\u001b[39;00minteresting_samples[i]\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m\"\u001b[39m,\n\u001b[32m 125\u001b[39m \u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[38;5;130;01m\\t\u001b[39;00m\u001b[33mLabel: \u001b[39m\u001b[38;5;132;01m{\u001b[39;00mclasses[pred_index]\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m\"\u001b[39m,\n\u001b[32m 126\u001b[39m ]\n\u001b[32m 127\u001b[39m \u001b[38;5;28;01mif\u001b[39;00m setting.lp_attributions:\n\u001b[32m 128\u001b[39m lp_block.append(\n\u001b[32m--> \u001b[39m\u001b[32m129\u001b[39m \u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[38;5;130;01m\\t\u001b[39;00m\u001b[33mAttributions: \u001b[39m\u001b[38;5;132;01m{\u001b[39;00m\u001b[38;5;28;43mself\u001b[39;49m\u001b[43m.\u001b[49m\u001b[43m_format_attribution_for_pred\u001b[49m\u001b[43m(\u001b[49m\u001b[43mcorresponding_attribution\u001b[49m\u001b[43m[\u001b[49m\u001b[43mi\u001b[49m\u001b[43m]\u001b[49m\u001b[43m,\u001b[49m\u001b[38;5;250;43m \u001b[39;49m\u001b[43mpred_index\u001b[49m\u001b[43m)\u001b[49m\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m\"\u001b[39m\n\u001b[32m 130\u001b[39m )\n\u001b[32m 131\u001b[39m lp_blocks.append(\u001b[33m\"\u001b[39m\u001b[38;5;130;01m\\n\u001b[39;00m\u001b[33m\"\u001b[39m.join(lp_block))\n\u001b[32m 133\u001b[39m system_prompt_parts.append(\u001b[33m\"\u001b[39m\u001b[38;5;130;01m\\n\u001b[39;00m\u001b[33m\"\u001b[39m.join(lp_blocks))\n", - "\u001b[36mFile \u001b[39m\u001b[32m~/dev/interpreto/docs/notebooks/../../interpreto/concepts/metrics/simulatability/attrsim.py:74\u001b[39m, in \u001b[36mAttrSim._format_attribution_for_pred\u001b[39m\u001b[34m(attribution_output, pred_index, top_k)\u001b[39m\n\u001b[32m 72\u001b[39m pred_attr = attributions\n\u001b[32m 73\u001b[39m \u001b[38;5;28;01melse\u001b[39;00m:\n\u001b[32m---> \u001b[39m\u001b[32m74\u001b[39m pred_attr = \u001b[43mattributions\u001b[49m\u001b[43m[\u001b[49m\u001b[43mpred_index\u001b[49m\u001b[43m]\u001b[49m\n\u001b[32m 76\u001b[39m top_k = \u001b[38;5;28mmin\u001b[39m(top_k, pred_attr.shape[-\u001b[32m1\u001b[39m])\n\u001b[32m 77\u001b[39m top_indices = torch.topk(pred_attr.abs(), k=top_k).indices.tolist()\n", - "\u001b[31mIndexError\u001b[39m: index 2 is out of bounds for dimension 0 with size 1" - ] - } - ], + "outputs": [], "source": [ "# AttrSim evaluation (same selected samples, but with token attributions as explanations)\n", "\n", @@ -156,10 +142,19 @@ }, { "cell_type": "code", - "execution_count": null, + "execution_count": 12, "id": "14d0a0fe", "metadata": {}, - "outputs": [], + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "AttrSim score: 0.2\n", + "AttrSim responses preview: ['The evaluation sample provided is a news article about the Nikkei stock index falling', 'of 3 minutes 13.17 seconds.\\n\\tLabel: \\nassistant\\nThe evaluation sample is a sentence from a news article about the 200']\n" + ] + } + ], "source": [ "llm = HuggingFaceLLM(\n", " model=\"HuggingFaceTB/SmolLM2-360M-Instruct\",\n", @@ -180,6 +175,27 @@ "print(\"AttrSim score:\", attr_score)\n", "print(\"AttrSim responses preview:\", attr_responses[:2])" ] + }, + { + "cell_type": "code", + "execution_count": 15, + "id": "a3f75f2b", + "metadata": {}, + "outputs": [ + { + "data": { + "text/plain": [ + "'You are a classifier. Predict the class for each evaluation sample.\\n\\nUse the provided learning examples and attribution explanations to infer the model behavior.\\n\\nOnly return the class name, no additional text.\\n\\nThe classes are: [World, Sports, Business, Sci/Tech]\\n\\nSample_0:\\n\\tText: Frail Pope Ends Tiring Lourdes Pilgrimage LOURDES, France (Reuters) - Pope John Paul, a sick man among the sick, wound up a emotional visit to this miracle shrine Sunday and struggled with iron determination to finish a sermon in order to encourage others suffering around him.\\n\\tLabel: World\\n\\tAttributions: {miracle: +0.005, lourdes: +0.004, lourdes: +0.003, pilgrimage: +0.003, shrine: +0.003, pope: +0.002}\\nSample_1:\\n\\tText: Two visions of Iraq struggle to take hold Fighting in Najaf threatened to undermine a conference to choose a national assembly.\\n\\tLabel: World\\n\\tAttributions: {iraq: +0.003, struggle: +0.001, assembly: +0.001, visions: +0.001, conference: +0.001, undermine: +0.001}\\nSample_2:\\n\\tText: Phish farewell attracts thousands \"Jam band\" Phish play their last gigs together at a special festival in the US which has attracted thousands of fans.\\n\\tLabel: World\\n\\tAttributions: {phish: +0.003, farewell: +0.003, phish: +0.002, gigs: +0.002, jam: +0.002, attracts: +0.002}\\nSample_3:\\n\\tText: Oldsmobile: The final parking lot Why General Motors dropped the Oldsmobile. The four brand paradoxes GM had to face - the name, the product, image re-positioning, and the consumer - all added up to a brand that had little hope of rebranding.\\n\\tLabel: Business\\n\\tAttributions: {gm: +0.003, oldsmobile: +0.002, parking: +0.002, oldsmobile: +0.002, motors: +0.002, lot: +0.001}\\nSample_4:\\n\\tText: Phelps, Rival Thorpe in 200M-Free Semis ATHENS, Greece - Michael Phelps took care of qualifying for the Olympic 200-meter freestyle semifinals Sunday, and then found out he had been added to the American team for the evening\\'s 400 freestyle relay final. Phelps\\' rivals Ian Thorpe and Pieter van den Hoogenband and teammate Klete Keller were faster than the teenager in the 200 free preliminaries...\\n\\tLabel: World\\n\\tAttributions: {greece: +0.002, -: +0.001, athens: +0.001, michael: +0.001, .: +0.001, phelps: +0.001}\\nSample_5:\\n\\tText: Shell \\'could be target for Total\\' Oil giant Shell could be bracing itself for a takeover attempt, possibly from French rival Total, a press report claims.\\n\\tLabel: Business\\n\\tAttributions: {takeover: +0.003, \\': +0.003, oil: +0.002, total: +0.002, target: +0.002, shell: +0.002}\\nSample_6:\\n\\tText: Oracle Sales Data Seen Being Released (Reuters) Reuters - Oracle Corp. sales documents\\\\detailing highly confidential information, such as which\\\\companies receive discounts on Oracle\\'s business software\\\\products and the size of the discounts, are likely to be made\\\\public, a federal judge said on Friday.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {oracle: +0.003, oracle: +0.003, corp: +0.003, reuters: +0.002, oracle: +0.002, \\\\: +0.002}\\nSample_7:\\n\\tText: Apple to open second Japanese retail store this month (MacCentral) MacCentral - Apple Computer Inc. will open its second Japanese retail store later this month in the western Japanese city of Osaka, it said Thursday.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {osaka: +0.001, apple: +0.001, apple: +0.001, maccentral: +0.001, japanese: +0.001, japanese: +0.001}\\nSample_8:\\n\\tText: Election-Year Rate Hike Puzzles Some WASHINGTON - Going against conventional wisdom, the Federal Reserve is raising interest rates in an election year. And it is Fed Chairman Alan Greenspan, a Republican, who is leading the charge even though an incumbent Republican in the White House is facing voter unrest about the state of the economy...\\n\\tLabel: World\\n\\tAttributions: {puzzles: +0.002, washington: +0.001, economy: +0.001, reserve: +0.001, going: +0.001, wisdom: +0.001}\\nSample_9:\\n\\tText: India Rethinks Plan for Manned Moon Mission By S. SRINIVASAN BANGALORE, India (AP) -- India is rethinking its plan to send a man to the moon by 2015, as the mission would cost a lot of money and yield very little in return, the national space agency said Thursday...\\n\\tLabel: Sci/Tech\\n\\tAttributions: {ap: +0.001, bangalore: +0.001, manned: +0.001, space: +0.001, india: +0.000, moon: +0.000}\\nSample_10:\\n\\tText: Delightful Dell The company\\'s results show that it\\'s not grim all over tech world. Just all of it that isn\\'t Dell.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {dell: +0.008, dell: +0.008, delightful: +0.006, tech: +0.005, grim: +0.003, company: +0.002}\\nSample_11:\\n\\tText: Antitrust Lawyer Takes Helm at FTC As Deborah P. Majoras takes over the Federal Trade Commission on Monday, she\\'s expected to build on the broad agenda set by her predecessor, Timothy J. Muris.\\n\\tLabel: Business\\n\\tAttributions: {ftc: +0.005, antitrust: +0.005, trade: +0.003, lawyer: +0.002, helm: +0.002, at: +0.002}\\nSample_12:\\n\\tText: Technology company sues five ex-employees A Marlborough-based technology company is suing five former employees, including three senior managers, for allegedly conspiring against their employer while working on opening a competing business.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {marlborough: +0.005, technology: +0.004, technology: +0.003, managers: +0.002, employer: +0.002, employees: +0.001}\\nSample_13:\\n\\tText: More Big Boobs in Playboy An interview with Google\\'s co-founders due out in the current issue of Playboy may delay the company\\'s IPO. Securities regulations restrict what executives can say while preparing to sell stock for the first time.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {google: +0.007, playboy: +0.006, playboy: +0.005, securities: +0.003, boobs: +0.003, interview: +0.002}\\nSample_14:\\n\\tText: Japan nuclear firm shuts plants The company running the Japanese nuclear plant hit by a fatal accident is to close its reactors for safety checks.\\n\\tLabel: World\\n\\tAttributions: {nuclear: +0.007, nuclear: +0.004, reactors: +0.004, japanese: +0.003, plant: +0.003, japan: +0.003}\\nSample_15:\\n\\tText: Autodesk tackles project collaboration Autodesk this week unwrapped an updated version of its hosted project collaboration service targeted at the construction and manufacturing industries. Autodesk Buzzsaw lets multiple, dispersed project participants -- including building owners, developers, architects, construction teams, and facility managers -- share and manage data throughout the life of a project, according to Autodesk officials.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {developers: +0.001, tackles: +0.001, collaboration: +0.001, collaboration: +0.001, autodesk: +0.001, project: +0.001}\\nSample_16:\\n\\tText: Barrel of Monkeys, 2004 Edition: Notes on Philippine Elections Well, it\\'s election time in the Republic of the Philippines, and that means the monkeys are rolling around in those political barrels, having as much fun as they can while laughing their heads off at the strange goings-on that characterize a democratic process loosely based on the American model but that de facto looks more like a Fellini movie crossed with a Tom and Jerry cartoon - column includes a useful election-year glossary!\\n\\tLabel: World\\n\\tAttributions: {philippine: +0.010, elections: +0.006, philippines: +0.005, notes: +0.003, edition: +0.003, barrel: +0.003}\\nSample_17:\\n\\tText: Fark Sells Out. France Surrenders Blogs are the hottest thing on the Net, but are they messing with traditional publishing principles? One of the most popular, Fark.com, is allegedly selling links. Is it the wave of the future? By Daniel Terdiman.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {blogs: +0.012, sells: +0.006, france: +0.005, publishing: +0.005, links: +0.004, net: +0.004}\\nSample_18:\\n\\tText: IT Myth 5: Most IT projects fail Do most IT projects fail? Some point to the number of giant consultancies such as IBM Global Services, Capgemini, and Sapient, who feed off bad experiences encountered by enterprises. Sapient is a company founded on the realization that IT projects are not successful, says Sapient CTO Ben Gaucherin.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {ibm: +0.003, myth: +0.002, enterprises: +0.002, it: +0.001, 5: +0.001, projects: +0.001}\\nSample_19:\\n\\tText: Eye on Athens, China stresses a \\'frugal\\' 2008 Olympics Amid a reevaluation, officials this week pushed the completion date for venues back to 2007.\\n\\tLabel: Sports\\n\\tAttributions: {olympics: +0.007, athens: +0.004, venues: +0.003, stresses: +0.003, officials: +0.003, amid: +0.002}'" + ] + }, + "execution_count": 15, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "attr_system_prompt" + ] } ], "metadata": { diff --git a/interpreto/concepts/metrics/simulatability/attrsim.py b/interpreto/concepts/metrics/simulatability/attrsim.py index c8b820bd..ab98161d 100644 --- a/interpreto/concepts/metrics/simulatability/attrsim.py +++ b/interpreto/concepts/metrics/simulatability/attrsim.py @@ -155,6 +155,10 @@ class AttrSim(AutomatedSimulatability): prompt_types: type[PromptTypes] = PromptTypes + @staticmethod + def _resolve_prompt_setting(setting: PromptTypes | PromptSetting) -> PromptSetting: + return setting.value if isinstance(setting, PromptTypes) else setting + @staticmethod def _format_attribution_for_pred( attribution_output: AttributionOutput, @@ -170,6 +174,8 @@ def _format_attribution_for_pred( attributions = attribution_output.attributions if attributions.ndim == 1: pred_attr = attributions + elif attributions.ndim == 2 and attributions.shape[0] == 1: + pred_attr = attributions[0] else: pred_attr = attributions[pred_index] @@ -192,19 +198,19 @@ def construct_prompt( *, corresponding_attribution: list[AttributionOutput], ) -> tuple[str, list[str], list[str]]: - if isinstance(setting, PromptTypes): - setting = setting.value + 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 len(interesting_samples) != len(corresponding_predictions): - raise ValueError("`interesting_samples` and `corresponding_predictions` must have the same length.") - if len(interesting_samples) != len(corresponding_labels): - raise ValueError("`interesting_samples` and `corresponding_labels` must have the same length.") - if len(interesting_samples) != len(corresponding_attribution): - raise ValueError("`interesting_samples` and `corresponding_attribution` 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.") - - classes = {i: c for i, c in enumerate(self.classes)} if setting.anonymize_classes: classes = {i: f"Class_{i}" for i in classes.keys()} @@ -243,3 +249,33 @@ def construct_prompt( 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: + 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.") From 86104f5029e4416f94ac3a5de108975b549ac871 Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Mon, 20 Apr 2026 10:17:11 +0000 Subject: [PATCH 09/17] add attrsim tests --- .../metrics/simulatability/attrsim.py | 4 +- tests/concepts/metrics/test_attrsim.py | 104 ++++++++++++++++++ 2 files changed, 106 insertions(+), 2 deletions(-) diff --git a/interpreto/concepts/metrics/simulatability/attrsim.py b/interpreto/concepts/metrics/simulatability/attrsim.py index ab98161d..f6210ea6 100644 --- a/interpreto/concepts/metrics/simulatability/attrsim.py +++ b/interpreto/concepts/metrics/simulatability/attrsim.py @@ -54,8 +54,8 @@ class PromptSetting(NamedTuple): Preventing the LLM from using knowledge on classes names. """ - lp_samples: bool = True - lp_attributions: bool = True + lp_samples: bool = False + lp_attributions: bool = False lp_contrastive_attributions: bool = False anonymize_classes: bool = False diff --git a/tests/concepts/metrics/test_attrsim.py b/tests/concepts/metrics/test_attrsim.py index 17a725ea..2a10cfc7 100644 --- a/tests/concepts/metrics/test_attrsim.py +++ b/tests/concepts/metrics/test_attrsim.py @@ -22,3 +22,107 @@ # 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_attribution_output( + attributions: torch.Tensor, + *, + elements: list[str] | None = None, +) -> AttributionOutput: + if elements is None: + elements = [f"t{i}" for i in range(attributions.shape[-1])] + + 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=torch.tensor([0, 1]), + ) # type: ignore + + +def test_prompt_settings_defaults_and_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 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.L2_baseline_with_lp.value.lp_attributions is False + assert AttrSim.prompt_types.E1_attribution_with_lp.value.lp_samples is True + assert AttrSim.prompt_types.E1_attribution_with_lp.value.lp_attributions is True + + +def test_format_attribution_for_pred_handles_singleton_class_axis(): + attribution = _build_attribution_output(torch.tensor([[0.2, -0.5, 0.1]])) + + rendered = AttrSim._format_attribution_for_pred(attribution_output=attribution, pred_index=2, top_k=2) + + assert "t1: -0.500" in rendered + assert "t0: +0.200" in rendered + + +def test_construct_prompt_with_enum_setting(): + 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_attribution_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_validates_lengths(): + metric = AttrSim(classes=["A", "B"]) + samples = ["s0", "s1", "s2"] + predictions = torch.tensor([0, 1, 0]) + labels = torch.tensor([0, 1, 0]) + attributions = [_build_attribution_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, + ) + + +def test_construct_prompt_rejects_too_many_learning_samples(): + metric = AttrSim(classes=["A", "B"]) + samples = ["s0", "s1"] + predictions = torch.tensor([0, 1]) + labels = torch.tensor([0, 1]) + attributions = [_build_attribution_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=2, + corresponding_attribution=attributions, + ) From d72cb6df0c7b8175d1114dd04f5907475b8697d1 Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Mon, 20 Apr 2026 14:00:25 +0000 Subject: [PATCH 10/17] add attrsim --- .../metrics/simulatability/attrsim.py | 90 +++++++++++++++---- tests/concepts/metrics/test_attrsim.py | 51 ++++++++++- 2 files changed, 123 insertions(+), 18 deletions(-) diff --git a/interpreto/concepts/metrics/simulatability/attrsim.py b/interpreto/concepts/metrics/simulatability/attrsim.py index f6210ea6..b0f09e37 100644 --- a/interpreto/concepts/metrics/simulatability/attrsim.py +++ b/interpreto/concepts/metrics/simulatability/attrsim.py @@ -160,35 +160,53 @@ def _resolve_prompt_setting(setting: PromptTypes | PromptSetting) -> PromptSetti return setting.value if isinstance(setting, PromptTypes) else setting @staticmethod - def _format_attribution_for_pred( + def _get_attr_vector( attribution_output: AttributionOutput, - pred_index: int, + class_index: int, + ) -> torch.Tensor: + 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, ) -> str: - elements = attribution_output.elements if isinstance(elements, torch.Tensor): elements = [str(e.item()) for e in elements] else: elements = [str(e) for e in elements] - attributions = attribution_output.attributions - if attributions.ndim == 1: - pred_attr = attributions - elif attributions.ndim == 2 and attributions.shape[0] == 1: - pred_attr = attributions[0] - else: - pred_attr = attributions[pred_index] - - top_k = min(top_k, pred_attr.shape[-1]) - top_indices = torch.topk(pred_attr.abs(), k=top_k).indices.tolist() + top_k = min(top_k, attr_vector.shape[-1]) + top_indices = torch.topk(attr_vector.abs(), 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}: {pred_attr[idx].item():+.3f}") + pieces.append(f"{token}: {attr_vector[idx].item():+.3f}") return "{" + ", ".join(pieces) + "}" - def construct_prompt( + @staticmethod + def _format_attribution_for_pred( + attribution_output: AttributionOutput, + pred_index: int, + top_k: int = 6, + ) -> str: + pred_attr = AttrSim._get_attr_vector(attribution_output, pred_index) + return AttrSim._format_attr_vector(attribution_output.elements, pred_attr, top_k=top_k) + + def construct_prompt( # type: ignore self, setting: PromptTypes | PromptSetting, interesting_samples: list[str], @@ -232,10 +250,26 @@ def construct_prompt( ] if setting.lp_attributions: lp_block.append( - f"\tAttributions: {self._format_attribution_for_pred(corresponding_attribution[i], pred_index)}" + f"\tAttributions for {classes[pred_index]}: {self._format_attribution_for_pred(corresponding_attribution[i], pred_index)}" ) if setting.lp_contrastive_attributions: - raise NotImplementedError("Contrastive attribution formatting is not implemented yet.") + 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 + + lp_block.append( + f"\t{text}: {self._format_attr_vector(corresponding_attribution[i].elements, attr_to_show)}" + ) + lp_blocks.append("\n".join(lp_block)) system_prompt_parts.append("\n".join(lp_blocks)) @@ -279,3 +313,25 @@ def _check_input_settings_correspondence( 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/tests/concepts/metrics/test_attrsim.py b/tests/concepts/metrics/test_attrsim.py index 2a10cfc7..f73fa41b 100644 --- a/tests/concepts/metrics/test_attrsim.py +++ b/tests/concepts/metrics/test_attrsim.py @@ -87,7 +87,7 @@ def test_construct_prompt_with_enum_setting(): corresponding_attribution=attributions, ) - assert "Attributions:" in system_prompt + assert "Attributions for A:" in system_prompt assert len(user_prompts) == 2 assert model_predictions == ["A", "B"] @@ -126,3 +126,52 @@ def test_construct_prompt_rejects_too_many_learning_samples(): nb_learning_samples=2, corresponding_attribution=attributions, ) + + +def test_construct_prompt_with_contrastive_attributions(): + 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 and 2 are misclassified + attributions = [ + _build_attribution_output(torch.tensor([[0.4, -0.1, 0.2], [0.1, 0.2, -0.3]])), + _build_attribution_output(torch.tensor([[0.2, -0.5, 0.1], [0.6, -0.1, -0.2]])), + _build_attribution_output(torch.tensor([[0.1, -0.3, 0.5], [-0.2, 0.7, 0.1]])), + _build_attribution_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 "Contrastive Attributions supporting B rather than A" in system_prompt + assert "Attributions for A" in system_prompt + assert len(user_prompts) == 2 + assert model_predictions == ["A", "B"] + + +def test_construct_prompt_contrastive_requires_classwise_attribution_for_miss(): + metric = AttrSim(classes=["A", "B"]) + samples = ["s0", "s1", "s2"] + predictions = torch.tensor([0, 1, 0]) + labels = torch.tensor([0, 0, 1]) # one miss in LP when nb_learning_samples=2 + attributions = [ + _build_attribution_output(torch.tensor([[0.1, -0.2, 0.3]])), + _build_attribution_output(torch.tensor([[0.2, -0.5, 0.1]])), # singleton axis, not classwise + _build_attribution_output(torch.tensor([[0.4, -0.1, 0.2]])), + ] + + with pytest.raises(ValueError, match="require 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, + ) From e66778696ce72f804616254a53f65599c69e5c3e Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Mon, 20 Apr 2026 14:20:46 +0000 Subject: [PATCH 11/17] add attrsim --- docs/api/attributions/metrics/attrsim.md | 3 +++ mkdocs.yml | 1 + 2 files changed, 4 insertions(+) create mode 100644 docs/api/attributions/metrics/attrsim.md diff --git a/docs/api/attributions/metrics/attrsim.md b/docs/api/attributions/metrics/attrsim.md new file mode 100644 index 00000000..8b7789fb --- /dev/null +++ b/docs/api/attributions/metrics/attrsim.md @@ -0,0 +1,3 @@ +# Simulatability Metric for Attribution Methods + +::: interpreto.concepts.metrics.simulatability.attrsim.AttrSim 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 From f034eed9077a0475f8b26c790d25047f9b9d8ccb Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Mon, 20 Apr 2026 14:40:27 +0000 Subject: [PATCH 12/17] add attrsim --- README.md | 3 +++ docs/api/attributions/metrics/attrsim.md | 2 +- docs/index.md | 2 ++ 3 files changed, 6 insertions(+), 1 deletion(-) 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 index 8b7789fb..d6b25360 100644 --- a/docs/api/attributions/metrics/attrsim.md +++ b/docs/api/attributions/metrics/attrsim.md @@ -1,3 +1,3 @@ -# Simulatability Metric for Attribution Methods +# 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/). From 2304954f591eabbc53b50adf6add1f9418358bac Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Mon, 20 Apr 2026 14:54:32 +0000 Subject: [PATCH 13/17] add attrsim --- docs/notebooks/test.ipynb | 243 ------------------------------ docs/notebooks/test_attrsim.ipynb | 222 --------------------------- 2 files changed, 465 deletions(-) delete mode 100644 docs/notebooks/test.ipynb delete mode 100644 docs/notebooks/test_attrsim.ipynb diff --git a/docs/notebooks/test.ipynb b/docs/notebooks/test.ipynb deleted file mode 100644 index 4b962831..00000000 --- a/docs/notebooks/test.ipynb +++ /dev/null @@ -1,243 +0,0 @@ -{ - "cells": [ - { - "cell_type": "code", - "execution_count": null, - "id": "274dccbe", - "metadata": {}, - "outputs": [], - "source": [ - "import sys\n", - "\n", - "sys.path.append(\"../..\")\n", - "import torch\n", - "from datasets import load_dataset\n", - "from transformers import AutoModelForSequenceClassification\n", - "\n", - "from interpreto import ModelWithSplitPoints\n", - "from interpreto.concepts import ICAConcepts\n", - "from interpreto.concepts.interpretations import TopKInputs\n", - "from interpreto.concepts.metrics import ConSim\n", - "from interpreto.model_wrapping.llm_interface import HuggingFaceLLM" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "44dfe233", - "metadata": {}, - "outputs": [], - "source": [ - "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", - "granularity = ModelWithSplitPoints.activation_granularities.CLS_TOKEN\n", - "\n", - "model_with_split_points = ModelWithSplitPoints(\n", - " model_or_repo_id=\"textattack/distilbert-base-uncased-ag-news\",\n", - " automodel=AutoModelForSequenceClassification,\n", - " split_points=[5],\n", - " device_map=device,\n", - " batch_size=64,\n", - ")\n", - "\n", - "dataset = load_dataset(\"fancyzhx/ag_news\")\n", - "\n", - "n_train = 700\n", - "train_inputs = dataset[\"train\"][\"text\"][:n_train]\n", - "classes_names = dataset[\"train\"].features[\"label\"].names\n" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "6616b5f7", - "metadata": {}, - "outputs": [], - "source": [ - "# Compute activations\n", - "train_activations = model_with_split_points.get_activations(\n", - " inputs=train_inputs,\n", - " activation_granularity=granularity,\n", - " include_predicted_classes=True,\n", - ")\n", - "\n", - "train_predictions = train_activations.pop(\"predictions\").cpu()\n", - "\n", - "# Fit ICA concepts\n", - "concept_explainer = ICAConcepts(\n", - " model_with_split_points,\n", - " nb_concepts=16,\n", - " device=device,\n", - ")\n", - "concept_explainer.fit(train_activations)\n", - "\n", - "# Extract top-k words per concept\n", - "topk_inputs = TopKInputs(\n", - " concept_explainer=concept_explainer,\n", - " k=5,\n", - " activation_granularity=granularity,\n", - " use_unique_words=True,\n", - ")\n", - "\n", - "topk_words = topk_inputs.interpret(\n", - " inputs=train_inputs,\n", - " concepts_indices=\"all\",\n", - ")\n", - "\n", - "# Build concept interpretations\n", - "concepts_interpretation = {}\n", - "\n", - "for concept_id, words_dict in topk_words.items():\n", - " if words_dict is None:\n", - " concepts_interpretation[concept_id] = \"No strong activation pattern\"\n", - " else:\n", - " concepts_interpretation[concept_id] = \", \".join(list(words_dict.keys())[:5])\n", - "\n", - "# Compute global importances\n", - "global_gradients = concept_explainer.concept_output_gradient(\n", - " inputs=train_inputs,\n", - " targets=None,\n", - " activation_granularity=granularity,\n", - " concepts_x_gradients=True,\n", - " batch_size=32,\n", - ")\n", - "\n", - "global_importances = torch.stack(global_gradients).abs().squeeze().mean(0).cpu()" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "7005d16f", - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Selected 25 samples total.\n", - "Learning phase: 10 good / 10 miss\n", - "Built 5 ConSim prompt(s).\n" - ] - } - ], - "source": [ - "consim = ConSim(classes=classes_names)\n", - "train_labels = torch.tensor(dataset[\"train\"][\"label\"][:n_train])\n", - "\n", - "# Select a balanced pool (good/miss) then force 10 good + 10 miss\n", - "# for the ConSim learning phase.\n", - "indices, samples, selected_labels, selected_predictions = consim.select_examples(\n", - " inputs=train_inputs,\n", - " labels=train_labels,\n", - " predictions=train_predictions,\n", - " nb_samples=30,\n", - " seed=0,\n", - ")\n", - "\n", - "good_mask = selected_labels == selected_predictions\n", - "good_idx = torch.where(good_mask)[0]\n", - "miss_idx = torch.where(~good_mask)[0]\n", - "\n", - "lp_good = good_idx[:10]\n", - "lp_miss = miss_idx[:10]\n", - "lp_idx = torch.cat([lp_good, lp_miss])\n", - "\n", - "all_idx = torch.arange(len(samples))\n", - "ep_idx = all_idx[~torch.isin(all_idx, lp_idx)]\n", - "ordered_idx = torch.cat([lp_idx, ep_idx])\n", - "\n", - "samples = [samples[i] for i in ordered_idx.tolist()]\n", - "selected_labels = selected_labels[ordered_idx]\n", - "selected_predictions = selected_predictions[ordered_idx]\n", - "indices = indices[ordered_idx]\n", - "\n", - "local_gradients = concept_explainer.concept_output_gradient(\n", - " inputs=samples,\n", - " targets=None,\n", - " activation_granularity=granularity,\n", - " concepts_x_gradients=True,\n", - " batch_size=16,\n", - ")\n", - "\n", - "local_importances = [g.squeeze(1).cpu() if g.ndim == 3 else g.cpu() for g in local_gradients]\n", - "\n", - "system_prompt, user_prompts, model_predictions = consim.construct_prompt(\n", - " setting=ConSim.prompt_types.E3_global_and_local_concepts_with_lp,\n", - " interesting_samples=samples,\n", - " corresponding_predictions=selected_predictions,\n", - " corresponding_labels=selected_labels,\n", - " nb_learning_samples=20,\n", - " concepts_interpretation=concepts_interpretation,\n", - " global_importances=global_importances,\n", - " local_importances=local_importances,\n", - ")\n", - "\n", - "print(f\"Selected {len(samples)} samples total.\")\n", - "print(\n", - " f\"Learning phase: \"\n", - " f\"{int((selected_labels[:20] == selected_predictions[:20]).sum())} good / \"\n", - " f\"{int((selected_labels[:20] != selected_predictions[:20]).sum())} miss\"\n", - ")\n", - "print(f\"Built {len(user_prompts)} ConSim prompt(s).\")" - ] - }, - { - "cell_type": "code", - "execution_count": 5, - "id": "8f022054", - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "ConSim score: 0.4\n", - "Responses preview: ['This evaluation sample is from the Nikkei, a Japanese stock exchange. The', 'of 3 minutes 13.17 seconds.\\n\\tLabel: \\nassistant\\nThe evaluation sample provided is a sports event report. The text mentions that No Gold']\n" - ] - } - ], - "source": [ - "# Prefer an instruction-tuned model for better ConSim performance\n", - "llm = HuggingFaceLLM(\n", - " model=\"HuggingFaceTB/SmolLM2-360M-Instruct\",\n", - " batch_size=2,\n", - " device=device,\n", - ")\n", - "\n", - "responses = llm.batch_generate(\n", - " system_prompt,\n", - " user_prompts,\n", - " max_new_tokens=16,\n", - " do_sample=False,\n", - ")\n", - "\n", - "# Compute score\n", - "score = consim.score_from_responses(responses, model_predictions)\n", - "\n", - "print(\"ConSim score:\", score)\n", - "print(\"Responses preview:\", responses[:2])" - ] - } - ], - "metadata": { - "kernelspec": { - "display_name": ".venv (3.12.3)", - "language": "python", - "name": "python3" - }, - "language_info": { - "codemirror_mode": { - "name": "ipython", - "version": 3 - }, - "file_extension": ".py", - "mimetype": "text/x-python", - "name": "python", - "nbconvert_exporter": "python", - "pygments_lexer": "ipython3", - "version": "3.12.3" - } - }, - "nbformat": 4, - "nbformat_minor": 5 -} diff --git a/docs/notebooks/test_attrsim.ipynb b/docs/notebooks/test_attrsim.ipynb deleted file mode 100644 index b6959531..00000000 --- a/docs/notebooks/test_attrsim.ipynb +++ /dev/null @@ -1,222 +0,0 @@ -{ - "cells": [ - { - "cell_type": "code", - "execution_count": 7, - "id": "bc397185", - "metadata": {}, - "outputs": [], - "source": [ - "import sys\n", - "\n", - "sys.path.append(\"../..\")\n", - "import torch\n", - "from datasets import load_dataset\n", - "from transformers import AutoModelForSequenceClassification, AutoTokenizer\n", - "\n", - "from interpreto.attributions import Saliency\n", - "from interpreto.concepts.metrics import AttrSim\n", - "from interpreto.model_wrapping.llm_interface import HuggingFaceLLM" - ] - }, - { - "cell_type": "code", - "execution_count": 8, - "id": "0ab1ae9f", - "metadata": {}, - "outputs": [], - "source": [ - "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", - "\n", - "model_name = \"textattack/distilbert-base-uncased-ag-news\"\n", - "tokenizer = AutoTokenizer.from_pretrained(model_name)\n", - "\n", - "model = AutoModelForSequenceClassification.from_pretrained(model_name).to(device)\n", - "dataset = load_dataset(\"fancyzhx/ag_news\")\n", - "\n", - "n_train = 700\n", - "train_inputs = dataset[\"train\"][\"text\"][:n_train]\n", - "classes_names = dataset[\"train\"].features[\"label\"].names" - ] - }, - { - "cell_type": "code", - "execution_count": 9, - "id": "281281a0", - "metadata": {}, - "outputs": [], - "source": [ - "def predict_in_batches(model, tokenizer, texts, device, batch_size=32):\n", - " model.eval()\n", - " predictions = []\n", - "\n", - " with torch.no_grad():\n", - " for i in range(0, len(texts), batch_size):\n", - " batch_texts = texts[i : i + batch_size]\n", - "\n", - " encoded = tokenizer(batch_texts, return_tensors=\"pt\", padding=True, truncation=True)\n", - " encoded = {k: v.to(device) for k, v in encoded.items()}\n", - "\n", - " outputs = model(**encoded)\n", - " preds = outputs.logits.argmax(dim=-1).cpu()\n", - " predictions.append(preds)\n", - "\n", - " del encoded, outputs, preds\n", - " torch.cuda.empty_cache()\n", - "\n", - " return torch.cat(predictions, dim=0)\n", - "\n", - "\n", - "train_predictions = predict_in_batches(model, tokenizer, train_inputs, device, batch_size=32)" - ] - }, - { - "cell_type": "code", - "execution_count": 10, - "id": "a3d39b0d", - "metadata": {}, - "outputs": [], - "source": [ - "attrsim = AttrSim(classes=classes_names)\n", - "train_labels = torch.tensor(dataset[\"train\"][\"label\"][:n_train])\n", - "\n", - "# Select a balanced pool (good/miss) then force 10 good + 10 miss\n", - "# for the ConSim learning phase.\n", - "indices, samples, selected_labels, selected_predictions = attrsim.select_examples(\n", - " inputs=train_inputs,\n", - " labels=train_labels,\n", - " predictions=train_predictions,\n", - " nb_samples=30,\n", - " seed=0,\n", - ")\n", - "\n", - "good_mask = selected_labels == selected_predictions\n", - "good_idx = torch.where(good_mask)[0]\n", - "miss_idx = torch.where(~good_mask)[0]\n", - "\n", - "lp_good = good_idx[:10]\n", - "lp_miss = miss_idx[:10]\n", - "lp_idx = torch.cat([lp_good, lp_miss])\n", - "\n", - "all_idx = torch.arange(len(samples))\n", - "ep_idx = all_idx[~torch.isin(all_idx, lp_idx)]\n", - "ordered_idx = torch.cat([lp_idx, ep_idx])\n", - "\n", - "samples = [samples[i] for i in ordered_idx.tolist()]\n", - "selected_labels = selected_labels[ordered_idx]\n", - "selected_predictions = selected_predictions[ordered_idx]\n", - "indices = indices[ordered_idx]" - ] - }, - { - "cell_type": "code", - "execution_count": 11, - "id": "1f0b5f0b", - "metadata": {}, - "outputs": [], - "source": [ - "# AttrSim evaluation (same selected samples, but with token attributions as explanations)\n", - "\n", - "saliency = Saliency(\n", - " model=model,\n", - " tokenizer=tokenizer,\n", - " batch_size=8,\n", - " device=device,\n", - ")\n", - "\n", - "# Build attributions on the selected samples for each sample predicted class.\n", - "attr_outputs = saliency.explain(\n", - " samples,\n", - " targets=selected_predictions,\n", - ")\n", - "\n", - "attr_system_prompt, attr_user_prompts, attr_model_predictions = attrsim.construct_prompt(\n", - " setting=AttrSim.prompt_types.E1_attribution_with_lp,\n", - " interesting_samples=samples,\n", - " corresponding_predictions=selected_predictions,\n", - " corresponding_labels=selected_labels,\n", - " nb_learning_samples=20,\n", - " corresponding_attribution=attr_outputs,\n", - ")" - ] - }, - { - "cell_type": "code", - "execution_count": 12, - "id": "14d0a0fe", - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "AttrSim score: 0.2\n", - "AttrSim responses preview: ['The evaluation sample provided is a news article about the Nikkei stock index falling', 'of 3 minutes 13.17 seconds.\\n\\tLabel: \\nassistant\\nThe evaluation sample is a sentence from a news article about the 200']\n" - ] - } - ], - "source": [ - "llm = HuggingFaceLLM(\n", - " model=\"HuggingFaceTB/SmolLM2-360M-Instruct\",\n", - " batch_size=2,\n", - " device=device,\n", - ")\n", - "\n", - "\n", - "attr_responses = llm.batch_generate(\n", - " attr_system_prompt,\n", - " attr_user_prompts,\n", - " max_new_tokens=16,\n", - " do_sample=False,\n", - ")\n", - "\n", - "attr_score = attrsim.score_from_responses(attr_responses, attr_model_predictions)\n", - "\n", - "print(\"AttrSim score:\", attr_score)\n", - "print(\"AttrSim responses preview:\", attr_responses[:2])" - ] - }, - { - "cell_type": "code", - "execution_count": 15, - "id": "a3f75f2b", - "metadata": {}, - "outputs": [ - { - "data": { - "text/plain": [ - "'You are a classifier. Predict the class for each evaluation sample.\\n\\nUse the provided learning examples and attribution explanations to infer the model behavior.\\n\\nOnly return the class name, no additional text.\\n\\nThe classes are: [World, Sports, Business, Sci/Tech]\\n\\nSample_0:\\n\\tText: Frail Pope Ends Tiring Lourdes Pilgrimage LOURDES, France (Reuters) - Pope John Paul, a sick man among the sick, wound up a emotional visit to this miracle shrine Sunday and struggled with iron determination to finish a sermon in order to encourage others suffering around him.\\n\\tLabel: World\\n\\tAttributions: {miracle: +0.005, lourdes: +0.004, lourdes: +0.003, pilgrimage: +0.003, shrine: +0.003, pope: +0.002}\\nSample_1:\\n\\tText: Two visions of Iraq struggle to take hold Fighting in Najaf threatened to undermine a conference to choose a national assembly.\\n\\tLabel: World\\n\\tAttributions: {iraq: +0.003, struggle: +0.001, assembly: +0.001, visions: +0.001, conference: +0.001, undermine: +0.001}\\nSample_2:\\n\\tText: Phish farewell attracts thousands \"Jam band\" Phish play their last gigs together at a special festival in the US which has attracted thousands of fans.\\n\\tLabel: World\\n\\tAttributions: {phish: +0.003, farewell: +0.003, phish: +0.002, gigs: +0.002, jam: +0.002, attracts: +0.002}\\nSample_3:\\n\\tText: Oldsmobile: The final parking lot Why General Motors dropped the Oldsmobile. The four brand paradoxes GM had to face - the name, the product, image re-positioning, and the consumer - all added up to a brand that had little hope of rebranding.\\n\\tLabel: Business\\n\\tAttributions: {gm: +0.003, oldsmobile: +0.002, parking: +0.002, oldsmobile: +0.002, motors: +0.002, lot: +0.001}\\nSample_4:\\n\\tText: Phelps, Rival Thorpe in 200M-Free Semis ATHENS, Greece - Michael Phelps took care of qualifying for the Olympic 200-meter freestyle semifinals Sunday, and then found out he had been added to the American team for the evening\\'s 400 freestyle relay final. Phelps\\' rivals Ian Thorpe and Pieter van den Hoogenband and teammate Klete Keller were faster than the teenager in the 200 free preliminaries...\\n\\tLabel: World\\n\\tAttributions: {greece: +0.002, -: +0.001, athens: +0.001, michael: +0.001, .: +0.001, phelps: +0.001}\\nSample_5:\\n\\tText: Shell \\'could be target for Total\\' Oil giant Shell could be bracing itself for a takeover attempt, possibly from French rival Total, a press report claims.\\n\\tLabel: Business\\n\\tAttributions: {takeover: +0.003, \\': +0.003, oil: +0.002, total: +0.002, target: +0.002, shell: +0.002}\\nSample_6:\\n\\tText: Oracle Sales Data Seen Being Released (Reuters) Reuters - Oracle Corp. sales documents\\\\detailing highly confidential information, such as which\\\\companies receive discounts on Oracle\\'s business software\\\\products and the size of the discounts, are likely to be made\\\\public, a federal judge said on Friday.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {oracle: +0.003, oracle: +0.003, corp: +0.003, reuters: +0.002, oracle: +0.002, \\\\: +0.002}\\nSample_7:\\n\\tText: Apple to open second Japanese retail store this month (MacCentral) MacCentral - Apple Computer Inc. will open its second Japanese retail store later this month in the western Japanese city of Osaka, it said Thursday.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {osaka: +0.001, apple: +0.001, apple: +0.001, maccentral: +0.001, japanese: +0.001, japanese: +0.001}\\nSample_8:\\n\\tText: Election-Year Rate Hike Puzzles Some WASHINGTON - Going against conventional wisdom, the Federal Reserve is raising interest rates in an election year. And it is Fed Chairman Alan Greenspan, a Republican, who is leading the charge even though an incumbent Republican in the White House is facing voter unrest about the state of the economy...\\n\\tLabel: World\\n\\tAttributions: {puzzles: +0.002, washington: +0.001, economy: +0.001, reserve: +0.001, going: +0.001, wisdom: +0.001}\\nSample_9:\\n\\tText: India Rethinks Plan for Manned Moon Mission By S. SRINIVASAN BANGALORE, India (AP) -- India is rethinking its plan to send a man to the moon by 2015, as the mission would cost a lot of money and yield very little in return, the national space agency said Thursday...\\n\\tLabel: Sci/Tech\\n\\tAttributions: {ap: +0.001, bangalore: +0.001, manned: +0.001, space: +0.001, india: +0.000, moon: +0.000}\\nSample_10:\\n\\tText: Delightful Dell The company\\'s results show that it\\'s not grim all over tech world. Just all of it that isn\\'t Dell.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {dell: +0.008, dell: +0.008, delightful: +0.006, tech: +0.005, grim: +0.003, company: +0.002}\\nSample_11:\\n\\tText: Antitrust Lawyer Takes Helm at FTC As Deborah P. Majoras takes over the Federal Trade Commission on Monday, she\\'s expected to build on the broad agenda set by her predecessor, Timothy J. Muris.\\n\\tLabel: Business\\n\\tAttributions: {ftc: +0.005, antitrust: +0.005, trade: +0.003, lawyer: +0.002, helm: +0.002, at: +0.002}\\nSample_12:\\n\\tText: Technology company sues five ex-employees A Marlborough-based technology company is suing five former employees, including three senior managers, for allegedly conspiring against their employer while working on opening a competing business.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {marlborough: +0.005, technology: +0.004, technology: +0.003, managers: +0.002, employer: +0.002, employees: +0.001}\\nSample_13:\\n\\tText: More Big Boobs in Playboy An interview with Google\\'s co-founders due out in the current issue of Playboy may delay the company\\'s IPO. Securities regulations restrict what executives can say while preparing to sell stock for the first time.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {google: +0.007, playboy: +0.006, playboy: +0.005, securities: +0.003, boobs: +0.003, interview: +0.002}\\nSample_14:\\n\\tText: Japan nuclear firm shuts plants The company running the Japanese nuclear plant hit by a fatal accident is to close its reactors for safety checks.\\n\\tLabel: World\\n\\tAttributions: {nuclear: +0.007, nuclear: +0.004, reactors: +0.004, japanese: +0.003, plant: +0.003, japan: +0.003}\\nSample_15:\\n\\tText: Autodesk tackles project collaboration Autodesk this week unwrapped an updated version of its hosted project collaboration service targeted at the construction and manufacturing industries. Autodesk Buzzsaw lets multiple, dispersed project participants -- including building owners, developers, architects, construction teams, and facility managers -- share and manage data throughout the life of a project, according to Autodesk officials.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {developers: +0.001, tackles: +0.001, collaboration: +0.001, collaboration: +0.001, autodesk: +0.001, project: +0.001}\\nSample_16:\\n\\tText: Barrel of Monkeys, 2004 Edition: Notes on Philippine Elections Well, it\\'s election time in the Republic of the Philippines, and that means the monkeys are rolling around in those political barrels, having as much fun as they can while laughing their heads off at the strange goings-on that characterize a democratic process loosely based on the American model but that de facto looks more like a Fellini movie crossed with a Tom and Jerry cartoon - column includes a useful election-year glossary!\\n\\tLabel: World\\n\\tAttributions: {philippine: +0.010, elections: +0.006, philippines: +0.005, notes: +0.003, edition: +0.003, barrel: +0.003}\\nSample_17:\\n\\tText: Fark Sells Out. France Surrenders Blogs are the hottest thing on the Net, but are they messing with traditional publishing principles? One of the most popular, Fark.com, is allegedly selling links. Is it the wave of the future? By Daniel Terdiman.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {blogs: +0.012, sells: +0.006, france: +0.005, publishing: +0.005, links: +0.004, net: +0.004}\\nSample_18:\\n\\tText: IT Myth 5: Most IT projects fail Do most IT projects fail? Some point to the number of giant consultancies such as IBM Global Services, Capgemini, and Sapient, who feed off bad experiences encountered by enterprises. Sapient is a company founded on the realization that IT projects are not successful, says Sapient CTO Ben Gaucherin.\\n\\tLabel: Sci/Tech\\n\\tAttributions: {ibm: +0.003, myth: +0.002, enterprises: +0.002, it: +0.001, 5: +0.001, projects: +0.001}\\nSample_19:\\n\\tText: Eye on Athens, China stresses a \\'frugal\\' 2008 Olympics Amid a reevaluation, officials this week pushed the completion date for venues back to 2007.\\n\\tLabel: Sports\\n\\tAttributions: {olympics: +0.007, athens: +0.004, venues: +0.003, stresses: +0.003, officials: +0.003, amid: +0.002}'" - ] - }, - "execution_count": 15, - "metadata": {}, - "output_type": "execute_result" - } - ], - "source": [ - "attr_system_prompt" - ] - } - ], - "metadata": { - "kernelspec": { - "display_name": ".venv (3.12.3)", - "language": "python", - "name": "python3" - }, - "language_info": { - "codemirror_mode": { - "name": "ipython", - "version": 3 - }, - "file_extension": ".py", - "mimetype": "text/x-python", - "name": "python", - "nbconvert_exporter": "python", - "pygments_lexer": "ipython3", - "version": "3.12.3" - } - }, - "nbformat": 4, - "nbformat_minor": 5 -} From 0a275b63e50e6a88cf224babe1e9af59b89ad522 Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Mon, 20 Apr 2026 14:55:22 +0000 Subject: [PATCH 14/17] add attrsim --- docs/notebooks/test_sentences.ipynb | 41548 -------------------------- 1 file changed, 41548 deletions(-) delete mode 100644 docs/notebooks/test_sentences.ipynb diff --git a/docs/notebooks/test_sentences.ipynb b/docs/notebooks/test_sentences.ipynb deleted file mode 100644 index abe03e5e..00000000 --- a/docs/notebooks/test_sentences.ipynb +++ /dev/null @@ -1,41548 +0,0 @@ -{ - "cells": [ - { - "cell_type": "code", - "execution_count": 8, - "id": "7eb999b1", - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "The autoreload extension is already loaded. To reload it, use:\n", - " %reload_ext autoreload\n" - ] - } - ], - "source": [ - "%load_ext autoreload\n", - "%autoreload 2\n", - "\n", - "import sys\n", - "\n", - "sys.path.append(\"../..\")\n", - "\n", - "from transformers import AutoModelForCausalLM, AutoTokenizer\n", - "\n", - "from interpreto import (\n", - " Granularity,\n", - " plot_attributions,\n", - ")\n", - "from interpreto.commons import GranularityAggregationStrategy" - ] - }, - { - "cell_type": "markdown", - "id": "ae0b758d", - "metadata": {}, - "source": [ - "Modèles testés qui marchent:\n", - "- gpt2\n", - "- \n", - "\n", - "\n", - "Modèles testés qui ne marchent pas:\n", - "- Qwen/Qwen3.5-0.8B\n", - "- mistralai/Mistral-7B-v0.1" - ] - }, - { - "cell_type": "code", - "execution_count": 9, - "id": "ff3d10c5", - "metadata": {}, - "outputs": [ - { - "data": { - "application/vnd.jupyter.widget-view+json": { - "model_id": "13d73f90100e428fbe60f339b08a0dd3", - "version_major": 2, - "version_minor": 0 - }, - "text/plain": [ - "Loading weights: 0%| | 0/64 [00:00

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 2\n", - "Attributions shape: torch.Size([5, 9])\n", - "Number of elements: 9\n", - "Attribution: tensor([[-9.2138e-03, -9.2428e-03, 7.6194e-03, -2.8685e-02, nan,\n", - " nan, nan, nan, nan],\n", - " [ 1.6982e-02, 1.1031e-02, -6.1722e-03, -4.1739e-02, -1.0123e-01,\n", - " nan, nan, nan, nan],\n", - " [-1.6126e-03, -4.8020e-04, 4.5713e-04, 6.0391e-05, 8.5584e-03,\n", - " -1.7940e-02, nan, nan, nan],\n", - " [ 3.1949e-03, -2.8566e-04, 7.6152e-03, -1.8436e-02, -1.2982e-02,\n", - " 9.0746e-03, 2.5908e-01, nan, nan],\n", - " [-1.9200e-02, 1.2491e-02, -3.2375e-03, -8.9267e-03, -2.3039e-03,\n", - " 1.1168e-03, -3.0267e-03, 6.4192e-02, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 3\n", - "Attributions shape: torch.Size([1, 4])\n", - "Number of elements: 4\n", - "Attribution: tensor([[8.9990e-05, 1.2292e-02, 1.7112e-02, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 4\n", - "Attributions shape: torch.Size([3, 5])\n", - "Number of elements: 5\n", - "Attribution: tensor([[-0.0237, -0.0976, nan, nan, nan],\n", - " [-0.0186, 0.0137, -0.0373, nan, nan],\n", - " [-0.0087, -0.0029, 0.0004, -0.0142, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 5\n", - "Attributions shape: torch.Size([3, 11])\n", - "Number of elements: 11\n", - "Attribution: tensor([[ 0.0047, 0.0133, 0.0140, 0.0023, -0.0094, -0.0006, 0.0075, 0.0023,\n", - " nan, nan, nan],\n", - " [-0.0004, -0.0060, 0.0020, -0.0049, -0.0020, 0.0011, 0.0049, -0.0042,\n", - " 0.0354, nan, nan],\n", - " [-0.0094, -0.0002, 0.0027, -0.0060, 0.0017, -0.0005, 0.0008, -0.0014,\n", - " -0.0002, -0.1018, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 6\n", - "Attributions shape: torch.Size([1, 4])\n", - "Number of elements: 4\n", - "Attribution: tensor([[-0.0026, 0.0024, 0.0007, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 7\n", - "Attributions shape: torch.Size([5, 9])\n", - "Number of elements: 9\n", - "Attribution: tensor([[-9.2138e-03, -9.2428e-03, 7.6194e-03, -2.8685e-02, nan,\n", - " nan, nan, nan, nan],\n", - " [ 1.6982e-02, 1.1031e-02, -6.1722e-03, -4.1739e-02, -1.0123e-01,\n", - " nan, nan, nan, nan],\n", - " [-1.6126e-03, -4.8020e-04, 4.5713e-04, 6.0391e-05, 8.5584e-03,\n", - " -1.7940e-02, nan, nan, nan],\n", - " [ 3.1949e-03, -2.8566e-04, 7.6152e-03, -1.8436e-02, -1.2982e-02,\n", - " 9.0746e-03, 2.5908e-01, nan, nan],\n", - " [-1.9200e-02, 1.2491e-02, -3.2375e-03, -8.9267e-03, -2.3039e-03,\n", - " 1.1168e-03, -3.0267e-03, 6.4192e-02, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 8\n", - "Attributions shape: torch.Size([1, 5])\n", - "Number of elements: 5\n", - "Attribution: tensor([[-0.0005, 0.0114, 0.0193, -0.0177, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 9\n", - "Attributions shape: torch.Size([3, 5])\n", - "Number of elements: 5\n", - "Attribution: tensor([[-0.0237, -0.0976, nan, nan, nan],\n", - " [-0.0186, 0.0137, -0.0373, nan, nan],\n", - " [-0.0087, -0.0029, 0.0004, -0.0142, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 10\n", - "Attributions shape: torch.Size([3, 11])\n", - "Number of elements: 11\n", - "Attribution: tensor([[ 0.0047, 0.0133, 0.0140, 0.0023, -0.0094, -0.0006, 0.0075, 0.0023,\n", - " nan, nan, nan],\n", - " [-0.0004, -0.0060, 0.0020, -0.0049, -0.0020, 0.0011, 0.0049, -0.0042,\n", - " 0.0354, nan, nan],\n", - " [-0.0094, -0.0002, 0.0027, -0.0060, 0.0017, -0.0005, 0.0008, -0.0014,\n", - " -0.0002, -0.1018, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 11\n", - "Attributions shape: torch.Size([1, 4])\n", - "Number of elements: 4\n", - "Attribution: tensor([[-0.0026, 0.0024, 0.0007, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 12\n", - "Attributions shape: torch.Size([5, 9])\n", - "Number of elements: 9\n", - "Attribution: tensor([[-9.2138e-03, -9.2428e-03, 7.6194e-03, -2.8685e-02, nan,\n", - " nan, nan, nan, nan],\n", - " [ 1.6982e-02, 1.1031e-02, -6.1722e-03, -4.1739e-02, -1.0123e-01,\n", - " nan, nan, nan, nan],\n", - " [-1.6126e-03, -4.8020e-04, 4.5713e-04, 6.0391e-05, 8.5584e-03,\n", - " -1.7940e-02, nan, nan, nan],\n", - " [ 3.1949e-03, -2.8566e-04, 7.6152e-03, -1.8436e-02, -1.2982e-02,\n", - " 9.0746e-03, 2.5908e-01, nan, nan],\n", - " [-1.9200e-02, 1.2491e-02, -3.2375e-03, -8.9267e-03, -2.3039e-03,\n", - " 1.1168e-03, -3.0267e-03, 6.4192e-02, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 13\n", - "Attributions shape: torch.Size([1, 4])\n", - "Number of elements: 4\n", - "Attribution: tensor([[8.9990e-05, 1.2292e-02, 1.7112e-02, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 14\n", - "Attributions shape: torch.Size([3, 5])\n", - "Number of elements: 5\n", - "Attribution: tensor([[-0.0237, -0.0976, nan, nan, nan],\n", - " [-0.0186, 0.0137, -0.0373, nan, nan],\n", - " [-0.0087, -0.0029, 0.0004, -0.0142, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Example 15\n", - "Attributions shape: torch.Size([3, 11])\n", - "Number of elements: 11\n", - "Attribution: tensor([[ 0.0047, 0.0133, 0.0140, 0.0023, -0.0094, -0.0006, 0.0075, 0.0023,\n", - " nan, nan, nan],\n", - " [-0.0004, -0.0060, 0.0020, -0.0049, -0.0020, 0.0011, 0.0049, -0.0042,\n", - " 0.0354, nan, nan],\n", - " [-0.0094, -0.0002, 0.0027, -0.0060, 0.0017, -0.0005, 0.0008, -0.0014,\n", - " -0.0002, -0.1018, nan]])\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - } - ], - "source": [ - "import torch\n", - "\n", - "from interpreto import Occlusion\n", - "\n", - "DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", - "\n", - "list_texts = [\n", - " \"I like this\",\n", - " \"Oh it's cool\",\n", - " [\"My dog is \", \"this is very\"],\n", - " \"Interpreto is\",\n", - " \"This is two sentences. The goal is\",\n", - "]\n", - "list_targets = [\"video\", \"and I like it.\", [\"nice\", \"good\"], \"a great library\", \"to test.\"]\n", - "\n", - "list_tokenized_texts = [\n", - " tokenizer_gen(text, return_tensors=\"pt\", padding=True, truncation=True, return_offsets_mapping=True)\n", - " for text in list_texts\n", - "]\n", - "\n", - "list_tokenized_targets = [\n", - " tokenizer_gen(target, return_tensors=\"pt\", padding=True, truncation=True, return_offsets_mapping=True)\n", - " for target in list_targets\n", - "]\n", - "list_texts_complete = list_texts + list_tokenized_texts + list_texts\n", - "list_targets_complete = list_targets + list_targets + list_tokenized_targets\n", - "\n", - "explainer = Occlusion(\n", - " model_gen,\n", - " tokenizer_gen,\n", - " granularity=Granularity.WORD,\n", - " granularity_aggregation_strategy=GranularityAggregationStrategy.MEAN,\n", - ")\n", - "\n", - "i = 0\n", - "\n", - "for text, target in zip(list_texts_complete, list_targets_complete):\n", - " i += 1\n", - " print(f\"Example {i}\")\n", - " attributions = explainer.explain(text, targets=target)\n", - " print(f\"Attributions shape: {attributions[0].attributions.shape}\")\n", - " print(f\"Number of elements: {len(attributions[0].elements)}\")\n", - " print(f\"Attribution: {attributions[0].attributions}\")\n", - " # print(\"Elements:\", attributions[0].elements)\n", - "\n", - " # there is a third visualization class for generation attributions\n", - " plot_attributions(attributions[0])" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "234a4c3c", - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "AttributionOutput(attributions=tensor([[-0.0407, 0.0271, -0.0071, nan, nan, nan, nan],\n", - " [-0.0039, -0.0123, -0.0051, -0.0076, nan, nan, nan],\n", - " [-0.0046, 0.0243, -0.0079, 0.0180, -0.0141, nan, nan],\n", - " [ 0.0095, 0.0375, 0.0051, 0.0175, 0.0005, -0.0044, nan]]), elements=['I', ' like', ' you', ' with', ' you', ' my', ' friend'], model_inputs_to_explain={'input_ids': tensor([[ 41, 250, 783, 407, 235, 296, 407, 235, 230, 89, 214, 337, 483]]), 'attention_mask': tensor([[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]]), 'offset_mapping': tensor([[[ 0, 1],\n", - " [ 1, 3],\n", - " [ 3, 6],\n", - " [ 6, 8],\n", - " [ 8, 10],\n", - " [10, 15],\n", - " [15, 17],\n", - " [17, 19],\n", - " [19, 21],\n", - " [21, 22],\n", - " [22, 24],\n", - " [24, 26],\n", - " [26, 29]]])}, targets=tensor([296, 407, 235, 230, 89, 214, 337, 483]), model_task=, classes=None, granularity=, granularity_aggregation_strategy=, inference_mode=)\n" - ] - }, - { - "data": { - "text/html": [ - "

Inputs

\n", - "

Outputs

\n", - "\n", - " \n", - " \n", - " " - ], - "text/plain": [ - "" - ] - }, - "metadata": {}, - "output_type": "display_data" - } - ], - "source": [ - "from interpreto import Occlusion\n", - "\n", - "explainer = Occlusion(\n", - " model_gen,\n", - " tokenizer_gen,\n", - " granularity=Granularity.WORD,\n", - " granularity_aggregation_strategy=GranularityAggregationStrategy.MAX,\n", - ")\n", - "\n", - "text = \"I like you\"\n", - "target = \"with you my friend\"\n", - "\n", - "tokenized_texts = tokenizer_gen(text, return_tensors=\"pt\", padding=True, truncation=True, return_offsets_mapping=True)\n", - "\n", - "tokenized_targets = tokenizer_gen(\n", - " target, return_tensors=\"pt\", padding=True, truncation=True, return_offsets_mapping=True\n", - ")[\"input_ids\"]\n", - "attributions = explainer.explain(tokenized_texts, targets=tokenized_targets)\n", - "\n", - "\n", - "print(attributions[0])\n", - "plot_attributions(attributions[0])" - ] - }, - { - "cell_type": "code", - "execution_count": 13, - "id": "8cb4d137", - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "torch.Size([4, 7])\n", - "7\n" - ] - } - ], - "source": [ - "print(attributions[0].attributions.shape)\n", - "print(len(attributions[0].elements))" - ] - }, - { - "cell_type": "code", - "execution_count": 14, - "id": "9d29d0cf", - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "7\n", - "['I', ' like', ' you', ' with', ' you', ' my', ' friend']\n" - ] - } - ], - "source": [ - "print(len(attributions[0].elements))\n", - "print(attributions[0].elements)" - ] - } - ], - "metadata": { - "kernelspec": { - "display_name": ".venv (3.12.3)", - "language": "python", - "name": "python3" - }, - "language_info": { - "codemirror_mode": { - "name": "ipython", - "version": 3 - }, - "file_extension": ".py", - "mimetype": "text/x-python", - "name": "python", - "nbconvert_exporter": "python", - "pygments_lexer": "ipython3", - "version": "3.12.3" - } - }, - "nbformat": 4, - "nbformat_minor": 5 -} From 14be4526fbb3e361b939f01f79c46431e189e0f4 Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Mon, 20 Apr 2026 15:18:42 +0000 Subject: [PATCH 15/17] add attrsim --- .../concepts/metrics/simulatability/consim.py | 8 ------- interpreto/model_wrapping/llm_interface.py | 10 -------- .../interpretation/test_llm_labels.py | 6 ----- tests/visualizations/test_concepts.py | 24 ------------------- 4 files changed, 48 deletions(-) diff --git a/interpreto/concepts/metrics/simulatability/consim.py b/interpreto/concepts/metrics/simulatability/consim.py index d42215ab..5387d2ee 100644 --- a/interpreto/concepts/metrics/simulatability/consim.py +++ b/interpreto/concepts/metrics/simulatability/consim.py @@ -517,14 +517,6 @@ def _setting_to_prompt( # type: ignore[override] # noqa: PLR0912 # ignore too "The most important concepts and their importance for each class are:\n" + "\n".join( [ - # f"\t{class_name}: { - # ConSim._concepts_to_string( - # global_importances[class_index], - # concepts_interpretation, - # top_k=top_k, - # threshold=importance_threshold, - # ) - # }" f"\t{class_name}: { ConSim._concepts_to_string( global_importances[class_index], diff --git a/interpreto/model_wrapping/llm_interface.py b/interpreto/model_wrapping/llm_interface.py index b9e471ad..931d263a 100644 --- a/interpreto/model_wrapping/llm_interface.py +++ b/interpreto/model_wrapping/llm_interface.py @@ -83,11 +83,6 @@ def __init__(self, model: str, batch_size: int = 8, device: str = "auto"): self.tokenizer = AutoTokenizer.from_pretrained(model) self.tokenizer.padding_side = "left" self.tokenizer.truncation_side = "left" - # self.model = AutoModelForCausalLM.from_pretrained( - # model, - # torch_dtype="auto", - # device_map=device, - # ) try: self.model = AutoModelForCausalLM.from_pretrained( model, @@ -149,11 +144,6 @@ def _format_prompt(self, system_prompt: str, user_prompt: str) -> str: {"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, diff --git a/tests/concepts/interpretation/test_llm_labels.py b/tests/concepts/interpretation/test_llm_labels.py index daffabf3..80d9e7df 100644 --- a/tests/concepts/interpretation/test_llm_labels.py +++ b/tests/concepts/interpretation/test_llm_labels.py @@ -273,11 +273,6 @@ 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", @@ -289,7 +284,6 @@ def splitted_encoder() -> ModelWithSplitPoints: 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" diff --git a/tests/visualizations/test_concepts.py b/tests/visualizations/test_concepts.py index f3f798dc..3eaad1ef 100644 --- a/tests/visualizations/test_concepts.py +++ b/tests/visualizations/test_concepts.py @@ -22,30 +22,6 @@ # OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE # SOFTWARE. -# 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, -# 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 json From 90cc3b4281b5cfb7f6a5712a0b1cf54b654ced03 Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Tue, 21 Apr 2026 10:55:49 +0000 Subject: [PATCH 16/17] add attrsim tests --- interpreto/concepts/metrics/simulatability/attrsim.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/interpreto/concepts/metrics/simulatability/attrsim.py b/interpreto/concepts/metrics/simulatability/attrsim.py index b0f09e37..50c7137f 100644 --- a/interpreto/concepts/metrics/simulatability/attrsim.py +++ b/interpreto/concepts/metrics/simulatability/attrsim.py @@ -234,7 +234,6 @@ def construct_prompt( # type: ignore system_prompt_parts = [ "You are a classifier. Predict the class for each evaluation sample.", - "Use the provided learning examples and attribution explanations to infer the model behavior.", "Only return the class name, no additional text.", f"The classes are: [{', '.join(list(classes.values()))}]", ] @@ -243,7 +242,14 @@ def construct_prompt( # type: ignore 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]}", From 0ba4c1e12f4c2a3b8165052a1999d80ea43c1344 Mon Sep 17 00:00:00 2001 From: fanny-jourdan Date: Wed, 22 Apr 2026 16:27:06 +0000 Subject: [PATCH 17/17] add attrsim docstring --- .../metrics/simulatability/attrsim.py | 151 ++++++++++++++++- tests/concepts/metrics/test_attrsim.py | 159 ++++++++++++------ 2 files changed, 246 insertions(+), 64 deletions(-) diff --git a/interpreto/concepts/metrics/simulatability/attrsim.py b/interpreto/concepts/metrics/simulatability/attrsim.py index 50c7137f..d73222e0 100644 --- a/interpreto/concepts/metrics/simulatability/attrsim.py +++ b/interpreto/concepts/metrics/simulatability/attrsim.py @@ -58,6 +58,8 @@ class PromptSetting(NamedTuple): 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, @@ -95,6 +97,8 @@ def validate( 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): @@ -157,6 +161,15 @@ class AttrSim(AutomatedSimulatability): @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 @@ -164,6 +177,24 @@ 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 @@ -182,29 +213,86 @@ 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] - top_k = min(top_k, attr_vector.shape[-1]) - top_indices = torch.topk(attr_vector.abs(), k=top_k).indices.tolist() + 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}: {attr_vector[idx].item():+.3f}") + 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) + 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, @@ -216,6 +304,23 @@ def construct_prompt( # type: ignore *, 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, @@ -255,9 +360,14 @@ def construct_prompt( # type: ignore f"\tLabel: {classes[pred_index]}", ] if setting.lp_attributions: - lp_block.append( - f"\tAttributions for {classes[pred_index]}: {self._format_attribution_for_pred(corresponding_attribution[i], pred_index)}" + 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) @@ -272,10 +382,15 @@ def construct_prompt( # type: ignore text = f"Contrastive Attributions supporting {pred_name} rather than {gold_name}" attr_to_show = pred_attr - gold_attr - lp_block.append( - f"\t{text}: {self._format_attr_vector(corresponding_attribution[i].elements, attr_to_show)}" + 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)) @@ -299,6 +414,26 @@ def _check_input_settings_correspondence( 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): diff --git a/tests/concepts/metrics/test_attrsim.py b/tests/concepts/metrics/test_attrsim.py index f73fa41b..3142877c 100644 --- a/tests/concepts/metrics/test_attrsim.py +++ b/tests/concepts/metrics/test_attrsim.py @@ -31,52 +31,83 @@ from interpreto.concepts.metrics.simulatability.attrsim import AttrSim, PromptSetting -def _build_attribution_output( +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=torch.tensor([0, 1]), - ) # type: ignore + 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.L2_baseline_with_lp.value.lp_attributions is False - assert AttrSim.prompt_types.E1_attribution_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]])) -def test_format_attribution_for_pred_handles_singleton_class_axis(): - attribution = _build_attribution_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 - rendered = AttrSim._format_attribution_for_pred(attribution_output=attribution, pred_index=2, top_k=2) - assert "t1: -0.500" in rendered - assert "t0: +0.200" 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_enum_setting(): +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_attribution_output(torch.tensor([[0.1, -0.2, 0.3]])) for _ in samples] + 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, @@ -87,18 +118,47 @@ def test_construct_prompt_with_enum_setting(): corresponding_attribution=attributions, ) - assert "Attributions for A:" in system_prompt + assert "Attributions:" in system_prompt assert len(user_prompts) == 2 assert model_predictions == ["A", "B"] -def test_construct_prompt_validates_lengths(): +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 = [_build_attribution_output(torch.tensor([[0.1, -0.2, 0.3]])) for _ in range(2)] + 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, @@ -106,69 +166,56 @@ def test_construct_prompt_validates_lengths(): corresponding_predictions=predictions, corresponding_labels=labels, nb_learning_samples=1, - corresponding_attribution=attributions, + corresponding_attribution=attributions_too_short, ) - -def test_construct_prompt_rejects_too_many_learning_samples(): - metric = AttrSim(classes=["A", "B"]) - samples = ["s0", "s1"] - predictions = torch.tensor([0, 1]) - labels = torch.tensor([0, 1]) - attributions = [_build_attribution_output(torch.tensor([[0.1, -0.2, 0.3]])) for _ in samples] - + 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=2, - corresponding_attribution=attributions, + nb_learning_samples=3, + corresponding_attribution=attributions_ok, ) -def test_construct_prompt_with_contrastive_attributions(): +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", "s3"] - predictions = torch.tensor([0, 1, 0, 1]) - labels = torch.tensor([0, 0, 1, 1]) # sample 1 and 2 are misclassified + samples = ["s0", "s1", "s2"] + predictions = torch.tensor([0, 1, 0]) + labels = torch.tensor([0, 0, 1]) + attributions = [ - _build_attribution_output(torch.tensor([[0.4, -0.1, 0.2], [0.1, 0.2, -0.3]])), - _build_attribution_output(torch.tensor([[0.2, -0.5, 0.1], [0.6, -0.1, -0.2]])), - _build_attribution_output(torch.tensor([[0.1, -0.3, 0.5], [-0.2, 0.7, 0.1]])), - _build_attribution_output(torch.tensor([[0.2, 0.2, -0.2], [0.3, -0.4, 0.1]])), + _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]])), ] - 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 "Contrastive Attributions supporting B rather than A" in system_prompt - assert "Attributions for A" in system_prompt - assert len(user_prompts) == 2 - assert model_predictions == ["A", "B"] + 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_construct_prompt_contrastive_requires_classwise_attribution_for_miss(): +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]) # one miss in LP when nb_learning_samples=2 - attributions = [ - _build_attribution_output(torch.tensor([[0.1, -0.2, 0.3]])), - _build_attribution_output(torch.tensor([[0.2, -0.5, 0.1]])), # singleton axis, not classwise - _build_attribution_output(torch.tensor([[0.4, -0.1, 0.2]])), - ] + 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="require class-wise attributions"): + with pytest.raises(ValueError, match="attribution_top_k"): metric.construct_prompt( - setting=AttrSim.prompt_types.C1_contrastive_attribution_with_lp, + setting=PromptSetting(lp_samples=True, lp_attributions=True, attribution_top_k=0), interesting_samples=samples, corresponding_predictions=predictions, corresponding_labels=labels,