safetensors 库的安装、核心语法与实战用法,安全保存模型权重

safetensors 库的安装、核心语法与实战用法,安全保存模型权重

本文是一篇 Python 库用法教程,结合 2026年09月 的热门 / 新发布 / 高频实用包,讲清安装、核心语法和能直接抄走的写法。

1. 安装与导入:pip/uv 安装 safetensors 及配套库

safetensors 是一个纯 Python 包,核心逻辑用 Rust 实现并预编译为二进制 wheel,因此安装时不需要本地 Rust 工具链。它唯一的硬性依赖只有 numpy,用于处理内存视图和 dtype 映射。根据你使用的包管理器,安装命令如下:

bash 复制代码
# 使用 pip
pip install safetensors

# 使用 uv(更快,支持 lockfile)
uv add safetensors

# 需要把 PyTorch 张量直接转换时,额外安装 torch
pip install safetensors torch

安装完成后,safetensors 会暴露三个核心入口:save_fileload_filesafe_open,以及两个异常类 SafetensorErrorsafe_open 返回的 safe_open 上下文对象。常用的导入写法是:

python 复制代码
from safetensors import safe_open, SafetensorError
from safetensors.numpy import save_file, load_file

注意这里有一个容易踩的坑:save_fileload_file 并不在 safetensors 顶层命名空间里,而是按后端拆分。safetensors.numpy 处理 np.ndarraysafetensors.torch 处理 torch.Tensorsafetensors.flaxsafetensors.paddle 等同理。如果你直接写 from safetensors import save_file,会得到 ImportError。正确做法是按你实际使用的张量框架选择子模块。

下面是一段最小可运行代码,验证安装是否成功,并观察 save_file / load_file 的签名与返回值:

python 复制代码
import numpy as np
from safetensors.numpy import save_file, load_file
from safetensors import safe_open, SafetensorError

tensors = {
    "weight": np.zeros((4, 4), dtype=np.float32),
    "bias": np.arange(4, dtype=np.float32),
}

# save_file(tensors, filename, metadata=None) -> None
save_file(tensors, "model.safetensors", metadata={"format": "pt"})

# load_file(filename, device="cpu") -> dict[str, np.ndarray]
loaded = load_file("model.safetensors")
print(loaded["weight"].shape, loaded["bias"].dtype)

# safe_open(filename, framework="numpy", device="cpu") 返回上下文管理器
with safe_open("model.safetensors", framework="numpy") as f:
    print(f.keys())                 # 列出所有张量名
    print(f.metadata())             # 读取文件级 metadata
    sub = f.get_tensor("bias")      # 按需加载单个张量
    print(sub)

关键参数说明:save_filemetadata 必须是 dict[str, str],值只能是字符串,传整数或列表会抛 SafetensorErrorload_filedevice 在 numpy 后端下只能是 "cpu",在 torch 后端可传 "cuda:0" 等;safe_openframework 取值包括 "pt""tf""flax""numpy""paddle",它不会一次性把全部权重读进内存,适合加载大模型。若文件损坏或格式不符,safe_openload_file 都会抛出 SafetensorError,捕获它即可做降级处理。

2. 核心对象与语法:save_file、load_file 与 safe_open

safetensors 的核心 API 非常精简,围绕三个对象展开:save_fileload_filesafe_open。先安装并导入:

bash 复制代码
pip install safetensors
python 复制代码
from safetensors import safe_open
from safetensors.torch import save_file, load_file

注意 save_fileload_file 位于框架子模块中(如 safetensors.torchsafetensors.numpy),而 safe_open 在顶层 safetensors 包中。保存时传入一个字典,键为字符串,值为张量,框架会从张量类型自动推断:

python 复制代码
import torch
from safetensors.torch import save_file, load_file

tensors = {
    "weight": torch.zeros((2, 2)),
    "bias": torch.ones(2),
}
save_file(tensors, "model.safetensors")

loaded = load_file("model.safetensors")
print(loaded["weight"].shape, loaded["bias"].dtype)

save_file 的签名是 save_file(tensors, filename, metadata=None)metadata 是字符串到字符串的字典,可用于记录版本、来源等信息。load_file 返回一个新的字典,键与保存时一致,值默认加载到 CPU。若需直接加载到 GPU,可传 device="cuda:0"

