You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何不使用默认包装函数下载并加载Unbabel COMET模型?

不依赖COMET默认函数手动下载加载模型

Unbabel COMET是一款机器翻译评分库,官方文档给出的默认加载方式如下:

from comet import download_model, load_from_checkpoint

model_path = download_model("Unbabel/wmt22-comet-da")
model = load_from_checkpoint(model_path)
data = [
    {
        "src": "Dem Feuer konnte Einhalt geboten werden",
        "mt": "The fire could be stopped",
        "ref": "They were able to control the fire."
    },
    {
        "src": "Schulen und Kindergärten wurden eröffnet.",
        "mt": "Schools and kindergartens were open",
        "ref": "Schools and kindergartens opened"
    }
]
model_output = model.predict(data, batch_size=8, gpus=1)
print(model_output)

其中download_model是huggingface_hub.snapshot_download的包装函数,实现逻辑如下:

from huggingface_hub import snapshot_download
...

def download_model(
    model: str, 
    saving_directory: Union[str, Path, None] = None
) -> str:
    model_path = snapshot_download(repo_id=model, cache_dir=saving_directory)
    checkpoint_path = os.path.join(*[model_path, "checkpoints", "model.ckpt"])
    return checkpoint_path

这个函数会返回模型检查点路径,逻辑看似直接。

从底层实现来看,COMET模型是基于PyTorch Lightning的Module类封装的,核心定义如下:

import pytorch_lightning as ptl
...

class CometModel(ptl.LightningModule, metaclass=abc.ABCMeta):
    """CometModel: Base class for all COMET models. ..."""
    def __init__(
        self,...
    ) -> None:
        super().__init__()
        self.save_hyperparameters()
        self.encoder = str2encoder[self.hparams.encoder_model].from_pretrained(
            self.hparams.pretrained_model, load_pretrained_weights
        )

另外,comet.load_from_checkpoint看起来和LightningModule.load_from_checkpoint类似,但实际上是一层封装,实现代码如下:

def load_from_checkpoint(checkpoint_path: str) -> CometModel:
    """Loads models from a checkpoint path.
    Args:
        checkpoint_path (str): Path to a model checkpoint.
    Return:
        COMET model.
    """
    checkpoint_path = Path(checkpoint_path)

    if not checkpoint_path.is_file():
        raise Exception(f"Invalid checkpoint path: {checkpoint_path}")
    
    parent_folder = checkpoint_path.parents[1] # .parent.parent
    hparams_file = parent_folder / "hparams.yaml"

    if hparams_file.is_file():
        with open(hparams_file) as yaml_file:
            hparams = yaml.load(yaml_file.read(), Loader=yaml.FullLoader)
        model_class = str2model[hparams["class_identifier"]]
        model = model_class.load_from_checkpoint(
            checkpoint_path, load_pretrained_weights=False
        )
        return model
    else:
        raise Exception(f"hparams.yaml file is missing from {parent_folder}!")

虽然这两个默认函数可以直接使用,但多层封装会模糊模型实际的存储和加载路径。

问题:是否可不使用默认的download_model和load_from_checkpoint下载加载COMET模型?

动机是明确模型存储与加载位置、防范恶意目录访问并限定COMET函数访问特定目录,需要自行指定下载路径并明确加载逻辑。


答案是可以,我们可以完全绕过这两个封装函数,手动完成模型下载和加载的全流程,具体步骤如下:

1. 手动下载模型文件

直接使用huggingface_hub.snapshot_download指定存储目录,下载完整的模型仓库:

from huggingface_hub import snapshot_download
import os
from pathlib import Path

# 指定自定义存储目录
custom_dir = "./comet_models"
# 下载模型仓库,替换为你需要的模型ID
model_repo = "Unbabel/wmt22-comet-da"

# 执行下载,指定cache_dir为自定义目录
repo_path = snapshot_download(repo_id=model_repo, cache_dir=custom_dir)
# 拼接得到检查点和配置文件路径
checkpoint_path = os.path.join(repo_path, "checkpoints", "model.ckpt")
hparams_path = os.path.join(repo_path, "hparams.yaml")

2. 手动加载模型

参照comet.load_from_checkpoint的逻辑,手动读取配置文件并加载模型:

import yaml
from comet.models import str2model

# 读取hparams配置文件
with open(hparams_path, "r") as f:
    hparams = yaml.load(f.read(), Loader=yaml.FullLoader)

# 获取对应的模型类
model_class = str2model[hparams["class_identifier"]]
# 加载模型,注意设置load_pretrained_weights=False,避免重复加载
model = model_class.load_from_checkpoint(
    checkpoint_path, 
    load_pretrained_weights=False
)

3. 验证模型功能

加载完成后可以直接使用模型进行预测,和官方流程一致:

data = [
    {
        "src": "Dem Feuer konnte Einhalt geboten werden",
        "mt": "The fire could be stopped",
        "ref": "They were able to control the fire."
    }
]
output = model.predict(data, batch_size=8)
print(output)

关键注意事项

  • 必须确保下载的模型仓库包含checkpoints/model.ckpt和hparams.yaml两个核心文件,这是COMET模型加载的必要条件。
  • 加载模型时设置load_pretrained_weights=False,因为模型检查点已经包含了预训练编码器的权重,重复加载会导致冲突。
  • 自定义目录需要有读写权限,避免下载或加载时出现权限错误。

这样操作后,你可以完全掌控模型的存储路径,避免封装函数带来的路径模糊问题,同时也能限定模型文件的存储范围,提升安全性。

内容的提问来源于stack exchange,提问作者alvas

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.26 08:07:04