如何重置HuggingFace模型generation_config默认值以消除UserWarning
TrOCR模型警告处理与配置调整方案
环境与问题背景
在Windows 11 Pro、Python 3.8.1环境下运行以下TrOCR模型代码时,触发了UserWarning:
原运行代码
from transformers import TrOCRProcessor, VisionEncoderDecoderModel from PIL import Image import requests # 加载IAM数据库中的图片 url = 'https://fki.tic.heia-fr.ch/static/img/a01-122-02-00.jpg' image = Image.open(requests.get(url, stream=True).raw).convert("RGB") # 加载模型 processor = TrOCRProcessor.from_pretrained('fhswf/TrOCR_german_handwritten') model = VisionEncoderDecoderModel.from_pretrained('fhswf/TrOCR_german_handwritten') # print(model.generation_config) # 处理图片并生成文本 pixel_values = processor(images=image, return_tensors="pt").pixel_values generated_ids = model.generate(pixel_values) generated_text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0] print(generated_text)
触发的警告信息
C:\Users\user\Projects\Project\venv\lib\site-packages\transformers\generation\utils.py:1376: UserWarning: You have modified the pretrained model configuration to control generation. This is a deprecated strategy to control generation and will be removed soon, in a future version. Please use and modify the model generation configuration (see https://huggingface.co/docs/transformers/generation_strategies#default-text-generation-configuration )
模型当前的generation_config
查看模型生成配置,发现以下参数与默认配置不同:
GenerationConfig { "bos_token_id": 0, "decoder_start_token_id": 2, "eos_token_id": 2, "pad_token_id": 1, "use_cache": false }
用户计划将model.generation_config恢复为默认值,同时在generate函数中传入这些参数以保留原模型性能,提出两个问题:
- 如何恢复模型配置为默认值?
- 该策略能否消除警告且不影响模型原有性能?
问题解答
1. 恢复模型配置为默认值的方法
可以通过加载VisionEncoderDecoderModel对应的默认生成配置来覆盖当前配置,代码如下:
from transformers import GenerationConfig # 将模型的生成配置恢复为默认值 model.generation_config = GenerationConfig.from_model_config(model.config)
这个方法会基于模型的基础配置生成官方默认的生成规则,确保配置回归初始状态。
2. 该策略的效果
- 能消除警告:警告的根源是模型自带的生成配置被修改(与默认不一致),属于官方弃用的旧策略。恢复默认配置后,在
generate()函数中显式传入原参数,符合官方推荐的「通过生成函数参数控制生成逻辑」的新方式,不会再触发警告。 - 不会影响原有性能:只要在
generate()中准确传入原generation_config里的所有特殊参数,生成逻辑和原模型完全一致,输出结果、性能表现都和之前保持一致。
修改后的完整代码
from transformers import TrOCRProcessor, VisionEncoderDecoderModel, GenerationConfig from PIL import Image import requests # 加载图片 url = 'https://fki.tic.heia-fr.ch/static/img/a01-122-02-00.jpg' image = Image.open(requests.get(url, stream=True).raw).convert("RGB") # 加载模型和处理器 processor = TrOCRProcessor.from_pretrained('fhswf/TrOCR_german_handwritten') model = VisionEncoderDecoderModel.from_pretrained('fhswf/TrOCR_german_handwritten') # 恢复模型生成配置为默认值 model.generation_config = GenerationConfig.from_model_config(model.config) # 处理图片并生成文本,显式传入原配置参数 pixel_values = processor(images=image, return_tensors="pt").pixel_values generated_ids = model.generate( pixel_values, bos_token_id=0, decoder_start_token_id=2, eos_token_id=2, pad_token_id=1, use_cache=False ) generated_text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0] print(generated_text)
内容的提问来源于stack exchange,提问作者user26598303
相关产品推荐
相关产品推荐