python 复制代码
loaded = load_file("model.safetensors", device="cuda:0")

device 参数只影响张量落点,不改变文件内容。load_file 会一次性把全部张量读入内存,适合文件较小或需要全量访问的场景。

当文件很大、只想读取其中几个键时,用 safe_open 上下文管理器实现懒加载:

python 复制代码
with safe_open("model.safetensors", framework="pt", device="cpu") as f:
    keys = f.keys()
    weight = f.get_tensor("weight")
    slice_ = f.get_slice("bias")

framework 参数指定返回张量的框架,取值 "pt"(PyTorch)、"tf"(TensorFlow)、"flax""np"(NumPy)等,必须与运行环境匹配。keys() 返回所有键;get_tensor(name) 读取单个张量;get_slice(name) 返回切片对象,可进一步用 slice_[:]slice_[0:2] 按需读取部分数据,避免整块载入。safe_open 还提供 metadata() 读取元数据。离开 with 块后文件句柄自动关闭。易错点:frameworkdevice 类型不匹配会抛错;get_slice 返回的是视图对象,需显式切片才得到张量;save_file 要求所有值形状连续,非连续张量应先 .contiguous()

3. 常用 API 详解:保存、加载、分片与元数据

save_file、load_file 与 metadata

save_file 接收一个字典(键为字符串,值为张量),并通过 metadata 参数保存任意字符串键值对。加载时用 load_file 返回张量字典,元数据需通过 safe_open 读取。

python 复制代码
import torch
from safetensors.torch import save_file, load_file, safe_open

tensors = {"weight": torch.randn(4, 4), "bias": torch.zeros(4)}
save_file(tensors, "model.safetensors", metadata={"format": "pt", "author": "demo"})

loaded = load_file("model.safetensors")
print(loaded["weight"].shape)  # torch.Size([4, 4])

with safe_open("model.safetensors", framework="pt") as f:
    print(f.metadata())          # {'format': 'pt', 'author': 'demo'}
    print(f.keys())              # ['bias', 'weight']

metadata 的值必须是字符串,传数字或列表会抛 TypeErrorsafe_open 支持惰性读取,只有调用 f.get_tensor(key) 时才真正加载对应张量。

save_model 与 load_model

save_model 直接保存 torch.nn.Moduleload_model 直接还原模型实例,省去手动提取 state_dict

python 复制代码
import torch.nn as nn
from safetensors.torch import save_model, load_model

model = nn.Linear(8, 2)
save_model(model, "linear.safetensors")
restored = load_model(model, "linear.safetensors")
print(restored.weight.shape)  # torch.Size([2, 8])

load_model 的第一个参数是目标模型实例,不是类。它会把文件中的权重加载进该实例并返回。若模型结构不匹配,会因缺失或多余键报错。

分片保存与加载

大模型单文件超过 5GB 时建议分片。用 save_file 配合分片字典,或使用 safetensors.torch.save_file_shared 机制。常见做法是手动切分:

python 复制代码
from safetensors.torch import save_file, load_file

shards = {"shard1": {"a": torch.randn(2, 2)}, "shard2": {"b": torch.randn(3, 3)}}
for name, data in shards.items():
    save_file(data, f"{name}.safetensors")

merged = {}
for name in shards:
    merged.update(load_file(f"{name}.safetensors"))
print(list(merged.keys()))  # ['a', 'b']

加载时按顺序合并即可。注意分片文件之间键名不能重复,否则合并会覆盖。

get_slice 按需读取

safe_openget_slice 允许只读取张量的一部分,避免加载整个大张量。

python 复制代码
with safe_open("model.safetensors", framework="pt") as f:
    slice_ = f.get_slice("weight")
    print(slice_.get_shape())          # [4, 4]
    part = slice_[:2, :2]              # 只取左上角 2x2
    print(part.shape)                  # torch.Size([2, 2])

get_slice 返回一个可切片的对象,支持 Python 切片语法。它适合在推理时按需加载部分权重,减少内存占用。注意切片维度必须与张量实际维度一致,否则报 IndexError

4. 完整小例子:保存并加载一个 PyTorch 模型

下面用一个最小可运行的 PyTorch 例子,把 safetensors.torch 的保存与加载串起来。先确保环境里有依赖:

