Print the attention implementation being used. Args: config: The model configuration
(config: "PretrainedConfig")
| 103 | |
| 104 | |
| 105 | def print_attn_implementation(config: "PretrainedConfig") -> None: |
| 106 | """ |
| 107 | Print the attention implementation being used. |
| 108 | |
| 109 | Args: |
| 110 | config: The model configuration |
| 111 | """ |
| 112 | attn_implementation = getattr(config, "_attn_implementation", None) |
| 113 | |
| 114 | if attn_implementation == "flash_attention_2": |
| 115 | logger.info("Using FlashAttention-2 for faster training and inference.") |
| 116 | elif attn_implementation == "sdpa": |
| 117 | logger.info("Using torch SDPA for faster training and inference.") |
| 118 | else: |
| 119 | logger.info("Using vanilla attention implementation.") |
| 120 | |
| 121 | # Print sub-config attention if it's a FunAudioChat model |
| 122 | if getattr(config, "model_type", None) == "funaudiochat": |
| 123 | if hasattr(config, "audio_config"): |
| 124 | audio_attn = getattr(config.audio_config, "_attn_implementation", None) |
| 125 | if audio_attn: |
| 126 | logger.info(f" - Audio encoder attention: {audio_attn}") |
| 127 | if hasattr(config, "text_config"): |
| 128 | text_attn = getattr(config.text_config, "_attn_implementation", None) |
| 129 | if text_attn: |
| 130 | logger.info(f" - Text config attention: {text_attn}") |
| 131 | |
| 132 | |
| 133 | __all__ = ["configure_attn_implementation", "print_attn_implementation"] |
nothing calls this directly
no outgoing calls
no test coverage detected