Skip to content

feat(representation_geometry): add causal inner products, concept directions, and categorical diagnostics - #1881

Open
janmenjayap wants to merge 6 commits into
TransformerLensOrg:devfrom
janmenjayap:feat/representation-geometry-core
Open

janmenjayap wants to merge 6 commits into
TransformerLensOrg:devfrom
janmenjayap:feat/representation-geometry-core

Conversation

@janmenjayap

Copy link
Copy Markdown
Contributor

Description

Add a reusable representation-geometry tool for TransformerBridge and tensor-level readouts. It derives an explicit metric from centered unembedding covariance, builds oriented concept directions from counterfactual token pairs, and analyzes supplied categorical vertices without assuming that they form regular simplices.

The feature is motivated by Park, Choe and Veitch, arXiv:2311.03658v2. The covariance-inverse metric is one choice under the paper's assumptions, not a universal guarantee of causal separability. Geometry, probe decodability, and a metric-derived intervention do not by themselves establish causal model use.

Adds the core representation-geometry, concept-direction, and categorical-diagnostics part of #1880.

This delivers the usable core of the representation-geometry proposal. Independent paper reproduction and controlled behavioral validation remain outside this PR, so #1880 should remain open after this feature merges.

New dependencies: none. The implementation uses existing PyTorch, transformers, and tokenizer dependencies. The demo and targeted integration tests require no pretrained-model or tokenizer downloads.


Type of change

  • New feature: non-breaking representation-geometry analysis API.
  • Bug fix: preserve locally constructed tokenizers when no pretrained reload source exists.
  • Documentation update: public guide, analysis-tool navigation, API listing, and executed demo.

What is implemented

Explicit dual-space geometry

Export RepresentationGeometry, GeometryBasis, ConceptDirection, and CategoricalGeometry from transformer_lens.tools.analysis.

Fit centered, uniformly weighted population covariance from a [d_model, d_vocab] unembedding tensor. Record token population, storage/compute dtype, numerical rank, thresholds, conditioning, and explicit ridge regularization.

For M = Sigma, or M = Sigma + epsilon I with explicit ridge, provide separate measurement and intervention transforms:

measurement:   g' = M^(-1/2) g
intervention:  h' = M^(1/2) h
pairing:       h'^T g' = h^T g

Include inverse maps, inner products, cosines, batched broadcasting, input validation, and vector-gradient support over detached fitted weights. Zero-direction cosines raise rather than returning an arbitrary value.

Singular/near-singular exact covariance fits raise. There is no silent pseudo-inverse or covariance-eigenvalue clamping. With ridge, whitening identity applies to the regularized covariance, not generally the original covariance.

Counterfactual concept directions

Pair (lo, hi) contributes W_U[:, hi] - W_U[:, lo]. Return an unnormalized mean, per-pair differences, population dispersion, frozen IDs/labels, and raw/whitened measurement and metric-derived intervention values.

Reject empty, invalid, self, duplicate, reversed-duplicate, zero, and non-finite contrasts. A single pair has zero population dispersion, not a confidence guarantee. Derived-coordinate alignment is labeled as an algebraic construction, not an independently observed intervention.

String endpoints use a snapshotted Hugging Face tokenizer, resolve to exactly one known token, and disable implicit BOS/EOS. Preserve leading-space spelling and support mixed string/ID inputs. Token IDs remain usable without a tokenizer.

Categorical diagnostics

Accept explicit concept vertices in a declared measurement or intervention space. Report centering, affine rank, singular values, Gram matrices, distances, cosines, angles, and validity masks.

  • Simplex checks require affine independence.
  • Regular-simplex checks additionally test relative pair-distance spread under a declared tolerance.
  • Duplicate and degenerate vertices are retained; undefined angles are NaN where the mask is false.
  • Unrepresentable transforms/Gram values raise rather than being misreported as coincident categories.