bash 复制代码
pip install safetensors torch

核心思路是:模型权重通过 state_dict() 取出,交给 save_model() 写入磁盘;新模型用 load_model() 把权重灌回去,再比较两次前向输出是否一致。

python 复制代码
import torch
import torch.nn as nn
from safetensors.torch import save_model, load_model

class LinearReg(nn.Module):
    def __init__(self, in_features=4, out_features=2):
        super().__init__()
        self.fc = nn.Linear(in_features, out_features)

    def forward(self, x):
        return self.fc(x)

torch.manual_seed(0)
model = LinearReg()
x = torch.randn(3, 4)
out_before = model(x)

# 保存:直接传模型对象,而不是 state_dict
save_model(model, "linear.safetensors")

# 加载:先构造同结构的新模型,再灌入权重
new_model = LinearReg()
load_model(new_model, "linear.safetensors")
out_after = new_model(x)

print(torch.allclose(out_before, out_after))  # True

关键语法说明:

save_model(model, filename, metadata=None) 的第一个参数是 nn.Module 实例,第二个是保存路径,第三个可选 metadatadict[str, str],用于写入自定义信息(如版本号、训练步数)。它内部会调用 model.state_dict(),把每个张量按名字写入文件,因此你不需要手动 torch.save。注意文件名后缀建议用 .safetensors,虽然库不强制,但便于识别。

load_model(model, filename, strict=True)model 必须已经实例化,且结构与保存时一致;strict=True 表示键必须完全匹配,缺键或多键都会抛错。若只想加载部分权重(例如迁移学习),可设 strict=False,此时未匹配的键会被忽略,但需要你自己确认哪些层被跳过。

容易踩的坑有三点。第一,不要传 model.state_dict()save_model,它期望的是模块对象,传字典会报类型错误。第二,加载前务必先构造模型,load_model 是原地修改,不会返回新对象。第三,torch.allclose 默认容差较小,多数情况下权重逐位相等,结果为 True;若你在 GPU 上保存、CPU 上加载,需先 model.to("cpu") 再比较,避免设备不一致导致报错。

如果想验证保存内容,可以用 safetensors.safe_open 读取键名:

python 复制代码
from safetensors import safe_open

with safe_open("linear.safetensors", framework="pt") as f:
    print(list(f.keys()))  # ['fc.bias', 'fc.weight']

这样就能确认权重键与 state_dict 的键一一对应,排查加载失败时非常有用。

5. 进阶写法:内存映射、零拷贝与跨框架加载

safe_open 是 safetensors 提供的惰性读取入口,它不会一次性把整个文件读进内存,而是按需把指定 key 映射为张量。核心签名如下:

python 复制代码
from safetensors import safe_open

with safe_open(filename, framework="pt", device="cpu") as f:
    ...

关键参数说明:

  • framework:支持 "pt""tf""flax""numpy",决定返回的张量类型。
  • device"cpu""cuda:0""cpu" 配合内存映射可实现零拷贝。
  • get_tensor(key):返回对应张量,底层直接指向 mmap 区域。
  • keys():列出文件中所有张量名,便于遍历。
  • get_slice(key):返回切片对象,可在不整体加载的情况下按行读取。

device="cpu" 时,PyTorch 张量直接引用内存映射页,不做数据复制,这就是零拷贝的含义。加载到 GPU 时才发生一次 H2D 拷贝。

python 复制代码
import torch
from safetensors import safe_open

# 零拷贝:直接从磁盘映射为 CPU 张量
with safe_open("model.safetensors", framework="pt", device="cpu") as f:
    names = list(f.keys())
    print("tensors:", names)

    w = f.get_tensor("weight")          # torch.Tensor,指向 mmap
    print(w.dtype, w.shape, w.device)   # float32 torch.Size([...]) cpu

    # 按行切片,避免整体读入
    sl = f.get_slice("weight")
    first_rows = sl[:4]
    print(first_rows.shape)

# 迁移到 GPU 时才发生拷贝
gpu_w = w.to("cuda:0", non_blocking=True)

跨框架加载同样简单。若文件由 NumPy 写入,可用 framework="numpy" 读回 np.ndarray,再转成 PyTorch 张量:

python 复制代码
import numpy as np
import torch
from safetensors import safe_open
from safetensors.numpy import save_file

