Skip to content

Base Classes & Utilities

The base interpretation class defines the shared interface for making concept dimensions interpretable. The extract_ngrams utility provides reusable preprocessing for text-based interpretations.

API Reference

interpreto.concepts.interpretations.base.BaseConceptInterpretationMethod

BaseConceptInterpretationMethod(concept_explainer, activation_granularity=None, aggregation_strategy=MEAN, concept_encoding_batch_size=1024, use_vocab=False, use_unique_words=0, unique_words_kwargs={})

Bases: ABC

Code: concepts/interpretations/base.py

Abstract class defining an interface for concept interpretation. Its goal is to make the dimensions of the concept space interpretable by humans.

Attributes:

Name Type Description
concept_explainer ConceptEncoderExplainer

The concept explainer used to compute the concept activations.

activation_granularity ActivationGranularity

The granularity of the activations to use for the interpretation. See :method:interpreto.concepts.splitters.model_with_split_points.ModelWithSplitPoints.get_activations for more details.

aggregation_strategy GranularityAggregationStrategy

The aggregation strategy to use for the activations. See :method:interpreto.concepts.splitters.model_with_split_points.ModelWithSplitPoints.get_activations for more details.

concept_encoding_batch_size int

The batch size to use for the concept encoding.

use_vocab bool

Whether to use the vocabulary to extract the granular inputs. If True, the granular inputs are extracted from the vocabulary. If False, the granular inputs are extracted from the inputs.

use_unique_words bool

If True, the interpretation will be computed from the unique words of the inputs. Incompatible with use_vocab=True. Default unique words selects all different word from the input. It can be tuned through the unique_words_kwargs argument.

unique_words_kwargs dict

The kwargs to pass to the extract_ngrams function. see interpreto.concepts.interpretations.topk_inputs.extract_ngrams for more details. Possible arguments are count_min_threshold, lemmatize, words_to_ignore.

Source code in interpreto/concepts/interpretations/base.py
def __init__(
    self,
    concept_explainer: ConceptEncoderExplainer,
    activation_granularity: ActivationGranularity | None = None,
    aggregation_strategy: GranularityAggregationStrategy = GranularityAggregationStrategy.MEAN,
    concept_encoding_batch_size: int = 1024,
    use_vocab: bool = False,
    use_unique_words: bool | int = 0,
    unique_words_kwargs: dict = {},
):
    if activation_granularity is None:
        if isinstance(concept_explainer.splitter, SplitterForClassification):
            activation_granularity = ActivationGranularity.CLS_TOKEN
        else:
            activation_granularity = ActivationGranularity.TOKEN
    elif activation_granularity not in (
        ActivationGranularity.CLS_TOKEN,
        ActivationGranularity.TOKEN,
        ActivationGranularity.WORD,
        ActivationGranularity.SENTENCE,
        ActivationGranularity.SAMPLE,
    ):
        raise ValueError(
            f"The granularity {activation_granularity} is not supported. "
            "Supported `activation_granularities`: CLS_TOKEN, TOKEN, WORD, SENTENCE, and SAMPLE"
        )

    if use_unique_words and use_vocab:
        raise ValueError("Cannot use both `use_unique_words` and `use_vocab`. Please use only one of them.")

    self.concept_explainer: ConceptEncoderExplainer = concept_explainer
    self.activation_granularity: ActivationGranularity = activation_granularity
    self.aggregation_strategy: GranularityAggregationStrategy = aggregation_strategy
    self.concept_encoding_batch_size: int = concept_encoding_batch_size
    self.use_vocab: bool = use_vocab
    self.use_unique_words: int = int(use_unique_words)
    self.unique_words_kwargs: dict = unique_words_kwargs

concepts_activations_from_source

concepts_activations_from_source(*, inputs=None, latent_activations=None, concepts_activations=None)

Computes the concepts activations from the given samples. Samples can be provided as raw text (inputs), latent activations (latent_activations), or directly concept activations (concepts_activations).

Parameters:

Name Type Description Default

inputs

list[str] | None

