如何从Google Cloud Storage加载Hugging Face基础模型(无需本地存储)
无需本地持久化从GCS加载XLM-RoBERTa基础模型
核心思路
借助Hugging Face模型加载接口对内存中配置/权重文件的支持,结合Google Cloud Storage(GCS)客户端直接将Blob内容读取到内存,完全绕开本地文件持久化操作。
具体实现步骤
1. 安装必要依赖
如果尚未安装相关库,执行以下命令:
pip install torch transformers google-cloud-storage
2. 编写GCS内存读取工具函数
实现一个函数,从指定GCS存储桶读取目标Blob内容,返回内存字节流对象:
from google.cloud import storage from io import BytesIO def read_gcs_blob_to_memory(bucket_name, blob_path): storage_client = storage.Client() bucket = storage_client.bucket(bucket_name) blob = bucket.blob(blob_path) byte_stream = BytesIO() blob.download_to_file(byte_stream) byte_stream.seek(0) # 重置文件指针至起始位置 return byte_stream
3. 加载XLM-RoBERTa基础模型
分三步完成模型加载,全程在内存中操作:
from transformers import XLMRobertaConfig, XLMRobertaModel import torch # 替换为你的GCS存储桶名称和模型文件前缀 BUCKET_NAME = "your-gcs-bucket" MODEL_PREFIX = "path/to/xlm-roberta-base/" # 1. 加载模型配置 config_stream = read_gcs_blob_to_memory(BUCKET_NAME, f"{MODEL_PREFIX}config.json") config = XLMRobertaConfig.from_json_file(config_stream) # 2. 加载模型权重 weights_stream = read_gcs_blob_to_memory(BUCKET_NAME, f"{MODEL_PREFIX}pytorch_model.bin") state_dict = torch.load(weights_stream, map_location="cpu") # 可根据需求指定设备,如"cuda" # 3. 初始化模型并加载权重 base_model = XLMRobertaModel(config) base_model.load_state_dict(state_dict) # 验证加载完成,设置为评估模式 base_model.eval()
4. 基于基础模型构建自定义PyTorch模型
直接在加载好的基础模型上叠加自定义结构即可:
import torch.nn as nn class CustomModel(nn.Module): def __init__(self, base_model, num_classes=10): super().__init__() self.base_model = base_model self.classifier = nn.Linear(base_model.config.hidden_size, num_classes) def forward(self, input_ids, attention_mask=None): outputs = self.base_model(input_ids=input_ids, attention_mask=attention_mask) pooled_output = outputs.last_hidden_state[:, 0, :] # 取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]> token输出 logits = self.classifier(pooled_output) return logits # 初始化自定义模型 custom_model = CustomModel(base_model)
关键注意事项
- 权限配置:运行代码的环境(本地/GCE/GCF等)需要具备目标GCS存储桶的读取权限,可通过服务账号密钥或默认应用凭据配置。
- 分片权重处理:如果模型权重分为多个分片文件(如
pytorch_model-00001-of-00002.bin),需循环读取所有分片并合并state_dict:from collections import OrderedDict def load_sharded_weights(bucket_name, prefix): storage_client = storage.Client() bucket = storage_client.bucket(bucket_name) blobs = bucket.list_blobs(prefix=prefix) state_dict = OrderedDict() for blob in blobs: if blob.name.endswith(".bin") and "pytorch_model-" in blob.name: byte_stream = BytesIO() blob.download_to_file(byte_stream) byte_stream.seek(0) shard_state_dict = torch.load(byte_stream, map_location="cpu") state_dict.update(shard_state_dict) return state_dict # 使用分片加载函数替代单文件权重加载 state_dict = load_sharded_weights(BUCKET_NAME, MODEL_PREFIX)
内容的提问来源于stack exchange,提问作者Ravi156
相关产品推荐
相关产品推荐