arr = np.random.rand(4, 8).astype("float32")
save_file({"weight": arr}, "numpy_model.safetensors")

with safe_open("numpy_model.safetensors", framework="numpy") as f:
    np_w = f.get_tensor("weight")       # np.ndarray
    pt_w = torch.from_numpy(np_w)       # 共享底层内存,无拷贝
    print(type(np_w), type(pt_w), pt_w.shape)

torch.from_numpy 与 NumPy 数组共享同一块内存,修改其中一个会影响另一个;若需要独立副本,用 torch.tensor(np_w)

常见易错点:

  • 忘记 with 上下文:离开作用域后 mmap 关闭,张量数据可能失效,务必在块内完成计算或 .clone()
  • device="cuda" 时不再零拷贝,会触发一次显式传输,需自行权衡 IO 与显存。
  • framework 与文件实际 dtype 不匹配会抛异常,读取前先用 safe_open(...).keys() 确认。
  • 多进程 DataLoader 中每个 worker 独立打开 safe_open,不要跨进程共享句柄。

结合 map_location 的写法在 safetensors 中并不直接存在,等价做法是先以 device="cpu" 零拷贝读取,再按需 .to(device),这样既保留 mmap 优势,又避免加载时立刻占满显存。

6. 注意事项:常见坑与避坑指南

save_file 只接受 Dict[str, Tensor],传入 Python 原生对象会在序列化阶段直接报错,而不是静默写入。常见错误是把配置字典、None、字符串混进张量字典:

python 复制代码
import torch
from safetensors.torch import save_file

tensors = {
    "weight": torch.randn(2, 2),
    "bias": torch.zeros(2),
    "config": {"lr": 1e-3},      # 错误:dict 不是张量
    "name": "layer0",            # 错误:str 不是张量
}
save_file(tensors, "bad.safetensors")
# TypeError: value for key 'config' must be a torch.Tensor

正确做法是把配置单独用 json 保存,只让张量进入 safetensors 文件。若确实需要保存 None,应显式跳过或用 torch.empty(0) 占位,并在加载端还原语义。

共享张量与内存布局

当多个参数共享同一块存储(如 torch.nn.Linear 权重转置、tie_weights 共享 embedding),save_file 默认会按张量分别写入,导致文件体积翻倍。此时应使用 save_model,它会通过 _shared_pointers 检测共享并只存一份:

python 复制代码
import torch
from safetensors.torch import save_model, load_model