The indices of the concepts to interpret.

None

latent_activations

Float[Tensor, 'nl d'] | None

The latent activations

None

concepts_activations

Float[Tensor, 'nl cpt'] | None

The concepts activations

None

Returns:

Type Description
Float[Tensor, 'nl cpt']

Float[torch.Tensor, "nl cpt"] :

Source code in interpreto/concepts/interpretations/base.py
def concepts_activations_from_source(
    self,
    *,
    inputs: list[str] | None = None,
    latent_activations: Float[torch.Tensor, "nl d"] | None = None,
    concepts_activations: Float[torch.Tensor, "nl cpt"] | None = None,
) -> Float[torch.Tensor, "nl cpt"]:
    """
    Computes the concepts activations from the given samples.
    Samples can be provided as raw text (`inputs`), latent activations (`latent_activations`),
    or directly concept activations (`concepts_activations`).

    Args:
        inputs (list[str] | None): The indices of the concepts to interpret.
        latent_activations (Float[torch.Tensor, "nl d"] | None): The latent activations
        concepts_activations (Float[torch.Tensor, "nl cpt"] | None): The concepts activations

    Returns:
        Float[torch.Tensor, "nl cpt"] :
    """

    if concepts_activations is not None:
        return concepts_activations

    if latent_activations is not None:
        # batch over latent activations for concept encoding
        concepts_activations_list = []
        with torch.no_grad():
            for batch_idx in range(0, latent_activations.shape[0], self.concept_encoding_batch_size):
                # concept model forward pass
                batch_concepts_activations = self.concept_explainer.activations_to_concepts(
                    latent_activations[batch_idx : batch_idx + self.concept_encoding_batch_size]
                ).cpu()
                concepts_activations_list.append(batch_concepts_activations)
        concepts_activations = torch.cat(concepts_activations_list, dim=0)
        return concepts_activations

    if inputs is not None:
        latent_activations, _ = self.concept_explainer.splitter.get_activations(
            inputs,
            activation_granularity=self.activation_granularity,
            aggregation_strategy=self.aggregation_strategy,
            forward_kwargs={"truncation": True},
        )
        return self.concepts_activations_from_source(latent_activations=latent_activations, inputs=inputs)

    raise ValueError(
        "No source provided. Please provide either `inputs`, `latent_activations`, or `concepts_activations`."
    )

concepts_activations_from_vocab

concepts_activations_from_vocab()

Computes the concepts activations for each token of the vocabulary

Returns:

Type Description
tuple[list[str], Float[Tensor, 'nl cpt']]

tuple[list[str], Float[torch.Tensor, "nl cpt"]]: - The list of tokens in the vocabulary - The concept activations for each token

