Reviewed by ByteBell

huggingface/transformers

πŸ€— Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.

#49325 fix(modeling_utils): propagate attention implementation to nested sub-configs (#49137)

Judged at indexed commit094f710145
Base080c288fe6
Headc9dc5b8a56

6 Oct 2026 at 00:14 UTC

deepseek-ai/DeepSeek-V4.1-Flash Β· $0.89

Two correctness regressions block this: the adjusted attention implementation never reaches submodules, and nested "" keys are silently dropped.

Fix the two blockers, make the flat-key depth rule deliberate and tested, and bring the experts setter in line.

Detailed review

1. Adjusted attention implementation is not propagated

The resolver maps sub-configs to implementations:

target_implementations = self._get_subconfig_target_implementations(attn_implementation)
...
sub_implementation = target_implementations.get(id(submodule.config), requested_implementation)

But attn_implementation is the caller's raw value. The adjusted value is produced earlier:

requested_implementation = self._check_and_adjust_attn_implementation(...)

Base handed that adjusted value straight to submodules:

sub_implementation = requested_implementation

Because the mapping wins for every config reachable through sub_configs, submodules receive the unadjusted string. That loses the sdpa→eager downgrade, the paged| strip and the _compatible_flash_implementations correction. For flash_attention_2 with the package absent, the root becomes FLASH_ATTN_KERNEL_FALLBACK[...], but a submodule re-resolving raw flash_attention_2 takes the get_correct_attn_implementation path and raises ImportError from _flash_attn_can_dispatch. The shipped test tests/utils/test_modeling_utils.py:3333-3337 pins that after set_attn_implementation("flash_attention_2") the value the model ends on is the KERNEL FALLBACK, not the raw requested name; that is now true of the root only. Pass requested_implementation into _get_subconfig_target_implementations so the mapping values are the adjusted implementation.

2. Nested "" keys are ignored

The resolver only reads a dict's own "" key at the root:

if is_root:
    target = spec.get("", current_attn)
else:
    target = inherited_impl

For {"vision_config": {"": "eager"}}, the nested dict's "" key is never read, so that branch stays at inherited_impl or its current value. The construction-time setter this PR uses as its reference does value.get("", current_attn) at every level and re-enters itself via subconfig._attn_implementation = sub_implementation (configuration_utils.py:462-478), so the same input would set the nested config eager at construction. Runtime and construction therefore disagree. Read spec.get("", ...) at every level, matching the construction-time setter.

3. Flat keys now bind at any depth

The resolver passes the same root dict down every level:

elif child_dict is not None:
    _resolve(subconfig, child_dict, child_inherited, is_root=False)

Base matched a flat key only against direct children:

for subconfig_key in self.config.sub_configs:
    ...
    getattr(self.config, subconfig_key) is submodule.config

Now a flat key naming a deeper sub-config is honored wherever it appears. In the shipped qwen3_omni_moe tree this is not hypothetical: Qwen3OmniMoeThinkerConfig.sub_configs declares text_config, and Qwen3OmniMoeTalkerConfig.sub_configs declares text_config again, one level deeper than the root. A single flat text_config key now applies to both branches; under the old one-level match it could only name a direct child of the root. Decide the intended depth rule before shipping; if flat keys should stay direct-child-only, restore the base matching, otherwise document and test the new behavior.

4. set_experts_implementation remains one level deep

The sibling setter still uses the old one-level match:

for subconfig_key in self.config.sub_configs:
    ...
    getattr(self.config, subconfig_key) is submodule.config

The config-side _experts_implementation setter is already recursive:

value.get(subconfig_key, current_subconfig_moe)
...
subconfig._experts_implementation = sub_implementation

So construction-time dispatch reaches nested experts configs while the runtime setter does not. get_experts_implementation is documented as the counterpart of set_experts_implementation, to snapshot and restore (modeling_utils.py:1761-1780), and that restore is what needs nested awareness. Make set_experts_implementation nested-aware like attention.

