Migrate to NNsight 0.8 and simplify concepts API - #164
AntoninPoche wants to merge 32 commits into
Conversation
Update the supported NNsight dependency range from 0.7.x to 0.8.x. Require nnsight>=0.8.0rc1,<0.9.0, enabling the splitter migration to the new TransformersModel, tracing, and module envoy APIs.
- Replace LanguageModel with TransformersModel and explicit task loading. - Forward init checks to nnsight.TransformersModel. - Remove the obsolete splitting_utils.py. - Resolve split points directly from NNsight Envoys. - Store the resolved split module alongside its path. - Remove tuple indices caching.
- Simplify init and rely on NNsight’s text-generation task. - Remove some nns_output check (changed in NNsight 0.8). - Reintegration is now a direct module output assignment. - Simplify the get_latent_shape method (relying on new version and other methods).
- Add token pooling after special token filtering (mean, max, min, signed_max, first, and last pooling.) - Factorize get_activations as a batched inputs_to_activations method. - Add tests for token pooling.
- The concept-output gradient was only used for one sample at a time, in the notebooks. So there is no need for fancy batching, we can just treat samples one by one. This solves the target indexing issue with padding (bug of previous implementation). - This also simplify reintegration of activations, these now rely on boolean indexing. - Finally, the gradient now uses `.backward()` instead of `torch.autograd.grad` as suggested by NNsight 0.8.
- Stop forcing extracted activations to float32. - Preserve dtype when passing activations to concept encoders. - Keep detached public activations copied to CPU. - Verify bfloat16 extraction and gradient execution.
- Load classifiers through the NNsight text-classification task. - Resolve classification heads using NNsight Envoy paths. - Support lazy dispatch for precomputed activations. - Rely on `activations_to_outputs` for gradients. - Determine latent shape through a real trace.
- Factorize batch preparation - Look at classification head dimensionality. 3D in and out were originally generation models. We use their right-most non-padding token.
- Remove try-except and always pass (batch, 1, hidden_dim) to the head. Therefore, all models have the same head input shape. - Force get_latent_shape to return (1, hidden_dim). - Extend tests to better verify heads and logits outputs.
- As the following commits will remove `ModelWithSplitPoints`, we needed a way to apply the TOKEN granularity to non-generation models. The `SplitterForGeneration` was already capable of doing so. Therefore, it was renamed to `TextTokensSplitter` and the `task` argument was removed. - `SplitterForGeneration` still exists as backward compatibility, but it is deprecated. - The `TextTokensSplitter` gradient method supposes gradients as it works with outputs. This is not the cleanest, but it is simpler than keeping both classes alive. - Some tests were added for the new splitter and some updated for the old one.
NNsight do not load models on device while trace is not called. Therefore, checking for a splitter device might lead to the meta device, hence we should not rely on it.
- Before, the splitter and concept models needed to be on the same dtype. It is not the case anymore, we can have bfloat16 splitters and float64 concepts. - The `_normalize_to_concept_model` helper is a bit complex because concept models can come from overcomplete and we do not have the full hand on their device and dtype. - The idea is that splitter or concept models model the activations to their device/dtype before inference.
The granularity in concepts was used by the `ModelWithSplitPoints`, but it will be removed in following commits. So we can simplify the concepts-to-outputs gradient method.
- It is now managed by the splitter class and the `token_pooling` argument. - A lot of the work is passed on to the splitter.
The previous implementation was too specific. In addition managing such tokenization manipulation was moved to the splitter.
Also adapt some tests by removing the notion of granularity in the concepts.
- Remove ModelWithSplitPoints - Add TextTokensSplitter (include pooling) - Mark SplitterForGeneration as deprecated - Remove mentions to the granularity - Update examples and descriptions
The method does not have any dtype before fitting. But we need to convert the activations to the model dtype before fitting. So we use a fallback dtype of torch.get_default_dtype().
- Remove the `Role` API, use a tuple (system_prompt, user_prompt) instead. - Simplify `OpenAILLM` in consequence. - Add a `batch_generate` method to `LLMInterface` to batch user_prompts over a common system_prompt. - Introduce an `HuggingFaceLLM` class, which can be initiated from a model or a repository ID. - Introduce a `_SplitterLLM` class, which can be initiated from a `BaseSplitter` instance. - Add a `_resolve_llm_interface` function to create a `LLMInterface` instance from a variety of inputs.
- support passing a repo id, a tuple of (model, tokenizer), a LLMInterface, or use the concept explainer's splitter when it is generation-capable - batch_generate instead of generate
The other concept classification notebooks use TopKInputs with `use_unique_words=3`. So a notebook solely for ngrams is redundant.
- Update to the splitters and concepts api post nnsight 0.8 migration - Re-run the concept notebooks to ensure they work
- Pass from python 3.10-3.13 to 3.11-3.14 - We need `accelerate` for some part of the new api - Our tests use the latest libraries versions, so some versions were not tested anymore (e.g. `transformers>=4.22,<5` while, the last version is `transformers==5.17`) - Clean some unnecessary dependencies, including nvidia binaries for torch (now torch resolves them automatically) - Smoother dependencies grouping (add test to docs, notebook, lint).
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
Multiple unresolved runtime and public API regressions affect core splitter and concept-labeling workflows.
Get a fresh assessment by requesting another Copilot review.
Review effort: Balanced
Findings: 4
Open (8)
Configure left padding and pad tokens in batch generation · New Preserve None for manual ConSim mode · New Use TransformerModel for NNsight 0.8 compatibility · New Restore pad token fallback and synchronize model config · New Support device selection for repository shorthand · New Honor per-call batch size in select_examples · New Use attention masks for last-token selection · New Build templates from a single known non-special token · New
What changed in this PR
Migrates the concept pipeline to NNsight 0.8 while simplifying splitters, interpretation APIs, model precision handling, and local LLM integration.
Changes:
- Replaces granularity-based splitting with classification and token-focused splitters.
- Adds pooling, dtype/device normalization, and batched local LLM support.
- Updates dependencies, documentation, exports, and tests for the new API.
| File | Description |
|---|---|
| tests/conftest.py | Updates shared splitter fixtures. |
| tests/concepts/splitters/test_text_tokens_splitter.py | Tests token and encoder splitting. |
| tests/concepts/splitters/test_splitter_for_generation.py | Expands generation splitter coverage. |
| tests/concepts/splitters/test_splitter_for_classification.py | Tests classification representations. |
| tests/concepts/metrics/test_sparsity_metrics.py | Migrates sparsity tests. |
| tests/concepts/metrics/test_reconstruction_metrics.py | Migrates reconstruction fixtures. |
| tests/concepts/metrics/test_dictionary_metrics.py | Migrates dictionary metric tests. |
| tests/concepts/metrics/test_consim.py | Updates ConSim and LLM tests. |
| tests/concepts/methods/test_torch_probe_explainer.py | Updates probe and precision tests. |
| tests/concepts/methods/test_sklearn_wrappers.py | Tests low-precision normalization. |
| tests/concepts/methods/test_sklearn_probe_explainer.py | Migrates sklearn probe tests. |
| tests/concepts/methods/test_neurons_as_concepts.py | Updates splitter typing. |
| tests/concepts/methods/test_concept_autoencoder_explainers.py | Tests gradients and dtype handling. |
| tests/concepts/interpretation/test_topk_inputs.py | Updates interpretation-mode tests. |
| tests/concepts/interpretation/test_llm_labels.py | Tests revised LLM labeling. |
| tests/concepts/interpretation/test_inputs_to_concepts_attributions.py | Tests lazy splitter handling. |
| tests/concepts/interpretation/test_base_interpretations.py | Adds representation-alignment tests. |
| README.md | Documents the new splitters. |
| pyproject.toml | Updates Python and dependencies. |
| mkdocs.yml | Revises splitter documentation navigation. |
| Makefile | Installs test dependencies for development. |
| interpreto/concepts/splitters/text_tokens_splitter.py | Adds generalized token splitting and pooling. |
| interpreto/concepts/splitters/splitting_utils.py | Removes legacy splitting utilities. |
| interpreto/concepts/splitters/splitter_for_classification.py | Generalizes classification representation extraction. |
| interpreto/concepts/splitters/base_splitter.py | Migrates the NNsight base wrapper. |
| interpreto/concepts/splitters/__init__.py | Exports new and compatibility splitters. |
| interpreto/concepts/probes/sklearn.py | Normalizes tensors for sklearn. |
| interpreto/concepts/probes/base.py | Normalizes probe inputs. |
| interpreto/concepts/metrics/consim.py | Integrates the revised LLM interface. |
| interpreto/concepts/methods/sklearn_wrappers.py | Normalizes sklearn wrapper precision. |
| interpreto/concepts/methods/overcomplete.py | Updates dtype handling and examples. |
| interpreto/concepts/methods/cockatiel.py | Updates splitter documentation. |
| interpreto/concepts/interpretations/topk_inputs.py | Replaces granularity with pooling. |
| interpreto/concepts/interpretations/llm_labels.py | Adds batched flexible LLM labeling. |
| interpreto/concepts/interpretations/base.py | Aligns examples with splitter representations. |
| interpreto/concepts/base.py | Centralizes concept-model normalization. |
| interpreto/concepts/__init__.py | Exports TextTokensSplitter. |
| interpreto/commons/llm_interface.py | Adds local and splitter-backed LLM adapters. |
| interpreto/__init__.py | Updates public exports and version lookup. |
| docs/index.md | Introduces the new splitter API. |
| docs/api/concepts/splitters/text_tokens_splitter.md | Documents token splitting and pooling. |
| docs/api/concepts/splitters/splitter_for_generation.md | Marks generation splitter deprecated. |
| docs/api/concepts/splitters/splitter_for_classification.md | Documents classification representations. |
| docs/api/concepts/splitters/model_with_split_points.md | Removes obsolete splitter documentation. |
| docs/api/concepts/probes.md | Revises probe workflows. |
| docs/api/concepts/overview.md | Rewrites the concept API overview. |
| docs/api/concepts/interpretations/topk_inputs.md | Documents interpretation modes. |
| docs/api/concepts/interpretations/llm_labels.md | Documents LLM selection and batching. |
| AGENTS.md | Updates repository architecture guidance. |
| .github/workflows/build.yml | Updates Python matrix and dependencies. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
| """ | ||
| self.splitter = splitter | ||
| self.user_llm: LLMInterface | None = user_llm | ||
| self.user_llm: LLMInterface = _resolve_llm_interface(user_llm) |
| The predictions of the model on the interesting samples. | ||
| """ | ||
| predictions = self._get_predictions(inputs, batch_size=batch_size, device=device) | ||
| _, predictions = self.splitter.get_activations(inputs, batch_size=batch_size) |
There was a problem hiding this comment.
The batch_size argument of ConSim.select_examples() should be removed.
| non_padding = input_ids != pad_token_id | ||
| token_indices = torch.arange(input_ids.shape[-1], device=input_ids.device) | ||
| positions = (token_indices * non_padding).argmax(-1).to(activations.device) |
There was a problem hiding this comment.
Should receive the prepared tensor mapping instead of the input_ids. Then, use the attention_mask when available.