Source code in interpreto/concepts/interpretations/base.py
@jaxtyped(typechecker=beartype)
def concepts_activations_from_vocab(
    self,
) -> tuple[list[str], Float[torch.Tensor, "nl cpt"]]:
    """
    Computes the concepts activations for each token of the vocabulary

    Returns:
        tuple[list[str], Float[torch.Tensor, "nl cpt"]]:
            - The list of tokens in the vocabulary
            - The concept activations for each token
    """
    # extract and sort the vocabulary
    vocab_dict: dict[str, int] = self.concept_explainer.splitter.tokenizer.get_vocab()
    inputs, input_ids = zip(*vocab_dict.items(), strict=True)  # type: ignore
    inputs: list[str] = list(inputs)  # type: ignore

    # unsqueeze for all ids to be considered as a single sample
    input_ids: Float[torch.Tensor, "v 1"] = torch.tensor(list(input_ids)).unsqueeze(1)
    vocab_size = input_ids.shape[0]

    if self.activation_granularity != ActivationGranularity.CLS_TOKEN:
        # compute the vocabulary's latent activations
        latent_activations, _ = self.concept_explainer.splitter.get_activations(
            input_ids,
            activation_granularity=ActivationGranularity.ALL_TOKENS,
            forward_kwargs={"truncation": True},
        )
    else:
        # we need to add the CLS token and maybe the EOS token to the ids
        # so that we can get correct CLS activations

        # first step extract the template
        template_ids = self.concept_explainer.splitter.tokenizer("a", return_tensors="pt")["input_ids"]

        # if we are not in a template [CLS] a [EOS]
        if len(template_ids) != 3:  # type: ignore
            warnings.warn(
                "When tokenizing a single character, the provided model does not output 3 token ids. "
                "Our implementation assumes that the model outputs is [CLS] a [EOS]. "
                "Indeed, when `aggregation_strategy` is `CLS_TOKEN`, the first token is considered as the CLS token. "
                "If the [CLS] token is still the first token, you can ignore this warning. "
                "Otherwise, either choose another model or contact the developers to find a workaround.",
                stacklevel=2,
            )

        # repeat the template and replace "a" token ids by the vocabulary ids
        repeated_template_ids = template_ids.repeat(vocab_size, 1)  # type: ignore
        repeated_template_ids[:, 1] = input_ids[:, 0]

        # compute the vocabulary's latent activations
        latent_activations, _ = self.concept_explainer.splitter.get_activations(
            repeated_template_ids,
            activation_granularity=self.activation_granularity,
            forward_kwargs={"truncation": True},
        )

    # compute the vocabulary's concepts activations
    with torch.no_grad():
        concepts_activations = self.concept_explainer.activations_to_concepts(latent_activations)
    return inputs, concepts_activations

get_granular_inputs

get_granular_inputs(inputs)

Split texts from the inputs based on the target granularity (for instance into tokens, words, sentences, ...)

Parameters:

Name Type Description Default

inputs

list[str]

n text samples

required

Returns:

Name Type Description
granular_flattened_texts list[str]

The granular texts elements from the inputs, flattened. [Example1_Tok1, Example1_Tok2, ... Example2_Tok1, Example2_Tok2, ...]

granular_flattened_sample_id list[int]

The sample id for each granular text, to keep track of which sample the text belongs to. It should have the same length as granular_flattened_texts. It elements indicates the sample if for the corresponding granular text in granular_flattened_texts. [0, 0, ... 1, 1, ...]

Source code in interpreto/concepts/interpretations/base.py
@jaxtyped(typechecker=beartype)
def get_granular_inputs(
    self,
    inputs: list[str],  # (n)
) -> tuple[list[str], list[int]]:  # (ng,)
    """Split texts from the inputs based on the target granularity
    (for instance into tokens, words, sentences, ...)

    Args:
        inputs (list[str]): n text samples

    Returns:
        granular_flattened_texts (list[str]):
            The granular texts elements from the inputs, flattened.
            [Example1_Tok1, Example1_Tok2, ... Example2_Tok1, Example2_Tok2, ...]

        granular_flattened_sample_id (list[int]):
            The sample id for each granular text, to keep track of which sample the text belongs to.
            It should have the same length as `granular_flattened_texts`.
            It elements indicates the sample if for the corresponding granular text in `granular_flattened_texts`.
            [0, 0, ... 1, 1, ...]
    """
    if self.activation_granularity in (
        ActivationGranularity.SAMPLE,
        ActivationGranularity.CLS_TOKEN,
    ):
        # no activation_granularity is needed
        return inputs, list(range(len(inputs)))

    if self.activation_granularity == ActivationGranularity.TOKEN:
        # we can use the tokenizer to split the inputs into tokens
        granular_texts: list[list[str]] = [
            self.concept_explainer.splitter.tokenizer.tokenize(text) for text in inputs
        ]
    else:
        # Get granular texts from the inputs
        tokens = self.concept_explainer.splitter.tokenizer(
            inputs,
            return_tensors="pt",
            padding=True,
            truncation=True,
            return_offsets_mapping=True,
        )
        granular_texts: list[list[str]] = self.activation_granularity.value.get_decomposition(  # type: ignore  (sure list[list[str]] with return_text=True)
            tokens,
            tokenizer=self.concept_explainer.splitter.tokenizer,
            return_text=True,
        )

    granular_flattened_texts = [text for sample_texts in granular_texts for text in sample_texts]
    granular_flattened_sample_id = [i for i, sample_texts in enumerate(granular_texts) for _ in sample_texts]
    return granular_flattened_texts, granular_flattened_sample_id