5. Minor: docs, annotation, and test coverage

  • modeling_utils.py:1964: the annotation says dict[int, str], but current_attn = getattr(cfg, "_attn_implementation", None) is the fallback target, and mapping[id(cfg)] = target stores it unchecked. A nested sub-config built with no implementation yields None. get_correct_attn_implementation accepts str | None and maps None to "sdpa", so the None is absorbed downstream rather than impossible. Make it dict[int, str | None].
  • modeling_utils.py:2008: the docstring still says dict keys are "the sub_configs name". The resolver now accepts nested dicts and binds a flat key at any depth (1973-2003), so the documented contract is wrong. Update it to describe the nested form and the depth rule.
  • tests/models/deepseek_ocr2/test_modeling_deepseek_ocr2.py:154-180: the four parts cover a branch-named string, a nested dict, a flat key and flag cleanup. The nested "" form and a name reused at two depths (qwen3_omni_moe's text_config) are absent, so the two cases the resolver disagrees on are untested. Add them.
0blocker
2major
4minor
4/4hunks reviewed
11base files read
12m 39stime
$0.89cost

1.98M input and 229k output tokens.

6 findings β€” 2 of 2 files reviewed.

Findings

majorBugsrc/transformers/modeling_utils.py:2042graph behind

`target_implementations` is built from the raw argument, so submodule configs receive the unadjusted string and the kernel and sdpa-to-eager fallbacks no longer reach them.

  • src/transformers/modeling_utils.py:2042 this change β€” `target_implementations = self._get_subconfig_target_implementations(attn_implementation)` passes the CALLER's value, not the `requested_implementation` lines 1991-1995 replaced with the adjusted one.
  • src/transformers/modeling_utils.py:2065 this change β€” `sub_implementation = target_implementations.get(id(submodule.config), requested_implementation)` β€” the mapping wins for every config reachable through `sub_configs`, so the raw value is what a submodule gets.
  • src/transformers/modeling_utils.py:1991-1995 base β€” `requested_implementation = self._check_and_adjust_attn_implementation(...)` is where the sdpa->eager downgrade, the `paged|` strip and the `_compatible_flash_implementations` correction are applied to the value submodels used to receive.
  • src/transformers/modeling_utils.py:2017 base β€” Base handed each submodule that ADJUSTED value (`sub_implementation = requested_implementation`), so a composite's submodules inherited the model-wide fallback.
  • src/transformers/modeling_utils.py:1792-1823 base β€” The kernel fallback: with `flash_attention_2` requested and the package absent the root becomes `FLASH_ATTN_KERNEL_FALLBACK[...]`; a submodule now re-resolving the raw `flash_attention_2` takes the `get_correct_attn_implementation` path instead and raises ImportError from `_flash_attn_can_dispatch`.
  • tests/utils/test_modeling_utils.py:3333-3337 base β€” The shipped test pins the intent that after `set_attn_implementation("flash_attention_2")` the value the model ends on is the KERNEL FALLBACK, not the raw requested name β€” which is now true of the root only.
majorBugsrc/transformers/modeling_utils.py:1974–1979graph behind

Nested `""` keys are ignored: `spec.get("", ...)` runs only when `is_root`, so `{"vision_config": {"": "eager"}}` leaves that branch unchanged while construction sets it eager.

  • src/transformers/modeling_utils.py:1974-1979 this change β€” `if is_root: target = spec.get("", current_attn) else: target = inherited_impl if ... ` β€” a nested dict's own `""` key is never read, so the branch is left at `inherited_impl`/its current value.
  • src/transformers/configuration_utils.py:462-478 base β€” The construction-time setter the PR names as its reference does `value.get("", current_attn)` at EVERY level and re-enters itself via `subconfig._attn_implementation = sub_implementation`, so a nested `""` IS honoured there.
minorCompatibilitysrc/transformers/modeling_utils.py:1998graph behind

Flat keys now bind at any depth, so `text_config`, reused at two levels in the shipped qwen3_omni_moe tree, is applied to both branches by one key.

  • src/transformers/modeling_utils.py:1998-2003 this change β€” `elif child_dict is not None: _resolve(subconfig, child_dict, child_inherited, is_root=False)` hands the same root dict to every descendant, so a key naming a deeper sub-config is honoured wherever it appears.
  • src/transformers/modeling_utils.py:2019-2025 base β€” Base matched `for subconfig_key in self.config.sub_configs` against `getattr(self.config, subconfig_key) is submodule.config`, so a flat key could only ever name a direct child of the root.
  • src/transformers/models/qwen3_omni_moe/configuration_qwen3_omni_moe.py:219-223 base β€” `Qwen3OmniMoeThinkerConfig.sub_configs` declares `text_config`.
  • src/transformers/models/qwen3_omni_moe/configuration_qwen3_omni_moe.py:470-473 base β€” `Qwen3OmniMoeTalkerConfig.sub_configs` declares `text_config` again, one level deeper than the root β€” a single flat key now selects both branches.
  • tests/models/deepseek_ocr2/test_modeling_deepseek_ocr2.py:154-180 this change β€” The four parts cover a branch-named string, a nested dict, a flat key and flag cleanup; the nested `""` form and a name reused at two depths (qwen3_omni_moe's `text_config`) are absent, so the two cases the resolver disagrees on are untested.
  • src/transformers/modeling_utils.py:1967-2003 this change β€” The resolver drops a nested `""` (`is_root` guard) and binds a flat key at any depth, so both untested inputs take a path no existing assertion pins.
minorStale textsrc/transformers/modeling_utils.py:2008graph behind

The docstring still says dict keys are "the sub_configs name"; the resolver now accepts nested dicts and binds a flat key at any depth, so the documented contract is wrong.

  • src/transformers/modeling_utils.py:1973-2003 this change β€” `spec` may be a dict, and `child_dict` is handed down every level, so a flat key such as `sam_config` now matches a sub-config name anywhere in the hierarchy.
  • src/transformers/modeling_utils.py:1969-1971 base β€” The unchanged docstring: 'a `dict` where keys are the sub_configs name, in which case each submodel will dispatch the corresponding value' β€” no nested form, no depth rule.
minorStale textsrc/transformers/modeling_utils.py:1964graph behind

Annotation `dict[int, str]` is wrong: `target` is `None` for any config without an `_attn_implementation`, and that `None` is what flows into `get_correct_attn_implementation`.

  • src/transformers/modeling_utils.py:1968-1985 this change β€” `current_attn = getattr(cfg, "_attn_implementation", None)` is the fallback `target`, and `mapping[id(cfg)] = target` stores it unchecked β€” a nested sub-config built with no implementation yields `None`.
  • src/transformers/modeling_utils.py:1848-1849 base β€” `get_correct_attn_implementation` accepts `str | None` and maps `None` to "sdpa", so the `None` is absorbed downstream rather than being impossible.
minorConventionsrc/transformers/modeling_utils.py:2116graph behind

`set_experts_implementation` still walks one level of `sub_configs`, so the sibling operation keeps the nested-config defect this pull fixes for attention.

  • src/transformers/modeling_utils.py:2126-2133 base β€” `set_experts_implementation` resolves a dict key only against `self.config.sub_configs` and `getattr(self.config, subconfig_key) is submodule.config` β€” the one-level match this pull replaced for attention; `get_experts_implementation` (2072-2075) reads back the same level.
  • src/transformers/configuration_utils.py:484-500 base β€” The config-side `_experts_implementation` setter is already recursive (`value.get(subconfig_key, current_subconfig_moe)` then `subconfig._experts_implementation = sub_implementation`), so construction-time dispatch reaches nested experts configs while the runtime setter does not β€” the identical asymmetry being fixed here.
  • src/transformers/modeling_utils.py:1761-1780 base β€” `get_experts_implementation` is documented as the counterpart of `set_experts_implementation`, to snapshot and restore β€” the restore is what a nested-aware setter has to apply.

How the analysis was done

Changes, file by file

src/transformers/modeling_utils.py+85 βˆ’35 Β· modified Β· 6 findings

Analysis of this file

## Closing Plan vs. what the fold found. The change is two fixes plus a reshape: recursive config→target resolution (H1), the map wired into the submodule loop and a new recursive dark-set walk over the whole tree (H2/H3), and flag cleanup everywhere. I checked it against the peer the PR itself names (the recursive _attn_implementation setter in configuration_utils.py), against every non-test caller (continuous_batching/continuous_api.py:779/855, quantizer_gguf.py:165), against the shipped nested trees (deepseek_ocr2, qwen3_omni_moe, sam2/sam3, instructblipvideo, colmodernvbert, gemma4_unified_assistant), and against the generic suite in tests/test_modeling_common.py / tests/utils/test_modeling_utils.py. Findings filed (9). Two majors, both in the new wiring: (1) target_implementations is computed from the raw attn_implementation instead of the adjusted requested_implementation, so the kernel fallback, the paged| strip and the sdpa→eager downgrade stop reaching submodules — the shipped test that pins the kernel fallback for a second call now describes the root only; (2) a nested "" key is silently dropped by the is_root guard, so {"vision_config": {"": "eager"}} sets the branch to one thing at construction and another at runtime. Then four minors on the resolver's new semantics: the docstring (2.008) still says "keys are the sub_configs name" with no depth rule; the dict[int, str] annotation admits None; the unguarded recursion turns an aliased/cyclic sub_configs graph into a RecursionError; and a flat key now binds at any depth, so qwen3_omni_moe's text_config (declared in both thinker and talker) is selected twice by one key. Concerns the analysis raised — disposition. - *Nested "" keys*: filed (bug, 1974). - *Flat-key collision across branches*: filed (compatibility, 1998) — confirmed real on the shipped qwen3_omni_moe tree rather than theoretical. - *Unvalidated values / non-config sub_configs entries*: refuted for shipped code — every sub_configs value in the survey is a config class and getattr(cfg, key) returns instances; colmodernvbert's PreTrainedConfig-valued entry resolves to an instance carrying _attn_implementation. Left as open-risk, not filed: a non-config value yields None, which get_correct_attn_implementation absorbs as "sdpa" (read at 1848-1849) rather than crashing. - *Nested dark-set validation now raising where base did not*: open. quantizer_gguf.py:165 already wraps the call in try/except and logs, so the blast radius is a changed warning message; I did not read a model whose nested sub-config would newly fail validation. - *_attn_was_changed leaking on nested configs (the reported defect)*: refuted as still broken — _clean_attn_was_changed (2107-2114) recurses from self.config, and no other file in the repo produces or reads the flag (shakedown: 5 hits, all in this file), so the invariant stands. - *changed_configs redundancy*: not filed. It is load-bearing only for a submodule config outside the root's sub_configs tree, which is exactly the "badly designed model" case the neighbouring comment names; overlap on the common path is harmless duplicated deletion. - *Cost / extra walk*: not filed — one O(#configs) walk and an id-keyed dict, released at exit; the added warnings are one per unmatched nested sub-config, and I found no test pinning log counts around this method. - *set_experts_implementation left one-level*: filed (convention, 2116), citing both the sibling in this file (2126-2133) and the already-recursive config-side experts setter (configuration_utils.py 484-500). - *Docs*: docs/source/en/attention_interface.md documents the flat and "" forms (47-128) and is now under-specified for nested use — covered by the docstring finding rather than filed twice. Tests. F2's four parts are the regression harness for the reported defect; I filed the two behaviours they omit (nested "", a name reused at two depths). tests/utils/test_modeling_utils.py:3335-3337 and test_modeling_common.py:4783-4817 encode the old contract and should be re-read against finding #2. Graph drift. modeling_utils.py moved 19 commits past the base; all base spans were resolved by symbol, never by hunk line number. Not reached. I did not open a per-model test for any other nested architecture (qwen3_omni_moe, sam3, instructblipvideo), so the two-branch flat-key consequence is argued from the config declarations alone; and I did not read the vision/text model files that hold the nested submodules, so the None-target path is reasoned from get_correct_attn_implementation, not exercised.
F1.H14 findings
def _can_set_experts_implementation(cls) -> bool:
19611961 cls._can_set_experts_implementation_cached_value = can_set
19621962 return can_set
19631963
1964+ def _get_subconfig_target_implementations(self, attn_implementation: str | dict) -> dict[int, str]:
1965+ mapping = {}
1966+
1967+ def _resolve(cfg, spec, inherited_impl, is_root=False):
1968+ current_attn = getattr(cfg, "_attn_implementation", None)
1969+ if isinstance(spec, str):
1970+ target = spec
1971+ child_inherited = spec
1972+ child_dict = None
1973+ elif isinstance(spec, dict):
1974+ if is_root:
1975+ target = spec.get("", current_attn)
1976+ else:
1977+ target = inherited_impl if inherited_impl is not None else current_attn
1978+ child_inherited = inherited_impl
1979+ child_dict = spec
1980+ else:
1981+ target = inherited_impl if inherited_impl is not None else current_attn
1982+ child_inherited = inherited_impl
1983+ child_dict = None
1984+
1985+ mapping[id(cfg)] = target
1986+
1987+ for subconfig_key in getattr(cfg, "sub_configs", {}):
1988+ subconfig = getattr(cfg, subconfig_key, None)
1989+ if subconfig is not None:
1990+ if child_dict is not None and subconfig_key in child_dict:
1991+ sub_spec = child_dict[subconfig_key]
1992+ _resolve(
1993+ subconfig,
1994+ sub_spec,
1995+ sub_spec if isinstance(sub_spec, str) else child_inherited,
1996+ is_root=False,
1997+ )
1998+ elif child_dict is not None:
1999+ _resolve(subconfig, child_dict, child_inherited, is_root=False)
2000+ elif isinstance(spec, str):
2001+ _resolve(subconfig, spec, child_inherited, is_root=False)
2002+ else:
2003+ _resolve(subconfig, None, child_inherited, is_root=False)
2004+
2005+ _resolve(self.config, attn_implementation, None, is_root=True)
2006+ return mapping
2007+
19642008 def set_attn_implementation(self, attn_implementation: str | dict, allow_all_kernels: bool = False):
19652009 """
19662010 Set the requested `attn_implementation` for this model.
majorBugsrc/transformers/modeling_utils.py:1974–1979graph behind

Nested `""` keys are ignored: `spec.get("", ...)` runs only when `is_root`, so `{"vision_config": {"": "eager"}}` leaves that branch unchanged while construction sets it eager.

  • src/transformers/modeling_utils.py:1974-1979 this change β€” `if is_root: target = spec.get("", current_attn) else: target = inherited_impl if ... ` β€” a nested dict's own `""` key is never read, so the branch is left at `inherited_impl`/its current value.
  • src/transformers/configuration_utils.py:462-478 base β€” The construction-time setter the PR names as its reference does `value.get("", current_attn)` at EVERY level and re-enters itself via `subconfig._attn_implementation = sub_implementation`, so a nested `""` IS honoured there.
minorCompatibilitysrc/transformers/modeling_utils.py:1998graph behind

Flat keys now bind at any depth, so `text_config`, reused at two levels in the shipped qwen3_omni_moe tree, is applied to both branches by one key.

  • src/transformers/modeling_utils.py:1998-2003 this change β€” `elif child_dict is not None: _resolve(subconfig, child_dict, child_inherited, is_root=False)` hands the same root dict to every descendant, so a key naming a deeper sub-config is honoured wherever it appears.
  • src/transformers/modeling_utils.py:2019-2025 base β€” Base matched `for subconfig_key in self.config.sub_configs` against `getattr(self.config, subconfig_key) is submodule.config`, so a flat key could only ever name a direct child of the root.
  • src/transformers/models/qwen3_omni_moe/configuration_qwen3_omni_moe.py:219-223 base β€” `Qwen3OmniMoeThinkerConfig.sub_configs` declares `text_config`.
  • src/transformers/models/qwen3_omni_moe/configuration_qwen3_omni_moe.py:470-473 base β€” `Qwen3OmniMoeTalkerConfig.sub_configs` declares `text_config` again, one level deeper than the root β€” a single flat key now selects both branches.
  • tests/models/deepseek_ocr2/test_modeling_deepseek_ocr2.py:154-180 this change β€” The four parts cover a branch-named string, a nested dict, a flat key and flag cleanup; the nested `""` form and a name reused at two depths (qwen3_omni_moe's `text_config`) are absent, so the two cases the resolver disagrees on are untested.
  • src/transformers/modeling_utils.py:1967-2003 this change β€” The resolver drops a nested `""` (`is_root` guard) and binds a flat key at any depth, so both untested inputs take a path no existing assertion pins.
minorStale textsrc/transformers/modeling_utils.py:2008graph behind

The docstring still says dict keys are "the sub_configs name"; the resolver now accepts nested dicts and binds a flat key at any depth, so the documented contract is wrong.

  • src/transformers/modeling_utils.py:1973-2003 this change β€” `spec` may be a dict, and `child_dict` is handed down every level, so a flat key such as `sam_config` now matches a sub-config name anywhere in the hierarchy.
  • src/transformers/modeling_utils.py:1969-1971 base β€” The unchanged docstring: 'a `dict` where keys are the sub_configs name, in which case each submodel will dispatch the corresponding value' β€” no nested form, no depth rule.
minorStale textsrc/transformers/modeling_utils.py:1964graph behind

Annotation `dict[int, str]` is wrong: `target` is `None` for any config without an `_attn_implementation`, and that `None` is what flows into `get_correct_attn_implementation`.

  • src/transformers/modeling_utils.py:1968-1985 this change β€” `current_attn = getattr(cfg, "_attn_implementation", None)` is the fallback `target`, and `mapping[id(cfg)] = target` stores it unchecked β€” a nested sub-config built with no implementation yields `None`.
  • src/transformers/modeling_utils.py:1848-1849 base β€” `get_correct_attn_implementation` accepts `str | None` and maps `None` to "sdpa", so the `None` is absorbed downstream rather than being impossible.
F1.H21 finding
def set_attn_implementation(self, attn_implementation: str | dict, allow_all_ker
19952039 # Apply the change (on the internal attr, to avoid setting it recursively)
19962040 self.config._attn_implementation_internal = requested_implementation
19972041
2042+ target_implementations = self._get_subconfig_target_implementations(attn_implementation)
2043+
19982044 # Apply it to all submodels as well
2045+ changed_configs = []
19992046 for submodule in self.modules():
20002047 # We found a submodel (which is not self) with a different config (otherwise, it may be the same "actual model",
20012048 # e.g. ForCausalLM has a Model inside, but no need to check it again)
majorBugsrc/transformers/modeling_utils.py:2042graph behind

`target_implementations` is built from the raw argument, so submodule configs receive the unadjusted string and the kernel and sdpa-to-eager fallbacks no longer reach them.

  • src/transformers/modeling_utils.py:2042 this change β€” `target_implementations = self._get_subconfig_target_implementations(attn_implementation)` passes the CALLER's value, not the `requested_implementation` lines 1991-1995 replaced with the adjusted one.
  • src/transformers/modeling_utils.py:2065 this change β€” `sub_implementation = target_implementations.get(id(submodule.config), requested_implementation)` β€” the mapping wins for every config reachable through `sub_configs`, so the raw value is what a submodule gets.
  • src/transformers/modeling_utils.py:1991-1995 base β€” `requested_implementation = self._check_and_adjust_attn_implementation(...)` is where the sdpa->eager downgrade, the `paged|` strip and the `_compatible_flash_implementations` correction are applied to the value submodels used to receive.
  • src/transformers/modeling_utils.py:2017 base β€” Base handed each submodule that ADJUSTED value (`sub_implementation = requested_implementation`), so a composite's submodules inherited the model-wide fallback.
  • src/transformers/modeling_utils.py:1792-1823 base β€” The kernel fallback: with `flash_attention_2` requested and the package absent the root becomes `FLASH_ATTN_KERNEL_FALLBACK[...]`; a submodule now re-resolving the raw `flash_attention_2` takes the `get_correct_attn_implementation` path instead and raises ImportError from `_flash_attn_can_dispatch`.
  • tests/utils/test_modeling_utils.py:3333-3337 base β€” The shipped test pins the intent that after `set_attn_implementation("flash_attention_2")` the value the model ends on is the KERNEL FALLBACK, not the raw requested name β€” which is now true of the root only.
F1.H31 finding
def set_attn_implementation(self, attn_implementation: str | dict, allow_all_ker
20152062 )
20162063 # Set the attn on the submodule
20172064 else:
2018- sub_implementation = requested_implementation
2019- if isinstance(attn_implementation, dict):
2020- for subconfig_key in self.config.sub_configs:
2021- # We need to check for exact object match here, with `is`
2022- if getattr(self.config, subconfig_key) is submodule.config:
2023- sub_implementation = attn_implementation.get(
2024- subconfig_key, submodule.config._attn_implementation
2025- )
2026- break
2065+ sub_implementation = target_implementations.get(id(submodule.config), requested_implementation)
20272066 # Check the module can use correctly, otherwise we raise an error if requested attention can't be set for submodule
20282067 sub_implementation = submodule.get_correct_attn_implementation(sub_implementation)
20292068 submodule.config._attn_implementation_internal = sub_implementation
20302069
20312070 # Still add it as "changed" even if it was skipped, as we would otherwise try to set it in the dark afterwards
20322071 # We need to set it on the config itself, to differentiate 2 subconfigs of the same __class__ potentially
20332072 submodule.config._attn_was_changed = True
2073+ changed_configs.append(submodule.config)
20342074
20352075 # We need this as some old and badly designed models use subconfigs without declaring the corresponding modules as PreTrainedModel
2036- for subconfig_key in self.config.sub_configs:
2037- if (subconfig := getattr(self.config, subconfig_key)) is not None:
2038- sub_implementation = (
2039- requested_implementation
2040- if not isinstance(attn_implementation, dict)
2041- else attn_implementation.get(subconfig_key, subconfig._attn_implementation)
2042- )
2043- # This means we did not perform any check above for this particular subconfig -> set it in the dark if it is registered
2044- if (
2045- not hasattr(subconfig, "_attn_was_changed")
2046- # If it's already the same, then no need to enter here and raise warnings
2047- and sub_implementation != subconfig._attn_implementation
2048- ):
2049- if sub_implementation not in ["eager"] + ALL_ATTENTION_FUNCTIONS.valid_keys():
2050- raise ValueError(
2051- f'Specified `attn_implementation="{sub_implementation}"` is not supported for {subconfig_key}. '
2052- 'The only possible arguments are "eager" (manual attention implementation)'
2053- f"or one of the following: {list(ALL_ATTENTION_FUNCTIONS.valid_keys())}"
2076+ def _apply_to_subconfigs(cfg):
2077+ for subconfig_key in getattr(cfg, "sub_configs", {}):
2078+ if (subconfig := getattr(cfg, subconfig_key)) is not None:
2079+ sub_implementation = target_implementations.get(id(subconfig), requested_implementation)
2080+ # This means we did not perform any check above for this particular subconfig -> set it in the dark if it is registered
2081+ if (
2082+ not hasattr(subconfig, "_attn_was_changed")
2083+ # If it's already the same, then no need to enter here and raise warnings
2084+ and sub_implementation != subconfig._attn_implementation
2085+ ):
2086+ if sub_implementation not in ["eager"] + ALL_ATTENTION_FUNCTIONS.valid_keys():
2087+ raise ValueError(
2088+ f'Specified `attn_implementation="{sub_implementation}"` is not supported for {subconfig_key}. '
2089+ 'The only possible arguments are "eager" (manual attention implementation)'
2090+ f"or one of the following: {list(ALL_ATTENTION_FUNCTIONS.valid_keys())}"
2091+ )
2092+ subconfig._attn_implementation_internal = sub_implementation
2093+ logger.warning(
2094+ f"We set the attention implementation for the sub-config `{subconfig_key}` to `{sub_implementation}` "
2095+ "without finding the associated sub-model. For this reason we could not check if the model supports it. "
2096+ "You may encounter undefined behavior."
20542097 )
2055- subconfig._attn_implementation_internal = sub_implementation
2056- logger.warning(
2057- f"We set the attention implementation for the sub-config `{subconfig_key}` to `{sub_implementation}` "
2058- "without finding the associated sub-model. For this reason we could not check if the model supports it. "
2059- "You may encounter undefined behavior."
2060- )
2061- # Unset the attribute in this case, to avoid issues in the future
2062- else:
2098+ _apply_to_subconfigs(subconfig)
2099+
2100+ _apply_to_subconfigs(self.config)
2101+
2102+ # Unset the attribute in all cases, to avoid issues in future calls
2103+ for cfg in changed_configs:
2104+ if hasattr(cfg, "_attn_was_changed"):
2105+ del cfg._attn_was_changed
2106+
2107+ def _clean_attn_was_changed(cfg):
2108+ for subconfig_key in getattr(cfg, "sub_configs", {}):
2109+ if (subconfig := getattr(cfg, subconfig_key)) is not None:
20632110 if hasattr(subconfig, "_attn_was_changed"):
20642111 del subconfig._attn_was_changed
2112+ _clean_attn_was_changed(subconfig)
2113+
2114+ _clean_attn_was_changed(self.config)
20652115
20662116 def get_experts_implementation(self) -> dict[str, str | None]:
20672117 """
minorConventionsrc/transformers/modeling_utils.py:2116graph behind

`set_experts_implementation` still walks one level of `sub_configs`, so the sibling operation keeps the nested-config defect this pull fixes for attention.

  • src/transformers/modeling_utils.py:2126-2133 base β€” `set_experts_implementation` resolves a dict key only against `self.config.sub_configs` and `getattr(self.config, subconfig_key) is submodule.config` β€” the one-level match this pull replaced for attention; `get_experts_implementation` (2072-2075) reads back the same level.
  • src/transformers/configuration_utils.py:484-500 base β€” The config-side `_experts_implementation` setter is already recursive (`value.get(subconfig_key, current_subconfig_moe)` then `subconfig._experts_implementation = sub_implementation`), so construction-time dispatch reaches nested experts configs while the runtime setter does not β€” the identical asymmetry being fixed here.
  • src/transformers/modeling_utils.py:1761-1780 base β€” `get_experts_implementation` is documented as the counterpart of `set_experts_implementation`, to snapshot and restore β€” the restore is what a nested-aware setter has to apply.
tests/models/deepseek_ocr2/test_modeling_deepseek_ocr2.py+33 βˆ’0 Β· modified
F2.H1
def _image_features_prepare_config_and_inputs(self):
147147 config.vision_config.hidden_size = config.vision_config.encoder_config.hidden_size
148148 return config, inputs_dict
149149
150+ def test_set_attn_implementation_nested_subconfigs(self):
151+ config = self.model_tester.get_config()
152+ model = DeepseekOcr2Model(config)
153+
154+ # 1. Propagate string implementation to nested subconfigs
155+ model.set_attn_implementation({"vision_config": "eager"})
156+ self.assertEqual(model.language_model.config._attn_implementation, "sdpa")
157+ self.assertEqual(model.vision_tower.config._attn_implementation, "eager")
158+ self.assertEqual(model.vision_tower.sam_encoder.config._attn_implementation, "eager")
159+ self.assertEqual(model.vision_tower.vision_encoder.config._attn_implementation, "eager")
160+
161+ # 2. Targeted nested subconfig via nested dict
162+ model.set_attn_implementation("sdpa")
163+ model.set_attn_implementation({"vision_config": {"sam_config": "eager"}})
164+ self.assertEqual(model.vision_tower.config._attn_implementation, "sdpa")
165+ self.assertEqual(model.vision_tower.sam_encoder.config._attn_implementation, "eager")
166+ self.assertEqual(model.vision_tower.vision_encoder.config._attn_implementation, "sdpa")
167+
168+ # 3. Targeted nested subconfig via flat dict key
169+ model.set_attn_implementation("sdpa")
170+ model.set_attn_implementation({"sam_config": "eager"})
171+ self.assertEqual(model.vision_tower.config._attn_implementation, "sdpa")
172+ self.assertEqual(model.vision_tower.sam_encoder.config._attn_implementation, "eager")
173+ self.assertEqual(model.vision_tower.vision_encoder.config._attn_implementation, "sdpa")
174+
175+ # 4. Consecutive calls work properly and do not leave _attn_was_changed behind
176+ model.set_attn_implementation("eager")
177+ self.assertEqual(model.vision_tower.sam_encoder.config._attn_implementation, "eager")
178+ model.set_attn_implementation("sdpa")
179+ self.assertEqual(model.vision_tower.sam_encoder.config._attn_implementation, "sdpa")
180+ self.assertFalse(hasattr(model.vision_tower.sam_encoder.config, "_attn_was_changed"))
181+ self.assertFalse(hasattr(model.vision_tower.config, "_attn_was_changed"))
182+
150183
151184@require_torch
152185class DeepseekOcr2IntegrationTest(unittest.TestCase):