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

如何从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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 12:11:03