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_file、load_file 和 safe_open,以及两个异常类 SafetensorError 和 safe_open 返回的 safe_open 上下文对象。常用的导入写法是:
python
from safetensors import safe_open, SafetensorError
from safetensors.numpy import save_file, load_file
注意这里有一个容易踩的坑:save_file 和 load_file 并不在 safetensors 顶层命名空间里,而是按后端拆分。safetensors.numpy 处理 np.ndarray,safetensors.torch 处理 torch.Tensor,safetensors.flax、safetensors.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_file 的 metadata 必须是 dict[str, str],值只能是字符串,传整数或列表会抛 SafetensorError;load_file 的 device 在 numpy 后端下只能是 "cpu",在 torch 后端可传 "cuda:0" 等;safe_open 的 framework 取值包括 "pt"、"tf"、"flax"、"numpy"、"paddle",它不会一次性把全部权重读进内存,适合加载大模型。若文件损坏或格式不符,safe_open 与 load_file 都会抛出 SafetensorError,捕获它即可做降级处理。
2. 核心对象与语法:save_file、load_file 与 safe_open
safetensors 的核心 API 非常精简,围绕三个对象展开:save_file、load_file 和 safe_open。先安装并导入:
bash
pip install safetensors
python
from safetensors import safe_open
from safetensors.torch import save_file, load_file
注意 save_file 与 load_file 位于框架子模块中(如 safetensors.torch、safetensors.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 块后文件句柄自动关闭。易错点:framework 与 device 类型不匹配会抛错;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 的值必须是字符串,传数字或列表会抛 TypeError。safe_open 支持惰性读取,只有调用 f.get_tensor(key) 时才真正加载对应张量。
save_model 与 load_model
save_model 直接保存 torch.nn.Module,load_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_open 的 get_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 实例,第二个是保存路径,第三个可选 metadata 是 dict[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 返回的字典,或使用 transformers 的 load_sharded_checkpoint。
框架版本与格式混用
safetensors 对 torch、numpy、tensorflow 分别提供 safetensors.torch、safetensors.numpy、safetensors.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 文件可以直接通过 transformers、diffusers 加载,也可以用 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_download 的 repo_id 是仓库名,filename 是仓库内路径,revision 指定分支或 commit,cache_dir 控制本地缓存目录。返回值是本地文件路径,随后交给 load_file。注意 load_file 的 device 参数只接受字符串,比如 "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/huggingface。hf_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_open 的 framework 参数必须与后续 get_tensor 期望的框架一致,取 "pt" 得到 torch.Tensor,取 "tf" 得到 TensorFlow 张量。
进一步学习方向:transformers 的 from_pretrained 已内置 safetensors 支持,会自动优先加载 .safetensors 而非 .bin,并可配合 accelerate 做 device_map 分片加载、配合 peft 加载 LoRA 权重。理解了本节的下载与惰性读取,再去看这些高层 API 的源码会顺畅很多。
9. 与 Web 服务集成:用 FastAPI 提供 safetensors 权重下载接口
模型训练完成后,最常见的需求之一是把权重通过 HTTP 暴露给下游服务。safetensors 的 safe_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改名。KeyError:get_tensor的 key 不存在,先用f.keys()打印确认,注意前缀如model.是否一致。RuntimeError: expected scalar type ...:dtype 不匹配,用t.to(dtype)显式转换,或保存时统一为 float32/float16。framework参数与张量类型不符:safe_open的framework="pt"返回 torch.Tensor,"np"返回 numpy 数组,二者不能混用。
与 frozenlist、websockets 的协作场景
在流式推理服务里,可以用 frozenlist 保存一个不可变的权重 key 列表,防止运行期被误改;用 websockets 把张量分块推送给客户端,而不是一次性 HTTP 响应。frozenlist 的 __setitem__ 会直接抛 RuntimeError,适合做「配置冻结」,而 websockets 的 send 接受 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 风险。