get_granular_inputs_and_concept_activations

get_granular_inputs_and_concept_activations(concepts_indices, inputs=None, latent_activations=None, concepts_activations=None)

Compute the granular inputs and concept activations for the specified concepts.

Parameters:

Name Type Description Default

concepts_indices

int | list[int] | Literal['all']

The indices of the concepts to interpret. If "all", all concepts are interpreted.

required

inputs

list[str] | None

The inputs to use for the interpretation. Necessary if not use_vocab,as examples are extracted from the inputs.

None

latent_activations

Float[Tensor, 'nl d'] | None

The latent activations matching the inputs. If not provided, it is computed from the inputs.

None

concepts_activations

Float[Tensor, 'nl cpt'] | None

The concepts activations matching the inputs. If not provided, it is computed from the inputs or latent activations.

None

Returns:

Name Type Description
sure_concepts_indices list[int]

The indices of the concepts to interpret.

granular_inputs list[str]

The granular inputs for the specified concepts. Each element of the list is a single granular input, such as a word.

sure_concepts_activations Float[Tensor, 'nl cpt']

The concepts activations matching the granular inputs.

granular_sample_ids list[int]

The granular sample ids for the specified concepts. Each element of the list is the index of the input sample from which the corresponding granular input was extracted. It has the same length as granular_inputs.

Source code in interpreto/concepts/interpretations/base.py
def get_granular_inputs_and_concept_activations(
    self,
    concepts_indices: int | list[int] | Literal["all"],
    inputs: list[str] | None = None,
    latent_activations: LatentActivations | None = None,
    concepts_activations: ConceptsActivations | None = None,
) -> tuple[list[int], list[str], Float[torch.Tensor, "nl cpt"], list[int]]:
    """
    Compute the granular inputs and concept activations for the specified concepts.

    Args:
        concepts_indices (int | list[int] | Literal["all"]):
            The indices of the concepts to interpret. If "all", all concepts are interpreted.

        inputs (list[str] | None):
            The inputs to use for the interpretation.
            Necessary if not `use_vocab`,as examples are extracted from the inputs.

        latent_activations (Float[torch.Tensor, "nl d"] | None):
            The latent activations matching the inputs. If not provided,
            it is computed from the inputs.

        concepts_activations (Float[torch.Tensor, "nl cpt"] | None):
            The concepts activations matching the inputs. If not provided,
            it is computed from the inputs or latent activations.

    Returns:
        sure_concepts_indices (list[int]):
            The indices of the concepts to interpret.

        granular_inputs (list[str]):
            The granular inputs for the specified concepts.
            Each element of the list is a single granular input, such as a word.

        sure_concepts_activations (Float[torch.Tensor, "nl cpt"]):
            The concepts activations matching the granular inputs.

        granular_sample_ids (list[int]):
            The granular sample ids for the specified concepts.
            Each element of the list is the index of the input sample from which the corresponding granular input was extracted.
            It has the same length as `granular_inputs`.

    """
    if concepts_indices == "all":
        concepts_indices = list(range(self.concept_explainer.concept_model.nb_concepts))

    # compute the concepts activations from the provided source, can also create inputs from the vocabulary
    if self.use_vocab:
        # --------------------------------------------------------------------------------------
        # Case 1: use_vocab=True
        granular_inputs: list[str]
        sure_concepts_activations: Float[torch.Tensor, "nl cpt"]
        granular_inputs, sure_concepts_activations = self.concepts_activations_from_vocab()

        granular_sample_ids: list[int] = list(range(len(granular_inputs)))
    else:
        if inputs is None:
            raise ValueError("Inputs must be provided when `use_vocab` is False.")

        if self.use_unique_words >= 1:
            # ----------------------------------------------------------------------------------
            # Case 2: use_unique_words >= 1
            # first list unique words/ngrams from the inputs and compute the activations from them
            if self.activation_granularity not in [
                ActivationGranularity.CLS_TOKEN,
                ActivationGranularity.SAMPLE,
            ]:
                raise ValueError(
                    f"`use_unique_words` requires `activation_granularity=CLS_TOKEN`, "
                    f"got `{self.activation_granularity}`. "
                    "Ngram-based interpretation relies on the CLS token activation "
                    "to represent each ngram as a single unit."
                )
            granular_inputs: list[str] = extract_ngrams(
                inputs=inputs,
                n=self.use_unique_words,
                return_counts=False,
                **self.unique_words_kwargs,
            )  # type: ignore  (sure list[str] with return_counts=False)
            if latent_activations is not None and concepts_activations is not None:
                warnings.warn(
                    "`latent_activations` or `concepts_activations` were provided, "
                    "but `use_unique_words` is True. "
                    "Therefore, the inputs and activations will likely mismatch. "
                    "Either do not provide `latent_activations` and `concepts_activations`, "
                    "or use `interpreto.concepts.interpretation.extract_ngrams` yourself, "
                    "and set `use_unique_words` to False.",
                    stacklevel=2,
                )
            sure_concepts_activations = self.concepts_activations_from_source(
                inputs=granular_inputs,
                latent_activations=latent_activations,
                concepts_activations=concepts_activations,
            )

            granular_sample_ids: list[int] = list(range(len(granular_inputs)))
        else:
            # ----------------------------------------------------------------------------------
            # Case 3: Default, use_vocab=False and use_unique_words=False
            sure_concepts_activations = self.concepts_activations_from_source(
                inputs=inputs,
                latent_activations=latent_activations,
                concepts_activations=concepts_activations,
            )
            granular_inputs: list[str]
            granular_sample_ids: list[int]
            granular_inputs, granular_sample_ids = self.get_granular_inputs(inputs)

    sure_concepts_indices = verify_concepts_indices(
        concepts_activations=sure_concepts_activations,
        concepts_indices=concepts_indices,
    )
    verify_granular_inputs(
        granular_inputs=granular_inputs,
        sure_concepts_activations=sure_concepts_activations,
        latent_activations=latent_activations,
        concepts_activations=concepts_activations,
    )

    return (
        sure_concepts_indices,
        granular_inputs,
        sure_concepts_activations,
        granular_sample_ids,
    )

