如何不使用默认包装函数下载并加载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
相关产品推荐
相关产品推荐