No token-to-category estimator or paper-specific categorical/hierarchical construction is implied.

Non-mutating TransformerBridge adapter

RepresentationGeometry.from_bridge(...) snapshots the raw linear vocabulary-head input after unfolded final LN/RMS. It does not fold weights, apply cached activation scales, enable compatibility mode, move the model, change training flags, or execute hooks/forward passes.

The adapter requires a raw causal decoder with a direct LN/RMS-to-plain-linear readout and consistent dimensions. It rejects processed/compatibility bases, final output projections, active soft caps, custom output transforms, parametrized/custom-forward heads, and meta/offloaded readouts. The exposed W_U must match the actual linear head.

Readout and tokenizer snapshots remain stable after source edits. Basis provenance is recorded but is not a unique model identity or permission to mix coordinate spaces.

Documentation, demo, and local-tokenizer fix

Add demos/RepresentationGeometry_Demo.ipynb and docs/source/content/representation_geometry.md, connect them to the analysis guide and docs navigation, and include the notebook in docs generation.

The demo uses known synthetic geometry and a tiny random GPT-2 Bridge. It demonstrates raw versus metric cosine on a deliberately constructed fixture, dual pairing, concept dispersion, ridge, regular/non-regular/degenerate categorical inputs, tokenizer behavior, and actual post-normalization readout pairing. It does not claim learned semantic geometry or empirical paper parity.

The demo exposed a BOS-setup bug: get_tokenizer_with_bos assumed a pretrained name_or_path existed for an in-memory tokenizer. Preserve local special-token behavior when no reload source exists, with download-free regression tests and a preloaded-Bridge integration check.


Validation

  • make format and make check-format: passed.
  • uv run mypy .: passed.
  • Unit tier: 7,536 passed, 58 skipped, 57 deselected, 3 xfailed.
  • Docstring tier: 17 passed, 24 skipped.
  • Acceptance tier: 163 passed, 69 skipped, 11 deselected.
  • Integration tier: 1,509 passed, 1 failed, 7 skipped, 185 deselected.
  • Demo notebook: 8 nbval checks passed.
  • Focused docs-copy/local-tokenizer/geometry-integration checks: 43 passed.
  • uv run build-docs: completed, with 529 diagnostics across the documentation tree; this was not a warning-free build.

No skips, xfails, or tolerance relaxations were added to obtain these results. The synthetic/tiny-model geometry checks are CPU-based; pretrained-model reproduction and model-scale GPU accuracy are not claimed.


Scope boundaries and interpretation

Not included: independent empirical reproduction, independently fitted contextual interventions, hierarchy tooling, automated contrast/category builders, steering/probing adapters, unrestricted multi-token concepts, or a universal architecture-support matrix.

Earlier residual hooks are not automatically in the post-normalization readout basis. Unembedding bias is excluded from the metric and must be accounted for in logit readouts. Ridge and token-population selection change the estimator; the paper's assumptions do not automatically transfer to arbitrary subsets.

Dense covariance/eigendecomposition and retained readout snapshots can be expensive. The guide documents precision, resource, basis, sparsity, and scientific-interpretation limits. Dense whitening changes coordinate sparsity, and a high geometric score is a hypothesis for controlled experiments, not causal evidence.


Checklist

  • I have documented non-obvious mathematical and coordinate-space contracts.
  • I have made corresponding changes to public documentation and the demo.
  • My changes generate no new warnings; the docs build has diagnostics and a warning-free comparison has not been established.
  • I have added value-based analytic, integration, invalid-input, and bug-regression tests.
  • New and existing unit tests pass locally with these changes.
  • I have not weakened key-interface tests or added skips/tolerance relaxations to bypass failures.
  • All required PR-review tiers pass; the MPS integration failure remains unresolved.
  • The tracking issue is linked with a non-closing reference: [Proposal] Representation Geometry: causal inner products, concept directions, and categorical diagnostics #1880.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant