SIGN IN SIGN UP

fix: MaskedCausalVisionTransformer build on all supported timm versions (#2043)

* fix: forward only attention kwargs in MaskedCausalBlock

MaskedCausalBlock forwarded its full block-level kwargs to
MaskedCausalAttention, including arguments the timm Attention
constructor does not accept (e.g. mlp_ratio). This made
MaskedCausalVisionTransformer, and therefore the AIM model, fail to
construct on every supported timm version (0.9.9-1.0.28) with
"Attention.__init__() got an unexpected keyword argument 'mlp_ratio'".

Forward only the arguments that Attention.__init__ defines, selected via
its signature so the fix stays correct as timm evolves. Add a regression
test covering the block and the vision transformer.

* fix: set global_pool for the AIM masked causal vision transformer

timm's VisionTransformer asserts `class_token or global_pool != 'token'`.
The AIM examples and benchmark build MaskedCausalVisionTransformer with
class_token=False but did not set global_pool, so construction failed on
current timm. AIM's self-supervised path uses forward_features and is not
affected by global_pool, so "avg" satisfies the assertion without changing
behaviour. Regenerate the AIM example notebooks accordingly.
L
Lőrincz-Molnár Szabolcs-Botond committed
8e995ec19cae68f81921164ffa37a165d7a39808
Parent: 413687d
Committed by GitHub <noreply@github.com> on 8/21/2026, 8:19:47 PM