interpret abstractmethod

Interpret the concepts dimensions in the latent space into a human-readable format. The interpretation is a mapping between the concepts indices and an object allowing to interpret them. It can be a label, a description, examples, etc.

Parameters:

Name Type Description Default

concepts_indices

int | list[int] | Literal['all']

The indices of the concepts to interpret. If "all", all concepts are interpreted.

required

inputs

list[str] | None

The inputs to use for the interpretation. Necessary if not use_vocab,as examples are extracted from the inputs.

None

latent_activations

Float[Tensor, 'nl d'] | None

The latent activations matching the inputs. If not provided, it is computed from the inputs.

None

concepts_activations

Float[Tensor, 'nl cpt'] | None

The concepts activations matching the inputs. If not provided, it is computed from the inputs or latent activations.

None

Returns:

Type Description
Mapping[int, Any]

Mapping[int, Any]: The interpretation of each of the specified concepts.

Source code in interpreto/concepts/interpretations/base.py
@abstractmethod
def interpret(
    self,
    concepts_indices: int | list[int],
    inputs: list[str] | None = None,
    latent_activations: LatentActivations | None = None,
    concepts_activations: ConceptsActivations | None = None,
) -> Mapping[int, Any]:
    """
    Interpret the concepts dimensions in the latent space into a human-readable format.
    The interpretation is a mapping between the concepts indices and an object allowing to interpret them.
    It can be a label, a description, examples, etc.

    Args:
        concepts_indices (int | list[int] | Literal["all"]):
            The indices of the concepts to interpret. If "all", all concepts are interpreted.

        inputs (list[str] | None):
            The inputs to use for the interpretation.
            Necessary if not `use_vocab`,as examples are extracted from the inputs.

        latent_activations (Float[torch.Tensor, "nl d"] | None):
            The latent activations matching the inputs. If not provided,
            it is computed from the inputs.

        concepts_activations (Float[torch.Tensor, "nl cpt"] | None):
            The concepts activations matching the inputs. If not provided,
            it is computed from the inputs or latent activations.

    Returns:
        Mapping[int, Any]:
            The interpretation of each of the specified concepts.
    """
    raise NotImplementedError

