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

如何高效存储BERT编码器生成的Embedding并实现快速随机访问?

针对Polars Parquet的优化方案
  • 调整Parquet行组大小:默认Parquet的行组设置不适合随机访问,建议把行组大小设为10万-100万条(对应单组数据量约300MB-3GB,适配磁盘IO特性)。用Polars写入时显式指定参数:
    pl.write_parquet("embeddings.parquet", row_group_size=100_000)
    
    这样按ID过滤时,Polars能快速定位到目标行所在的行组,无需扫描全量数据。
  • 显式存储ID列并优化索引:不要依赖默认行号,单独添加id列(从0到1.59999999亿),写入前按id排序Parquet文件,或者按id范围分区(比如每100万条一个分区)。查询时用df.filter(pl.col("id").is_in(target_ids)),Polars会直接定位到对应分区/行组,大幅提速。
  • 开启Predicate Pushdown:用scan_parquet时确保predicate_pushdown=True(默认已开启),让过滤逻辑下推到Parquet读取层,只加载需要的行组数据,配合collect(streaming=True)处理大批次查询,避免内存过载。
其他高效存储+快速随机访问方案

HDF5格式

专门为大规模数值数据设计,支持高效随机访问与无损压缩,存储效率和访问速度都优于Parquet的随机场景:

import h5py
import numpy as np

# 写入(支持gzip/lzf无损压缩)
with h5py.File("embeddings.h5", "w") as f:
    f.create_dataset(
        "embeddings",
        data=embeddings_array,
        dtype="float32",
        compression="gzip",  # 压缩等级可选1-9,推荐4-6平衡速度与压缩率
        compression_opts=4
    )

# 批量读取指定ID
with h5py.File("embeddings.h5", "r") as f:
    batch_embeddings = f["embeddings"][target_ids, :]

HDF5会记录数据块的位置索引,读取时直接定位目标数据,无需扫描全量文件,压缩后空间占用比原.pt减少30%-50%。

LMDB键值存储

适合以ID为键的随机访问场景,读写性能优异,支持批量操作:

import lmdb
import numpy as np

# 初始化环境(预分配足够空间,比如500GB)
env = lmdb.open("embeddings_lmdb", map_size=500 * 1024 ** 3)

# 批量写入(每1万条提交一次提升效率)
batch_size = 10_000
with env.begin(write=True) as txn:
    for idx in range(0, len(embeddings_array), batch_size):
        batch_ids = range(idx, min(idx+batch_size, len(embeddings_array)))
        for id in batch_ids:
            txn.put(str(id).encode(), embeddings_array[id].tobytes())

# 批量读取指定ID
with env.begin() as txn:
    batch_embeddings = np.array([
        np.frombuffer(txn.get(str(id).encode()), dtype=np.float32)
        for id in target_ids
    ])

LMDB的随机读取延迟极低,适合高频小批量查询场景,若配合msgpack序列化压缩,还能进一步节省空间。

TileDB数组存储

支持稠密/稀疏数组,原生支持快速随机访问与分布式存储,特性类似HDF5但更适配云环境:

import tiledb
import numpy as np

# 创建数组 schema
dom = tiledb.Domain(
    tiledb.Dim(name="id", domain=(0, 159999999), dtype=np.int64),
    tiledb.Dim(name="dim", domain=(0, 767), dtype=np.int64)
)
schema = tiledb.ArraySchema(
    domain=dom,
    sparse=False,
    attrs=[tiledb.Attr(name="emb", dtype=np.float32, compression="gzip")]
)
tiledb.Array.create("embeddings.tdb", schema)

# 写入全量数据
with tiledb.open("embeddings.tdb", "w") as arr:
    arr[:] = embeddings_array

# 批量读取指定ID
with tiledb.open("embeddings.tdb", "r") as arr:
    batch_embeddings = arr[target_ids, :]

TileDB支持多进程/多线程并发读写,压缩效率与HDF5相当,适合需要分布式访问的场景。

空间节省补充建议
  • 坚持用无损压缩:所有推荐的格式都支持gzip/snappy等无损压缩,不会损失float32精度,能有效降低磁盘占用(通常压缩比在1.5-2倍左右)。
  • 避免拆分embedding:不要把768维embedding拆成多列存储,保持二维数组形式,能大幅提升压缩效率与访问速度。

内容的提问来源于stack exchange,提问作者Noam

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 14:51:01