Kaggle加载Hugging Face预训练BERT模型失败求助
修复Kaggle加载gaunernst/bert-small-uncased模型的AttributeError问题
1. 解决缓存路径与权限问题
Kaggle默认缓存目录可能存在读写限制,导致模型文件下载不完整或读取失败。手动指定可写的缓存目录并清理残留缓存:
import os import shutil from transformers import BertTokenizer, BertModel # 创建自定义缓存目录 cache_dir = "/kaggle/working/bert_cache" os.makedirs(cache_dir, exist_ok=True) # 清理旧缓存(可选) shutil.rmtree(cache_dir, ignore_errors=True) # 加载Tokenizer与模型,指定缓存目录 tokenizer = BertTokenizer.from_pretrained("gaunernst/bert-small-uncased", cache_dir=cache_dir) model = BertModel.from_pretrained("gaunernst/bert-small-uncased", cache_dir=cache_dir)
2. 确保自定义模型类的可访问性
如果使用了自定义封装的BERT模型(比如添加了自定义分类头),必须在加载模型前完整定义该类,避免反序列化时找不到类定义:
from transformers import BertPreTrainedModel, BertModel import torch.nn as nn class CustomBERTClassifier(BertPreTrainedModel): def __init__(self, config): super().__init__(config) self.bert = BertModel(config) self.classifier = nn.Linear(config.hidden_size, 2) # 示例二分类头 def forward(self, input_ids, attention_mask=None): bert_output = self.bert(input_ids, attention_mask=attention_mask) logits = self.classifier(bert_output.pooler_output) return logits # 加载自定义模型,确保类已提前定义 model = CustomBERTClassifier.from_pretrained("gaunernst/bert-small-uncased", cache_dir=cache_dir)
若模型依赖仓库中的自定义代码,需添加trust_remote_code=True参数:
model = BertModel.from_pretrained("gaunernst/bert-small-uncased", trust_remote_code=True, cache_dir=cache_dir)
3. 对齐transformers与torch版本
Kaggle默认环境的torch版本可能与指定的transformers版本不兼容,手动安装匹配版本:
# 针对transformers 4.37.0,安装兼容的torch版本 !pip install --upgrade torch==2.1.0 transformers==4.37.0 # 针对transformers 4.31.0,安装兼容的torch版本 !pip install --upgrade torch==2.0.1 transformers==4.31.0
4. 修正GridSearch中的模型实例化逻辑
如果GridSearch中动态创建模型,需将模型加载逻辑封装在estimator的初始化方法内,避免全局作用域下重复加载导致的状态异常:
from sklearn.base import BaseEstimator, ClassifierMixin from sklearn.model_selection import GridSearchCV class BERTGridEstimator(BaseEstimator, ClassifierMixin): def __init__(self, model_name="gaunernst/bert-small-uncased", cache_dir="/kaggle/working/bert_cache"): self.model_name = model_name self.cache_dir = cache_dir # 初始化时加载模型 self.model = BertModel.from_pretrained(self.model_name, cache_dir=self.cache_dir) # 其他训练相关初始化... # 初始化GridSearch param_grid = {"some_hyperparam": [0.001, 0.01]} estimator = BERTGridEstimator() grid_search = GridSearchCV(estimator, param_grid, cv=3)
内容的提问来源于stack exchange,提问作者F. M.
相关产品推荐
相关产品推荐

