如何为Unbabel/COMET指定本地xlm-roberta-large分词器路径
离线环境下让COMET使用本地xlm-roberta-large分词器的解决方案
问题背景
在无法下载模型的离线环境中,已将xlm-roberta-large完整模型文件放置在HF_HOME目录/opt/huggingface/xlm-roberta-large下,能成功加载wmt22-comet-da模型,但COMET仍尝试从huggingface.co远程拉取xlm-roberta-large分词器文件,导致代理连接超时错误。需要指定本地分词器路径,避免远程请求。
解决方案
方法1:加载模型后手动替换本地编码器
直接初始化本地XLMR编码器,替换COMET模型中原有的编码器:
from comet import load_from_checkpoint from comet.encoders.xlmr import XLMREncoder # 加载COMET预训练模型 model = load_from_checkpoint("third_party/Unbabel/wmt22-comet-da/checkpoints/model.ckpt") # 指定本地xlm-roberta-large路径,初始化编码器(无需重复加载权重) local_xlmr_path = "/opt/huggingface/xlm-roberta-large" local_encoder = XLMREncoder(local_xlmr_path, load_pretrained_weights=False) # 替换模型的编码器 model.encoder = local_encoder # 正常执行预测 data = [ { "src": "这是个句子。", "mt": "This is a sentence.", "ref": "It is a sentence." }, { "src": "这是另一个句子。", "mt": "This is another sentence.", "ref": "It is another sentence." } ] model_output = model.predict(data, batch_size=8)
方法2:加载模型时覆盖编码器路径参数
通过override_kwargs直接修改模型的encoder_model参数为本地路径,让COMET初始化时直接使用本地文件:
from comet import load_from_checkpoint # 加载模型时指定本地xlm-roberta路径 model = load_from_checkpoint( "third_party/Unbabel/wmt22-comet-da/checkpoints/model.ckpt", override_kwargs={"encoder_model": "/opt/huggingface/xlm-roberta-large"} ) # 正常执行预测 data = [ { "src": "这是个句子。", "mt": "This is a sentence.", "ref": "It is a sentence." }, { "src": "这是另一个句子。", "mt": "This is another sentence.", "ref": "It is another sentence." } ] model_output = model.predict(data, batch_size=8)
注意事项
- 确保本地
/opt/huggingface/xlm-roberta-large目录下包含完整的分词器文件:tokenizer_config.json、vocab.json、merges.txt、special_tokens_map.json等 - 若使用方法2,需确认COMET版本支持通过
override_kwargs修改encoder_model参数(unbabel-comet>=2.0版本均支持)
问题详情
运行代码
from comet import load_from_checkpoint model = load_from_checkpoint("third_party/Unbabel/wmt22-comet-da/checkpoints/model.ckpt") data = [ { "src": "这是个句子。", "mt": "This is a sentence.", "ref": "It is a sentence." }, { "src": "这是另一个句子。", "mt": "This is another sentence.", "ref": "It is another sentence." } ] model_output = model.predict(data, batch_size=8)
报错信息
/home/usr1/.local/lib/python3.10/site-packages/torchvision/io/image.py:13: UserWarning: Failed to load image Python extension: '/home/usr1/.local/lib/python3.10/site-packages/torchvision/image.so: undefined symbol: _ZN3c1017RegisterOperatorsD1Ev'If you don't plan on using image functionality from `torchvision.io`, you can ignore this warning. Otherwise, there might be something wrong with your environment. Did you have `libjpeg` or `libpng` installed before building `torchvision` from source? warn( Lightning automatically upgraded your loaded checkpoint from v1.8.3.post1 to v2.4.0. To apply the upgrade to your files permanently, run `python -m pytorch_lightning.utilities.upgrade_checkpoint third_party/Unbabel/wmt22-comet-da/checkpoints/model.ckpt` Traceback (most recent call last): File "/home/usr1/.local/lib/python3.10/site-packages/urllib3/connectionpool.py", line 712, in urlopen self._prepare_proxy(conn) File "/home/usr1/.local/lib/python3.10/site-packages/urllib3/connectionpool.py", line 1014, in _prepare_proxy conn.connect() File "/home/usr1/.local/lib/python3.10/site-packages/urllib3/connection.py", line 374, in connect self._tunnel() File "/usr/lib/python3.10/http/client.py", line 925, in _tunnel raise OSError(f"Tunnel connection failed: {code} {message.strip()}") OSError: Tunnel connection failed: 504 Gateway Time-out During handling of the above exception, another exception occurred: Traceback (most recent call last): File "/home/usr1/.local/lib/python3.10/site-packages/requests/adapters.py", line 667, in send resp = conn.urlopen( File "/home/usr1/.local/lib/python3.10/site-packages/urllib3/connectionpool.py", line 801, in urlopen retries = retries.increment( File "/home/usr1/.local/lib/python3.10/site-packages/urllib3/util/retry.py", line 594, in increment raise MaxRetryError(_pool, url, error or ResponseError(cause)) urllib3.exceptions.MaxRetryError: HTTPSConnectionPool(host='huggingface.co', port=443): Max retries exceeded with url: /xlm-roberta-large/resolve/main/tokenizer_config.json (Caused by ProxyError('Cannot connect to proxy.', OSError('Tunnel connection failed: 504 Gateway Time-out'))) During handling of the above exception, another exception occurred: Traceback (most recent call last): File "/home/usr1/MIR-DEV/mtmc/comet_test.py", line 3, in <module> model = load_from_checkpoint("third_party/Unbabel/wmt22-comet-da/checkpoints/model.ckpt") File "/home/usr1/.local/lib/python3.10/site-packages/comet/models/__init__.py", line 88, in load_from_checkpoint model = model_class.load_from_checkpoint( File "/home/usr1/.local/lib/python3.10/site-packages/pytorch_lightning/utilities/model_helpers.py", line 125, in wrapper return self.method(cls, *args, **kwargs) File "/home/usr1/.local/lib/python3.10/site-packages/pytorch_lightning/core/module.py", line 1582, in load_from_checkpoint loaded = _load_from_checkpoint( File "/home/usr1/.local/lib/python3.10/site-packages/pytorch_lightning/core/saving.py", line 91, in _load_from_checkpoint model = _load_state(cls, checkpoint, strict=strict, **kwargs) File "/home/usr1/.local/lib/python3.10/site-packages/pytorch_lightning/core/saving.py", line 165, in _load_state obj = instantiator(cls, _cls_kwargs) if instantiator else cls(**_cls_kwargs) File "/home/usr1/.local/lib/python3.10/site-packages/comet/models/regression/regression_metric.py", line 98, in __init__ super().__init__( File "/home/usr1/.local/lib/python3.10/site-packages/comet/models/base.py", line 119, in __init__ self.encoder = str2encoder[self.hparams.encoder_model].from_pretrained( File "/home/usr1/.local/lib/python3.10/site-packages/comet/encoders/xlmr.py", line 78, in from_pretrained return XLMREncoder(pretrained_model, load_pretrained_weights) File "/home/usr1/.local/lib/python3.10/site-packages/comet/encoders/xlmr.py", line 42, in __init__ self.tokenizer = XLMRobertaTokenizerFast.from_pretrained(pretrained_model) File "/home/usr1/.local/lib/python3.10/site-packages/transformers/tokenization_utils_base.py", line 2190, in from_pretrained resolved_config_file = cached_file( File "/home/usr1/.local/lib/python3.10/site-packages/transformers/utils/hub.py", line 402, in cached_file resolved_file = hf_hub_download( File "/home/usr1/.local/lib/python3.10/site-packages/huggingface_hub/utils/_deprecation.py", line 101, in inner_f return f(*args, **kwargs) File "/home/usr1/.local/lib/python3.10/site-packages/huggingface_hub/utils/_validators.py", line 114, in _inner_fn return fn(*args, **kwargs) File "/home/usr1/.local/lib/python3.10/site-packages/huggingface_hub/file_download.py", line 1232, in hf_hub_download return _hf_hub_download_to_cache_dir( File "/home/usr1/.local/lib/python3.10/site-packages/huggingface_hub/file_download.py", line 1295, in _hf_hub_download_to_cache_dir (url_to_download, etag, commit_hash, expected_size, head_call_error) = _get_metadata_or_catch_error( File "/home/usr1/.local/lib/python3.10/site-packages/huggingface_hub/file_download.py", line 1746, in _get_metadata_or_catch_error metadata = get_hf_file_metadata( File "/home/usr1/.local/lib/python3.10/site-packages/huggingface_hub/utils/_validators.py", line 114, in _inner_fn return fn(*args, **kwargs) File "/home/usr1/.local/lib/python3.10/site-packages/huggingface_hub/file_download.py", line 1666, in get_hf_file_metadata r = _request_wrapper( File "/home/usr1/.local/lib/python3.10/site-packages/huggingface_hub/file_download.py", line 364, in _request_wrapper response = _request_wrapper( File "/home/usr1/.local/lib/python3.10/site-packages/huggingface_hub/file_download.py", line 387, in _request_wrapper response = get_session().request(method=method, url=url, **params) File "/home/usr1/.local/lib/python3.10/site-packages/requests/sessions.py", line 589, in request resp = self.send(prep, **send_kwargs) File "/home/usr1/.local/lib/python3.10/site-packages/requests/sessions.py", line 703, in send r = adapter.send(request, **kwargs) File "/home/usr1/.local/lib/python3.10/site-packages/huggingface_hub/utils/_http.py", line 93, in send return super().send(request, *args, **kwargs) File "/home/usr1/.local/lib/python3.10/site-packages/requests/adapters.py", line 694, in send raise ProxyError(e, request=request) requests.exceptions.ProxyError: (MaxRetryError("HTTPSConnectionPool(host='huggingface.co', port=443): Max retries exceeded with url: /xlm-roberta-large/resolve/main/tokenizer_config.json (Caused by ProxyError('Cannot connect to proxy.', OSError('Tunnel connection failed: 504 Gateway Time-out')))"), '(Request ID: ###############)')
使用的库版本
- unbabel-comet==2.2.2
- torch==2.5.0
- transformers==4.44.2
内容的提问来源于stack exchange,提问作者MrVocabulary
相关产品推荐
相关产品推荐