Many lines are shown because the notebooks were rerun.
interpreto/: 19 files changed, 903 insertions(+), 2368 deletions(-)tests/: 18 files changed, 800 insertions(+), 1223 deletions(-)This should be reviewed commit by commit, the commit message explains the changes in detail.
Description
This PR originated from a migration to NNsight 0.8 but resulted in some big modifications. Overall, it simplifies the API and makes things more robust.
Migration to NNsight 0.8
nnsight.TransformerModelinstead ofnnsight.LanguageModel.taskargument, but removes theautomodel.splitting_utils.pywas removed; we rely onTransformerModel.named_modules()instead.nns_outputwarning.split_pointand asplit_module.Update classification and generation splitters
I use the old granularity names for it to make sense, but these are not used anymore.
SplitterForClassificationhas a more robust extraction of theCLS_TOKEN.SplitterForGenerationbecomes more general and is renamedTextTokensSplitterTOKENSandALL_TOKENSgranularity through theinclude_special_tokensargument.SAMPLEgranularity by passing atokens_pooling.TextTokensSplitter.concept_output_gradientsfunction was simplified by forcing it to treat inputs one sample at a time. (In the notebooks)Removing
ModelWithSplitPointsandGranularity(from concepts)ModelWithSplitPointsbecause theWORDSandSENTENCESwere not used for concepts. The other granularities are better addressed and covered by the other splitters.Granularityfrom the splitters, it is replaced by the choice of the splitter and theinclude_special_tokensandtokens_poolingparameters.Should we rename the
TextTokensSplittertoModelWithSplitPointsand act as it was a refacto?dtypeanddevicemanagementPreviously, everything was forced to
torch.float32and there were somedeviceconflicts.So now, it is the
splitterandconcept_modelroles to convert activations to theirdtypeanddevicebefore inference.LLM Interface
OpenAILLManymore.Roleand addbatch_generateHuggingFaceLLMand_SplitterLLM_resolve_llm_interfaceto create the correctLLMInterfacefrom diverse inputs (str, tuple(model, tokenizer), LLMInterface, or None -> the splitter).Interpretations
token_poolingparameter.use_vocabclassification trick toSplitterForClassification.LLMLabelsllm_interface.LLMLabelsbase promptsNotebooks
use_unique_words=3.OpenAILLM, either instantiate anHuggingFaceLLMor pass the repo id directly to thellm_interfaceattribute. Now everything runs locally.Dependencies update
In addition, the
acceleratepackage was missing from the dependencies for a new update. So I took the time to clean them a little.3.10-3.13->3.11-3.14.transformersType of Change
Checklist
CODE_OF_CONDUCT.mddocument.CONTRIBUTING.mdguide.make lint.make test.