在Amazon SageMaker训练Llama3.1-8B-Instruct时遇rope_scaling配置错误
问题描述
在Amazon SageMaker上训练Llama3.1-8B-Instruct模型时任务失败,报错信息如下:
./usr/local/lib/python3.10/dist-packages/huggingface_hub/file_download.py:1150: FutureWarning: `resume_download` is deprecated and will be removed in version 1.0.0. Downloads always resume when possible. If you want to force a new download, use `force_download=True`. warnings.warn( Traceback (most recent call last): File "/workspace/train.py", line 85, in <module> main() File "/workspace/train.py", line 48, in main config = AutoConfig.from_pretrained(model_name, token=use_auth_token) File "/usr/local/lib/python3.10/dist-packages/transformers/models/auto/configuration_auto.py", line 1124, in from_pretrained return config_class.from_dict(config_dict, **unused_kwargs) File "/usr/local/lib/python3.10/dist-packages/transformers/configuration_utils.py", line 764, in from_dict config = cls(**config_dict) File "/usr/local/lib/python3.10/dist-packages/transformers/models/llama/configuration_llama.py", line 160, in __init__ self._rope_scaling_validation() File "/usr/local/lib/python3.10/dist-packages/transformers/models/llama/configuration_llama.py", line 180, in _rope_scaling_validation raise ValueError( ValueError: `rope_scaling` must be a dictionary with with two fields, `type` and `factor`, got {'factor': 8.0, 'low_freq_factor': 1.0, 'high_freq_factor': 4.0, 'original_max_position_embeddings': 8192, 'rope_type': 'llama3'}
尝试修改config.rope_scaling后仍未解决问题,修改代码如下:
# Load model configuration config = AutoConfig.from_pretrained(model_name, token=use_auth_token) # Modify the rope_scaling config config.rope_scaling = { "type": "llama3", "factor": 8.0 } # Initialize the model with the modified config model = LlamaForCausalLM.from_pretrained(modek_name, token=use_auth_token, config=config)
解决方案
问题核心是加载原始配置时就触发了校验逻辑,修改代码根本没执行到后续步骤。正确做法是在加载配置阶段就直接覆盖rope_scaling参数:
# 加载配置时直接指定合法的rope_scaling,跳过原始错误配置的校验 config = AutoConfig.from_pretrained( model_name, token=use_auth_token, rope_scaling={"type": "llama3", "factor": 8.0} ) # 初始化模型,注意修正拼写错误:modek_name → model_name model = LlamaForCausalLM.from_pretrained( model_name, token=use_auth_token, config=config )
额外注意两个细节:
- 原代码存在拼写错误:
modek_name需改为model_name,否则会导致模型加载失败 - 确保
transformers版本支持Llama3的ROPE类型,建议升级到最新稳定版:
pip install --upgrade transformers
内容的提问来源于stack exchange,提问作者MrYammers
相关产品推荐
相关产品推荐

