Skip to content

fix: include Granite embedding scale in residual decomposition - #1878

Open
emerardd wants to merge 1 commit into
TransformerLensOrg:devfrom
emerardd:fix/granite-residual-embedding-scale
Open

emerardd wants to merge 1 commit into
TransformerLensOrg:devfrom
emerardd:fix/granite-residual-embedding-scale

Conversation

@emerardd

Copy link
Copy Markdown
Contributor

Description

Granite applies embedding_multiplier after its embedding module returns. The native embedding hook therefore correctly captures the unscaled embedding, but ActivationCache.decompose_resid() and get_full_resid_decomposition() previously treated it as the full additive embedding contribution. With a non-unit multiplier, component stacks do not sum to the residual stream and component-level direct logit attribution assigns an incorrect embedding contribution, even when output scaling is identity.

This change propagates embedding_multiplier through the existing HF config passthrough and applies it when constructing embedding contributions in both decomposition paths. It does not change the cached native hook, hook editing semantics, model weights, or forward execution. Models without this attribute, and identity multipliers, keep the existing behavior. The API docstrings now distinguish the residual contribution from the native hook value.

The offline regressions use tiny randomly initialized models, without Hub downloads. They cover multipliers 1, 0.5 and 2; raw and compatibility-mode dense Granite; batched and batchless caches; layer-zero, final and position-sliced reconstruction; full decomposition; component DLA with logits_scaling=1; native hook edits and forward parity; and Granite MoE component reconstruction. This PR does not address non-identity output-transform semantics in DLA or claim runtime validation of Granite hybrid models.

Type of change

  • Bug fix (non-breaking change which fixes an issue)

Validation

The initial 18-case regression selection produced 10 failed, 8 passed before the production fix. The affected validation selection after the fix and additional MoE coverage produced 377 passed, 2 deselected. The two deselected cases are unrelated wide-tensor gradient stress tests; no skip or xfail was added.

python -m pytest tests/integration/model_bridge/test_granite_embedding_decomposition.py tests/unit/test_activation_cache.py tests/unit/test_activation_cache_batch_dim.py tests/unit/tools/test_direct_logit_attribution_final_norm.py tests/unit/tools/test_direct_logit_attribution.py tests/unit/tools/test_projection_kernel.py -k "not wide" -q

Pycln, isort, Black and git diff --check passed for the change. Full-repository mypy passed: no issues found in 391 source files. The complete unit suite and full PR suite were not run locally. The two import-time SWIG deprecation warnings also appeared in the pre-fix regression run.

Checklist

  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes
  • I have not rewritten tests relating to key interfaces which would affect backward compatibility

The unit-test checkbox remains unchecked because only the affected selection, not the complete unit suite, was run locally.

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