interpreto.concepts.interpretations.extract_ngrams

extract_ngrams(inputs, n=1, count_min_threshold=1, return_counts=False, lemmatize=False, words_to_ignore=None)

Extract n-grams (from 1-gram up to n-gram of words) from a list of texts.

If n=3, it extracts 1-grams, 2-grams, and 3-grams.

Parameters:

Name Type Description Default

inputs

Iterable[str]

The texts to extract n-grams from.

required

n

int

The maximum n-gram size. All sizes from 1 to n are extracted.

1

count_min_threshold

int

The minimum total number of occurrences of an n-gram in the whole inputs.

1

return_counts

bool

Whether to return the counts of each n-gram. Defaults to False.

False

lemmatize

bool

Whether to lemmatize words before counting.

False

words_to_ignore

list[str] | None

A list of words to ignore (applied to individual tokens before forming n-grams).

None

Returns:

Type Description
list[str] | Counter[str]

list[str] | Counter[str]: The list of unique n-grams or the counts of each n-gram.

Source code in interpreto/concepts/interpretations/base.py
@jaxtyped(typechecker=beartype)
def extract_ngrams(
    inputs: Iterable[str],
    n: int = 1,
    count_min_threshold: int = 1,
    return_counts: bool = False,
    lemmatize: bool = False,
    words_to_ignore: list[str] | None = None,
) -> list[str] | Counter[str]:
    """
    Extract n-grams (from 1-gram up to n-gram of words) from a list of texts.

    If n=3, it extracts 1-grams, 2-grams, and 3-grams.

    Args:
        inputs (Iterable[str]):
            The texts to extract n-grams from.

        n (int):
            The maximum n-gram size. All sizes from 1 to n are extracted.

        count_min_threshold (int, optional):
            The minimum total number of occurrences of an n-gram in the whole `inputs`.

        return_counts (bool, optional):
            Whether to return the counts of each n-gram.
            Defaults to False.

        lemmatize (bool, optional):
            Whether to lemmatize words before counting.

        words_to_ignore (list[str] | None, optional):
            A list of words to ignore (applied to individual tokens before forming n-grams).

    Returns:
        list[str] | Counter[str]:
            The list of unique n-grams or the counts of each n-gram.
    """
    _ensure_nltk_resources(lemmatize=lemmatize)

    if lemmatize:
        lemmatizer = WordNetLemmatizer()

    tuple_ngram_counts: Counter[tuple[str]] = Counter()

    for text in inputs:
        tokens = word_tokenize(text)

        # preprocess tokens
        processed = []
        for word in tokens:
            if lemmatize:
                word = lemmatizer.lemmatize(word.lower())  # noqa: PLW2901  # type: ignore  (ignore possibly unbound)
            if words_to_ignore is not None and word in words_to_ignore:
                continue
            processed.append(word)
            tuple_ngram_counts[(word,)] += 1  # unigram tuple

        for size in range(2, n + 1):  # skips size 1 as covered over
            for i in range(len(processed) - size + 1):
                tuple_ngram_counts[tuple(processed[i : i + size])] += 1  # >1-gram tuples

    str_ngram_counts: Counter[str] = Counter(
        {
            " ".join(key): count  # convert ngram tuples to strings
            for key, count in tuple_ngram_counts.items()
            if count >= count_min_threshold  # filter too rare n-grams
        }
    )

    if return_counts:
        return str_ngram_counts

    return list(str_ngram_counts.keys())