class Net(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.emb = torch.nn.Embedding(10, 4)
        self.head = torch.nn.Linear(4, 10, bias=False)
        self.head.weight = self.emb.weight  # 共享存储

net = Net()
save_model(net, "shared.safetensors")
loaded = load_model(Net(), "shared.safetensors")
print(loaded.head.weight.data_ptr() == loaded.emb.weight.data_ptr())  # True

save_model 的第二个参数是路径,第三个可选 metadata。加载时 load_model 需要传入已构造的模型实例,它按 state_dict 的 key 顺序回填,因此模型结构必须与保存时一致。

文件大小与分片阈值

单个 safetensors 文件没有硬性大小限制,但实践中超过 5GB 会拖慢加载和传输。save_file 不支持自动分片,需要手动按张量切分:

python 复制代码
from safetensors.torch import save_file, load_file
import torch

def save_sharded(state, prefix, max_bytes=5 * 1024**3):
    shard, size, idx = {}, 0, 0
    for k, v in state.items():
        n = v.numel() * v.element_size()
        if size + n > max_bytes and shard:
            save_file(shard, f"{prefix}-{idx:05d}.safetensors")
            shard, size, idx = {}, 0, idx + 1
        shard[k] = v
        size += n
    if shard:
        save_file(shard, f"{prefix}-{idx:05d}.safetensors")

save_sharded({"w": torch.randn(1000, 1000)})

加载端需按分片顺序合并 load_file 返回的字典,或使用 transformersload_sharded_checkpoint

框架版本与格式混用

safetensorstorchnumpytensorflow 分别提供 safetensors.torchsafetensors.numpysafetensors.tensorflow 子模块,跨框架保存后必须用对应子模块加载:

python 复制代码
from safetensors.torch import save_file
from safetensors.numpy import load_file  # 错误:格式不匹配

另外不要用 pickle 打开 safetensors 文件,也不要试图把 safetensors 内容塞进 torch.save 的 pickle 流。safetensors 的设计目标就是零反序列化执行,混用会完全丧失安全性,并可能触发 InvalidHeaderDeserialization 或静默损坏。

7. 适用场景:何时该用 safetensors

模型权重分发与 Hub 上传下载

当你需要把训练好的模型权重分发给别人,或者上传到 Hugging Face Hub 时,safetensors 几乎是默认选择。Hub 上的 model.safetensors 文件可以直接通过 transformersdiffusers 加载,也可以用 huggingface_hub 手动拉取:

python 复制代码
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file

path = hf_hub_download(
    repo_id="bert-base-uncased",
    filename="model.safetensors",
    revision="main",
    cache_dir="./hf_cache",
)
state = load_file(path, device="cpu")
print(list(state.keys())[:3])

hf_hub_downloadrepo_id 是仓库名,filename 是仓库内路径,revision 指定分支或 commit,cache_dir 控制本地缓存目录。返回值是本地文件路径,随后交给 load_file。注意 load_filedevice 参数只接受字符串,比如 "cpu""cuda:0",传 torch.device 对象会报错。

生产环境加载大模型

在生产服务里加载几十 GB 的权重,内存峰值是核心问题。safetensors 支持按张量惰性读取,配合 safe_open 可以只取需要的部分:

python 复制代码
from safetensors import safe_open
import torch

with safe_open("big_model.safetensors", framework="pt", device="cuda:0") as f:
    keys = list(f.keys())
    print(len(keys))
    w = f.get_tensor(keys[0])
    print(w.shape, w.dtype)

safe_open 是上下文管理器,framework 支持 "pt""tf""flax""numpy"device 决定张量落地设备。get_tensor(name) 按需加载,get_slice(name) 还能切分大张量。相比 torch.load 一次性反序列化,这种方式在显存受限时更可控。

与 pickle、np.savez 的对比

pickle/torch.load 会执行任意代码,加载来源不明的 .pt.pkl 文件存在反序列化攻击风险;safetensors 只存储张量数据与 JSON 头部,不执行任何代码,天然规避这类问题,因此需要安全审计的场合应优先选用。np.savez 虽然也不执行代码,但它不支持 GPU 张量、不保留框架元信息,且读取时需先 np.load 再手动转 torch.from_numpy,对大模型权重分发不够方便。safetensors 则原生支持 PyTorch/TF/Flax,零拷贝映射(load_file 底层用 mmap),加载速度通常快于 pickle。唯一需要注意的是它只存张量,优化器状态、超参等非张量内容要另存 JSON。

8. 总结与扩展:结合 huggingface-hub 使用

回顾一下 safetensors 的核心用法:save_file(tensors, filename, metadata=None) 把字典形式的张量写入磁盘,load_file(filename, device="cpu") 读回字典,safe_open(filename, framework="pt", device="cpu") 则支持惰性读取单个键。真正在工程里,权重几乎不会由你手动打包,而是直接从 Hugging Face Hub 拉取------官方发布的模型大多已经是 .safetensors 格式。这时需要 huggingface-hub 来负责下载与缓存。

安装:

bash 复制代码
pip install huggingface-hub safetensors

最常用的入口是 hf_hub_download,它返回本地缓存路径(字符串),配合 load_file 即可完成加载:

python 复制代码
from huggingface_hub import hf_hub_download
from safetensors import safe_open
import torch

path = hf_hub_download(
    repo_id="bert-base-uncased",
    filename="model.safetensors",
    revision="main",
    cache_dir="./hf_cache",
)
print(path)  # 本地缓存文件路径

with safe_open(path, framework="pt", device="cpu") as f:
    print(f.keys())                    # 所有张量名
    emb = f.get_tensor("bert.embeddings.word_embeddings.weight")
    print(emb.shape, emb.dtype)

关键参数:repo_id用户名/仓库名filename 是仓库内相对路径,需与仓库文件列表一致;revision 可传分支、tag 或 commit hash,生产环境建议锁定 commit 以保证可复现;cache_dir 指定缓存目录,默认在 ~/.cache/huggingfacehf_hub_download 有缓存机制,重复调用不会重新下载。safe_open 返回上下文管理器,只有 get_tensor 时才真正读取该张量,适合大模型按需取权重,避免一次性把整个文件载入内存。

如果整个仓库都要下载,用 snapshot_download

python 复制代码
from huggingface_hub import snapshot_download
from safetensors.torch import load_file

local_dir = snapshot_download(
    repo_id="distilbert-base-uncased",
    allow_patterns=["*.safetensors", "config.json"],
    cache_dir="./hf_cache",
)
state = load_file(f"{local_dir}/model.safetensors", device="cpu")
print(len(state), list(state)[:3])

allow_patterns 支持通配符,只拉取需要的文件,能显著减少流量。load_file 会把所有张量一次性读入内存,适合中小模型;超大模型请改用上面的 safe_open 逐张读取。

易错点:一是 filename 写错会抛 EntryNotFoundError,可先用 HfApi().list_repo_files(repo_id) 确认;二是私有仓库需先 huggingface-cli login 或设置 HF_TOKEN 环境变量;三是 safe_openframework 参数必须与后续 get_tensor 期望的框架一致,取 "pt" 得到 torch.Tensor,取 "tf" 得到 TensorFlow 张量。

进一步学习方向:transformersfrom_pretrained 已内置 safetensors 支持,会自动优先加载 .safetensors 而非 .bin,并可配合 acceleratedevice_map 分片加载、配合 peft 加载 LoRA 权重。理解了本节的下载与惰性读取,再去看这些高层 API 的源码会顺畅很多。

9. 与 Web 服务集成:用 FastAPI 提供 safetensors 权重下载接口

模型训练完成后,最常见的需求之一是把权重通过 HTTP 暴露给下游服务。safetensorssafe_open 支持惰性读取单个张量,配合 FastAPI 的 StreamingResponse 可以做到「按需取张量、不整包加载进内存」。

python 复制代码
# pip install fastapi uvicorn safetensors torch
from fastapi import FastAPI, HTTPException
from fastapi.responses import Response
from safetensors import safe_open
from safetensors.torch import save_file
import torch, io

app = FastAPI()
CKPT = "model.safetensors"

# 一次性写入示例权重
save_file({"layer0.weight": torch.randn(4, 4),
           "layer1.weight": torch.randn(4, 4)}, CKPT)

@app.get("/keys")
def list_keys():
    with safe_open(CKPT, framework="pt") as f:
        return {"keys": list(f.keys())}

@app.get("/tensor/{name}")
def get_tensor(name: str):
    with safe_open(CKPT, framework="pt") as f:
        if name not in f.keys():
            raise HTTPException(404, f"tensor {name} not found")
        t = f.get_tensor(name)          # 只反序列化这一个张量
    buf = io.BytesIO()
    torch.save(t, buf)                  # 或改用 numpy.tobytes 自定义协议
    return Response(buf.getvalue(), media_type="application/octet-stream")

启动 uvicorn main:app 后,GET /tensor/layer0.weight 只会把目标张量读进内存。这种模式在权重体积远大于单卡显存时尤其有用,避免了「先整包 load 再切片」的浪费。

为什么不用 pickle 直接返回

如果接口返回的是 Python 对象,很多人会顺手用 pickle.dumps。但 pickle 反序列化会执行任意代码,一旦服务端被投毒,攻击面极大。safetensors 的格式头部只包含 dtype、shape、data_offsets 等纯数据字段,天然免疫这类问题,这也是它适合做跨进程、跨网络传输格式的原因。

10. 与配置文件、校验库的组合用法

实际项目中,权重文件常伴随一份描述「有哪些张量、期望什么 shape」的元数据。可以用 tomlkit 写配置、jsonschema-specifications 做校验,再用 safetensors 校验实际权重是否匹配。

python 复制代码
# pip install tomlkit jsonschema-specifications jsonschema safetensors torch
import tomlkit, torch
from jsonschema import validate
from safetensors.torch import save_file
from safetensors import safe_open

cfg = tomlkit.parse("""
[model]
name = "demo"
[model.tensors.layer0]
shape = [4, 4]
dtype = "float32"
""")

schema = {
    "type": "object",
    "properties": {"model": {"type": "object"}},
    "required": ["model"],
}
validate(instance=cfg.unwrap(), schema=schema)

save_file({"layer0": torch.randn(4, 4)}, "demo.safetensors")

with safe_open("demo.safetensors", framework="pt") as f:
    for key in f.keys():
        t = f.get_tensor(key)
        spec = cfg["model"]["tensors"][key]
        assert list(t.shape) == spec["shape"], f"shape mismatch on {key}"
        assert str(t.dtype).replace("torch.", "") == spec["dtype"]
print("权重与 TOML 配置一致")

tomlkit 的好处是保留注释与格式,回写配置时不会把人工维护的文档结构冲掉;jsonschema 负责结构合法性,safetensors 负责二进制层面的事实一致性,三者职责不重叠。

与同类库的语法对比

操作 safetensors torch.save/load numpy.savez
保存 save_file(dict, path) torch.save(obj, path) np.savez(path, **arrs)
读取全部 load_file(path) torch.load(path) np.load(path)
惰性读单个 safe_open(...).get_tensor(k) 不支持(整包) 不支持(整包)
安全反序列化 否(pickle)
跨框架 pt/tf/flax/np 仅 torch 仅 numpy

可以看到,safetensors 在「惰性读取」和「跨框架」两点上是明显优势,代价是不支持保存任意 Python 对象------它只存张量,这恰恰是它安全性的来源。

常见报错与排查

  • SafetensorError: Error while deserializing header:文件被截断或非 safetensors 格式,检查下载是否完整、是否误把 .bin 改名。
  • KeyErrorget_tensor 的 key 不存在,先用 f.keys() 打印确认,注意前缀如 model. 是否一致。
  • RuntimeError: expected scalar type ...:dtype 不匹配,用 t.to(dtype) 显式转换,或保存时统一为 float32/float16。
  • framework 参数与张量类型不符:safe_openframework="pt" 返回 torch.Tensor,"np" 返回 numpy 数组,二者不能混用。

与 frozenlist、websockets 的协作场景

在流式推理服务里,可以用 frozenlist 保存一个不可变的权重 key 列表,防止运行期被误改;用 websockets 把张量分块推送给客户端,而不是一次性 HTTP 响应。frozenlist__setitem__ 会直接抛 RuntimeError,适合做「配置冻结」,而 websocketssend 接受 bytes,正好与 safetensors 的字节流衔接。

python 复制代码
# pip install frozenlist websockets safetensors torch
from frozenlist import FrozenList
from safetensors import safe_open
import asyncio, websockets

with safe_open("demo.safetensors", framework="pt") as f:
    keys = FrozenList(f.keys())
keys.freeze()                       # 之后 keys[0] = "x" 会抛 RuntimeError

async def handler(ws):
    for k in keys:
        with safe_open("demo.safetensors", framework="pt") as f:
            t = f.get_tensor(k)
        await ws.send(f"{k}:{t.shape}".encode())

async def main():
    async with websockets.serve(handler, "localhost", 8765):
        await asyncio.Future()

# asyncio.run(main())

这样权重服务既保持了只读语义(frozenlist),又通过异步通道高效分发(websockets),而 safetensors 始终只暴露张量数据本身,不引入 pickle 风险。

相关推荐
千里码aicood1 小时前
基于Python语言的测量程序设计
开发语言·python
2601_962300811 小时前
用PyTorch Lightning简化空间分析:Python和机器学习的力量
python·深度学习·机器学习·空间分析·pytorchlightning
@小匠1 小时前
1Password-开源计划申请免费使用资格
开源
jayson.h2 小时前
python-可视化提取表格-通用
开发语言·python
悟天特斯2 小时前
AI驱动的楼宇节能:从粗放管控到精准降碳的实践路径
开发语言·人工智能·python·物联网
名字还没想好☜2 小时前
Python 用 tempfile 安全创建临时文件:NamedTemporaryFile、TemporaryDirectory 与别自己拼 /tmp 的坑
开发语言·后端·python·安全·编程语言
Behaviour3 小时前
DeepSeek V4.1 Flash 发布开源,Harness v0.1.5同日适配
人工智能·语言模型·开源·aigc·ai编程
北城bot3 小时前
把 Python 脚本打包成 exe
python
狗凯之家源码网3 小时前
全开源 H5 棋牌对战系统修复优化与二次开发实测
开源·php·棋牌对战