| 2015 | 2062 | | ) |
| 2016 | 2063 | | # Set the attn on the submodule |
| 2017 | 2064 | | 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) |
| 2027 | 2066 | | # Check the module can use correctly, otherwise we raise an error if requested attention can't be set for submodule |
| 2028 | 2067 | | sub_implementation = submodule.get_correct_attn_implementation(sub_implementation) |
| 2029 | 2068 | | submodule.config._attn_implementation_internal = sub_implementation |
| 2030 | 2069 | | |
| 2031 | 2070 | | # Still add it as "changed" even if it was skipped, as we would otherwise try to set it in the dark afterwards |
| 2032 | 2071 | | # We need to set it on the config itself, to differentiate 2 subconfigs of the same __class__ potentially |
| 2033 | 2072 | | submodule.config._attn_was_changed = True |
| 2073 | + | changed_configs.append(submodule.config) |
| 2034 | 2074 | | |
| 2035 | 2075 | | # 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." |
| 2054 | 2097 | | ) |
| 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: |
| 2063 | 2110 | | if hasattr(subconfig, "_attn_was_changed"): |
| 2064 | 2111 | | del subconfig._attn_was_changed |
| 2112 | + | _clean_attn_was_changed(subconfig) |
| 2113 | + | |
| 2114 | + | _clean_attn_was_changed(self.config) |
| 2065 | 2115 | | |
| 2066 | 2116 | | def get_experts_implementation(self) -> dict[str, str | None]: |
| 2067 | 2117 | | """ |