flowchart TB
subgraph "对象存储在 AI 中的角色"
A[原始数据桶<br/>raw-data/]
B[处理数据桶<br/>processed/]
C[模型权重桶<br/>models/]
D[Checkpoint 桶<br/>checkpoints/]
E[特征存储桶<br/>features/]
end
F[数据采集] --> A
A --> G[处理管线]
G --> B
B --> H[训练集群]
C --> I[推理服务]
D --> J[容灾恢复]
第15章 存储架构
第15章 存储架构
存储是 AI 基础设施的地基。模型越大、数据越多,存储架构的选择就越深刻地影响你的训练效率和推理延迟。
一个典型的 AI 工作负载涉及从 GB 到 PB 跨度的数据:从几个 GB 的模型权重到数百 TB 的训练数据,从毫秒级的向量检索到分钟级的 checkpoint 写入。没有一种存储方案能同时满足所有需求。本章讨论如何为不同场景选择和组合存储方案。
15.1 对象存储在 AI 中的应用
为什么对象存储是默认选择
对象存储(Object Storage)——S3、GCS、Azure Blob、MinIO——已经成为 AI 数据的默认存储层。原因很简单:
- 无限容量:不需要预先分配空间
- 低成本:S3 Standard 约 $0.023/GB/月
- 高持久性:11 个 9 的数据持久性
- HTTP API:任何地方都能访问
但对象存储有一个根本性的限制:它不是文件系统。没有目录层级结构、没有原子重命名、不支持文件追加、一致性模型是”写后读”(read-after-write)。
AI 场景下的典型用法
S3 性能优化实践
import boto3
from botocore.config import Config
from concurrent.futures import ThreadPoolExecutor
import multiprocessing as mp
class S3DataManager:
"""优化的 S3 数据管理器"""
def __init__(self, bucket: str):
self.bucket = bucket
# 优化的 S3 客户端
self.s3 = boto3.client(
's3',
config=Config(
max_pool_connections=50, # 连接池大小
tcp_keepalive=True,
request_timeout=60,
retries={'max_attempts': 5},
)
)
def upload_large_file(self, local_path: str, s3_key: str):
"""大文件分片上传(自动多线程)"""
from boto3.s3.transfer import TransferConfig
transfer_config = TransferConfig(
multipart_threshold=100 * 1024 * 1024, # 100MB 触发分片
multipart_chunksize=100 * 1024 * 1024, # 每片 100MB
max_concurrency=10, # 并发上传
use_threads=True,
)
self.s3.upload_file(
local_path, self.bucket, s3_key,
Config=transfer_config
)
def parallel_download_shards(self, s3_prefix: str,
local_dir: str, num_workers: int = None):
"""并行下载多个 shard 文件"""
if num_workers is None:
num_workers = mp.cpu_count()
# 列出所有文件
files = self._list_objects(s3_prefix)
def download_one(s3_key):
local_path = f"{local_dir}/{s3_key.split('/')[-1]}"
self.s3.download_file(self.bucket, s3_key, local_path)
return local_path
with ThreadPoolExecutor(max_workers=num_workers) as pool:
paths = list(pool.map(download_one, files))
return paths
def _list_objects(self, prefix: str) -> list[str]:
"""递归列出所有对象(自动分页)"""
keys = []
paginator = self.s3.get_paginator('list_objects_v2')
for page in paginator.paginate(Bucket=self.bucket, Prefix=prefix):
for obj in page.get('Contents', []):
keys.append(obj['Key'])
return keysS3 Express One Zone:低延迟场景
2024 年 AWS 推出的 S3 Express One Zone 将对象存储的访问延迟从毫秒级降到了亚毫秒级(p99 < 10ms),这对 GPU 直读训练数据的场景意义重大:
# 配置 S3 Express One Zone
s3_express = boto3.client(
's3',
config=Config(
# S3 Express 使用不同的 endpoint
s3={'bucket_name_style': 'directory'},
)
)
# 创建 Express bucket
s3_express.create_bucket(
Bucket='ai-training-hot-data--use1-az1--x-s3',
CreateBucketConfiguration={
'Bucket': {
'Type': 'Directory',
'DataRedundancy': 'SingleAvailabilityZone',
},
'Location': {
'Type': 'AvailabilityZone',
'Name': 'use1-az1',
}
}
)性能对比:标准 S3 的 GET 延迟约 30-100ms(p99),S3 Express One Zone 约 5-10ms(p99)。对于训练数据加载,这意味着你可以在不复制到本地的情况下直接从 S3 读取 shard。
成本提醒:S3 Express One Zone 的存储价格是 Standard 的 2 倍左右(~$0.05/GB/月),但 request 价格便宜 50%。适合频繁访问的热数据,不适合冷存储。
15.2 分布式文件系统
为什么 GPU 集群需要分布式文件系统
大规模分布式训练有一个硬性需求:所有 worker 节点需要同时访问同一份数据。对象存储可以做到,但延迟和吞吐量无法满足密集读取的场景。
这就是分布式文件系统的领域——Lustre、GPFS(Spectrum Scale)和新兴的 JuiceFS。
主流方案对比
| 特性 | Lustre | GPFS | JuiceFS |
|---|---|---|---|
| 架构 | 内核态,独立存储 | 内核态,集成存储 | 用户态,对象存储后端 |
| 吞吐 | 极高(TB/s 级) | 极高 | 中高(受限于后端) |
| ** POSIX 兼容** | 完全 | 完全 | 完全 |
| 部署复杂度 | 高(需要专用 MDS/OSS) | 高 | 低 |
| 云原生 | AWS FSx for Lustre | AWS FSx for OpenZFS | 原生多云 |
| 成本 | 高 | 高 | 低(利用现有对象存储) |
| 最佳场景 | HPC、超大规模训练 | 企业 HPC | 中小团队、云原生 |
Lustre:HPC 领域的王者
Lustre 是超算中心和大模型训练集群的首选。它的架构设计为高吞吐量优化:
flowchart TB
subgraph "Lustre 架构"
Client[Client Nodes<br/>GPU 训练节点]
MDS[Metadata Server<br/>MDS + MDT]
OSS1[Object Storage Target 1<br/>OST]
OSS2[Object Storage Target 2<br/>OST]
OSS3[Object Storage Target 3<br/>OST]
Client -->|metadata 操作| MDS
Client -->|数据读写| OSS1
Client -->|数据读写| OSS2
Client -->|数据读写| OSS3
end
在 AWS 上使用 FSx for Lustre:
import boto3
fsx = boto3.client('fsx')
# 创建 Lustre 文件系统
response = fsx.create_file_system(
FileSystemType='LUSTRE',
StorageCapacity=4800, # GB,最小 1200 GB
StorageType='SSD',
LustreConfiguration={
'DeploymentType': 'PERSISTENT_2', # 持久型
'PerUnitStorageThroughput': 200, # MB/s/TiB
'MetadataConfiguration': {
'MetadataCapacity': 2400, # iops:capacity = 3:1
}
},
SubnetIds=['subnet-xxxxx'],
SecurityGroupIds=['sg-xxxxx'],
)
# 挂载到 GPU 节点
# mount -t lustre -o noatime,flock fs-xxxxx.fsx.us-east-1.amazonaws.com@tcp:/fsx /mnt/fsx关键调优参数:
# /etc/modprobe.d/lustre.conf
# 增大 OSC 缓存
options lnet lnd_timeout=120
options llite llite_max_cached_mb=16384 # 16GB client cache
# 挂载选项优化
mount -t lustre -o noatime,nodiratime,flock,user_xattr \
fs-xxxxx.fsx.us-east-1.amazonaws.com@tcp:/fsx /mnt/fsx
# stripe 优化(大文件场景)
lfs setstripe /mnt/fsx/training_data -c 32 -s 256M
# -c: stripe count,跨多少个 OST
# -s: stripe size,每个条带的大小JuiceFS:云原性的新选择
JuiceFS 采用创新的架构:元数据引擎(Redis/TiKV/MySQL)+ 对象存储后端(S3/MinIO)。这意味着你可以用已有的对象存储获得文件系统的体验。
# juicefs-ce.yaml — JuiceFS Community Edition 配置
apiVersion: v1
kind: PersistentVolume
metadata:
name: juicefs-pv
spec:
capacity:
storage: 10Ti
accessModes:
- ReadWriteMany # 关键:多节点同时读写
persistentVolumeReclaimPolicy: Retain
csi:
driver: csi.juicefs.com
volumeHandle: ai-training-data
volumeAttributes:
juicefs/metadata-url: "redis://redis:6379/0"
juicefs/bucket: "https://s3.us-east-1.amazonaws.com/ai-data"
juicefs/access-key: "${AWS_ACCESS_KEY_ID}"
juicefs/secret-key: "${AWS_SECRET_ACCESS_KEY}"
---
apiVersion: v1
kind: PersistentVolumeClaim
metadata:
name: juicefs-pvc
spec:
accessModes:
- ReadWriteMany
resources:
requests:
storage: 10Ti
volumeName: juicefs-pv# 快速格式化和挂载
juicefs format \
--storage s3 \
--bucket https://s3.us-east-1.amazonaws.com/ai-data \
redis://redis:6379/0 \
myjfs
juicefs mount redis://redis:6379/0 /mnt/jfs \
--buffer-size 1024 \
--prefetch 1 \
--max-uploads 20
# 性能基准测试
juicefs bench /mnt/jfsJuiceFS 的杀手锏:和 S3 不同,JuiceFS 提供完整的 POSIX 语义。你的训练代码不需要改动——像使用本地文件系统一样使用它。同时数据实际存储在 S3 上,成本和持久性与 S3 一致。
15.3 向量数据库与 Embedding 检索
向量检索:RAG 和语义搜索的基础
大模型时代,向量数据库从一个小众工具变成了基础设施的核心组件。无论是 RAG(Retrieval-Augmented Generation)、语义搜索还是推荐系统,都需要在海量 embedding 中快速找到最相似的结果。
向量数据库选型
flowchart TD
A[需求分析] --> B{数据规模?}
B -->|< 100万| C{需要过滤?}
B -->|100万-1亿| D[Pinecone / Weaviate]
B -->|> 1亿| E[Milvus / Qdrant Cluster]
C -->|否| F[FAISS / ScaNN]
C -->|是| G[Chroma / Qdrant]
F --> H{延迟要求?}
H -->|< 10ms| I[FAISS GPU]
H -->|> 10ms| J[FAISS CPU]
Milvus:大规模向量检索
Milvus 是目前最流行的开源向量数据库,支持十亿级向量检索:
from pymilvus import MilvusClient, DataType
import numpy as np
# 连接 Milvus
client = MilvusClient(uri="http://localhost:19530")
# 创建 Collection
collection_name = "documents"
schema = client.create_schema(auto_id=True)
schema.add_field("id", DataType.INT64, is_primary=True)
schema.add_field("embedding", DataType.FLOAT_VECTOR, dim=1024)
schema.add_field("text", DataType.VARCHAR, max_length=8192)
schema.add_field("source", DataType.VARCHAR, max_length=256)
schema.add_field("timestamp", DataType.INT64)
# 索引参数(HNSW 是最常用的近似最近邻索引)
index_params = client.prepare_index_params()
index_params.add_index(
field_name="embedding",
index_type="HNSW",
metric_type="COSINE",
params={
"M": 16, # 每个节点的最大邻居数
"efConstruction": 256, # 构建时候选集大小
}
)
client.create_collection(
collection_name=collection_name,
schema=schema,
index_params=index_params,
)
# 插入数据
def insert_documents(texts: list[str], embeddings: np.ndarray, source: str):
data = [
{
"embedding": emb.tolist(),
"text": text,
"source": source,
"timestamp": int(time.time()),
}
for text, emb in zip(texts, embeddings)
]
client.insert(collection_name, data)
# 检索
def search_similar(
query_embedding: np.ndarray,
top_k: int = 10,
filter_expr: str = "",
):
results = client.search(
collection_name=collection_name,
data=[query_embedding.tolist()],
limit=top_k,
filter=filter_expr, # e.g., 'source == "arxiv"'
output_fields=["text", "source", "timestamp"],
search_params={
"params": {"ef": 128} # 搜索时候选集大小,越大越精确
}
)
return results[0]索引类型选择
| 索引类型 | 构建速度 | 查询速度 | 召回率 | 内存占用 | 最佳场景 |
|---|---|---|---|---|---|
| FLAT | 极快 | 慢 | 100% | 低 | 小数据集(< 100K) |
| IVF_FLAT | 中 | 快 | 高 | 中 | 中等数据集 |
| IVF_PQ | 中 | 极快 | 中 | 低 | 内存受限 |
| HNSW | 慢 | 极快 | 高 | 高 | 通用首选 |
| DISKANN | 慢 | 中 | 高 | 极低 | 超大规模(内存不够) |
实践经验:对于大多数应用(< 1 亿向量),HNSW 是最佳选择。它的查询延迟在 1-5ms(p99),召回率可达 98% 以上。只有当内存装不下时才考虑 DiskANN。
Embedding 管线
一个完整的 embedding 检索系统不只是向量数据库,还包括 embedding 生成管线:
from sentence_transformers import SentenceTransformer
import torch
class EmbeddingPipeline:
"""端到端 Embedding 管线"""
def __init__(self, model_name: str = "BAAI/bge-large-en-v1.5"):
self.device = "cuda" if torch.cuda.is_available() else "cpu"
self.model = SentenceTransformer(
model_name, device=self.device
)
self.batch_size = 64 if self.device == "cuda" else 16
def encode(self, texts: list[str],
normalize: bool = True) -> np.ndarray:
"""批量编码文本为 embedding"""
embeddings = self.model.encode(
texts,
batch_size=self.batch_size,
normalize_embeddings=normalize,
show_progress_bar=True,
convert_to_numpy=True,
)
return embeddings
def encode_streaming(self, text_stream, batch_size: int = 64):
"""流式编码,适用于大数据集"""
batch = []
for text in text_stream:
batch.append(text)
if len(batch) >= batch_size:
yield self.encode(batch)
batch = []
if batch:
yield self.encode(batch)
# 异步批量编码写入
import asyncio
async def build_index_async(
pipeline: EmbeddingPipeline,
documents: list[dict],
milvus_client: MilvusClient,
):
"""异步构建向量索引"""
sem = asyncio.Semaphore(4) # 并发限制
async def insert_batch(batch):
async with sem:
texts = [d['text'] for d in batch]
embeddings = pipeline.encode(texts)
insert_documents(texts, embeddings, batch[0]['source'])
# 分批处理
batch_size = 256
tasks = [
insert_batch(documents[i:i+batch_size])
for i in range(0, len(documents), batch_size)
]
await asyncio.gather(*tasks)15.4 Feature Store 与特征服务
Feature Store:连接离线和在线的桥梁
Feature Store 解决了一个核心矛盾:训练时用的是批量特征,推理时需要实时特征。如果两边的特征计算逻辑不一致,就会出现”训练-推理偏差”(Training-Serving Skew)。
flowchart LR
subgraph "离线(训练)"
A[原始数据] --> B[批处理特征计算]
B --> C[离线特征存储]
end
subgraph "在线(推理)"
D[实时事件] --> E[流式特征计算]
E --> F[在线特征存储]
end
G[Feature Store] --> C
G --> F
H[模型训练] --> C
I[在线推理] --> F
style G fill:#faa
Feast:开源 Feature Store
Feast(Feature Store)是目前最流行的开源方案,它的核心理念是”一次定义,离线在线通用”:
# feature_repo/example.py — Feast 特征定义
from feast import Entity, FeatureView, Field, FileSource, RequestSource
from feast.types import Float32, Int64, String
from datetime import timedelta
# 定义数据源(离线)
batch_source = FileSource(
name="user_features_batch",
path="s3://feature-store/user_features.parquet",
timestamp_field="event_ts",
)
# 定义实时请求源
request_source = RequestSource(
name="user_request",
schema=[
Field(name="user_id", dtype=String),
]
)
# 实体定义
user_entity = Entity(
name="user_id",
join_keys=["user_id"],
)
# 特征视图
user_features = FeatureView(
name="user_features",
entities=[user_entity],
ttl=timedelta(days=30),
schema=[
Field(name="total_purchases", dtype=Int64),
Field(name="avg_order_value", dtype=Float32),
Field(name="days_since_last_visit", dtype=Int64),
Field(name="engagement_score", dtype=Float32),
],
online=True, # 启用在线服务
source=batch_source,
tags={"team": "growth"},
)# 训练时:从离线存储获取历史特征
from feast import FeatureStore
store = FeatureStore(repo_path="feature_repo")
# 获取训练数据的历史特征
training_data = store.get_historical_features(
entity_df=training_entities, # 包含 user_id 和 event_ts
features=[
"user_features:total_purchases",
"user_features:avg_order_value",
"user_features:engagement_score",
],
).to_df()
# 推理时:从在线存储获取实时特征
features = store.get_online_features(
features=[
"user_features:total_purchases",
"user_features:avg_order_value",
"user_features:engagement_score",
],
entity_rows=[{"user_id": "user_12345"}],
).to_dict()
print(features)
# {'user_id': ['user_12345'],
# 'total_purchases': [42],
# 'avg_order_value': [89.5],
# 'engagement_score': [0.73]}特征计算管线
# Airflow DAG: 定时计算特征并写入 Feature Store
from airflow import DAG
from airflow.operators.python import PythonOperator
from datetime import datetime, timedelta
default_args = {
'owner': 'data-platform',
'retries': 2,
'retry_delay': timedelta(minutes=5),
}
dag = DAG(
'feature_pipeline_daily',
default_args=default_args,
schedule_interval='0 2 * * *', # 每天凌晨 2 点
start_date=datetime(2026, 1, 1),
catchup=False,
)
def compute_user_features(**context):
"""计算用户聚合特征"""
from pyspark.sql import SparkSession
from pyspark.sql.functions import (
col, count, avg, max as spark_max,
datediff, current_date
)
spark = SparkSession.builder.getOrCreate()
# 读取原始事件
events = spark.read.parquet(
"s3://events/dt=" + context['ds'] # Airform 日期
)
# 计算特征
features = events.groupBy("user_id").agg(
count("order_id").alias("total_purchases"),
avg("order_amount").alias("avg_order_value"),
spark_max("event_ts").alias("last_visit_ts"),
).withColumn(
"days_since_last_visit",
datediff(current_date(), col("last_visit_ts"))
)
# 写入离线存储
features.write.mode("overwrite").parquet(
f"s3://feature-store/user_features/dt={context['ds']}"
)
# 物化到在线存储(Redis/DynamoDB)
push_to_online_store(features)
compute_task = PythonOperator(
task_id='compute_user_features',
python_callable=compute_user_features,
provide_context=True,
dag=dag,
)Training-Serving Skew 检测:定期对比离线特征和在线特征的统计分布。如果偏差超过阈值(如 PSI > 0.1),说明计算逻辑不一致或数据延迟。
15.5 多级缓存与数据预热策略
存储层次结构
大规模 AI 系统的存储层次遵循一个简单的原则:热数据在最近的存储层,冷数据在便宜的存储层。
flowchart TB
subgraph "存储层次(从热到冷)"
L1[L1: GPU HBM<br/>80-192 GB<br/>~30 TB/s 带宽]
L2[L2: CPU 内存<br/>256-2048 GB<br/>~200 GB/s 带宽]
L3[L3: NVMe 本地盘<br/>3.6-15 TB<br/>~7 GB/s 带宽]
L4[L4: 分布式缓存<br/>集群级<br/>~1-3 GB/s 带宽]
L5[L5: 对象存储<br/>无限<br/>~0.1-1 GB/s 带宽]
end
L1 --> L2 --> L3 --> L4 --> L5
style L1 fill:#f66
style L5 fill:#6f6
每一层的带宽差大约 10-30 倍,延迟差也类似。关键目标是让数据尽量待在离 GPU 近的层。
GPU 显存管理
import torch
class GPUMemoryManager:
"""GPU 显存管理策略"""
def __init__(self, total_memory_gb: float = 80.0):
self.total = total_memory_gb * 1024**3 # bytes
# 预留 10% 给系统/框架开销
self.usable = self.total * 0.9
def compute_optimal_batch_size(
self,
model_memory_gb: float,
sample_memory_mb: float,
gradient_factor: float = 2.0, # Adam 需要 2x 模型大小
) -> int:
"""
计算最优 batch size(充分利用显存)
"""
model_bytes = model_memory_gb * 1024**3
sample_bytes = sample_memory_mb * 1024**2
# 可用于 batch 的显存
batch_budget = (
self.usable
- model_bytes * gradient_factor # 模型 + 梯度 + 优化器状态
)
batch_size = int(batch_budget // sample_bytes)
# 对齐到 8 的倍数(Tensor Core 优化)
batch_size = (batch_size // 8) * 8
return max(batch_size, 8)
# 使用
mgr = GPUMemoryManager(total_memory_gb=80)
batch_size = mgr.compute_optimal_batch_size(
model_memory_gb=14, # 7B 模型 fp16
sample_memory_mb=200, # 每个样本 ~200MB
)
print(f"Optimal batch size: {batch_size}")数据预热策略
在训练开始前,提前将数据从慢速存储加载到快速存储:
import asyncio
import aiofiles
from pathlib import Path
class DataWarmer:
"""数据预热器:将即将使用的数据预加载到 NVMe/内存"""
def __init__(self, cache_dir: str = "/mnt/nvme/cache",
max_cache_size_gb: float = 1000):
self.cache_dir = Path(cache_dir)
self.max_cache_bytes = int(max_cache_size_gb * 1024**3)
async def warm_for_epoch(
self,
shard_paths: list[str],
look_ahead: int = 3,
):
"""
预热接下来 look_ahead 个 shard。
与训练循环并行执行。
"""
for i, shard_path in enumerate(shard_paths[:look_ahead]):
local_path = self.cache_dir / Path(shard_path).name
if not local_path.exists():
await self._download(shard_path, local_path)
async def _download(self, remote: str, local: Path):
"""异步下载文件"""
# 实际实现可以用 boto3 async 或 s5cmd
local.parent.mkdir(parents=True, exist_ok=True)
# 示例:使用 s5cmd 并行下载
proc = await asyncio.create_subprocess_exec(
"s5cmd", "cp", remote, str(local),
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
await proc.communicate()
def evict_old(self, current_index: int, shard_paths: list[str],
keep_behind: int = 2):
"""清理已经使用过的 shard"""
for i in range(max(0, current_index - keep_behind)):
local_path = self.cache_dir / Path(shard_paths[i]).name
if local_path.exists():
local_path.unlink()
# 检查缓存总大小
self._enforce_size_limit()
def _enforce_size_limit(self):
"""LRU 清理"""
files = sorted(
self.cache_dir.rglob("*"),
key=lambda f: f.stat().st_mtime
)
total = sum(f.stat().st_size for f in files if f.is_file())
while total > self.max_cache_bytes and files:
f = files.pop(0)
size = f.stat().st_size
f.unlink()
total -= size
# 在训练脚本中集成
async def train_with_warming(dataloader, shard_paths):
warmer = DataWarmer(cache_dir="/mnt/nvme/cache")
# 开始预热前 3 个 shard
warming_task = asyncio.create_task(
warmer.warm_for_epoch(shard_paths, look_ahead=3)
)
for batch_idx, batch in enumerate(dataloader):
# 训练逻辑
train_step(batch)
# 每处理 N 个 batch,检查并更新预热
if batch_idx % 100 == 0 and batch_idx > 0:
current_shard = batch_idx // 1000 # 假设每 1000 batch 一个 shard
warmer.evict_old(current_shard, shard_paths)
warming_task = asyncio.create_task(
warmer.warm_for_epoch(
shard_paths[current_shard:current_shard+3]
)
)
await warming_taskCheckpoint 缓存策略
模型 checkpoint 是另一种需要缓存的数据。一个 70B 模型的 fp16 checkpoint 约 140GB,从 S3 恢复需要几分钟,从本地 NVMe 恢复只需几秒:
class CheckpointCache:
"""Checkpoint 多级缓存"""
def __init__(self, s3_bucket: str, local_dir: str = "/mnt/nvme/ckpt"):
self.s3_bucket = s3_bucket
self.local_dir = Path(local_dir)
self.local_dir.mkdir(parents=True, exist_ok=True)
def save_checkpoint(self, model_state: dict, step: int):
"""保存 checkpoint:先写本地,再异步上传 S3"""
local_path = self.local_dir / f"checkpoint-{step}.pt"
# 1. 同步写本地(快)
torch.save(model_state, local_path)
# 2. 异步上传 S3(后台)
import subprocess
subprocess.Popen(
["s5cmd", "cp", str(local_path),
f"s3://{self.s3_bucket}/checkpoints/checkpoint-{step}.pt"],
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
)
def load_checkpoint(self, step: int) -> dict:
"""加载 checkpoint:优先本地,其次 S3"""
local_path = self.local_dir / f"checkpoint-{step}.pt"
if local_path.exists():
print(f"Loading from local cache: {local_path}")
return torch.load(local_path)
print(f"Downloading from S3...")
import subprocess
subprocess.run([
"s5cmd", "cp",
f"s3://{self.s3_bucket}/checkpoints/checkpoint-{step}.pt",
str(local_path)
], check=True)
return torch.load(local_path)
def cleanup(self, keep_steps: list[int]):
"""只保留指定 step 的 checkpoint"""
for f in self.local_dir.glob("checkpoint-*.pt"):
step = int(f.stem.split("-")[1])
if step not in keep_steps:
f.unlink()生产环境关键指标:监控 数据加载吞吐量(GB/s)和 GPU 利用率(nvidia-smi)。如果 GPU 利用率低于 90% 且数据加载是瓶颈,说明你的存储层次或缓存策略需要优化。一个简单的诊断方法:把数据集完全加载到内存(如果内存够大),看 GPU 利用率是否提升。如果是,问题就在存储/加载层。
小结
存储架构是 AI 基础设施中组件最多、权衡最复杂的一层。本章覆盖了五个核心主题:
- 对象存储——AI 数据的默认存储层。S3 Express One Zone 将延迟降到了亚毫秒级,让”训练数据直接存 S3”成为可能。
- 分布式文件系统——大规模训练的标准配置。Lustre 追求极致吞吐,JuiceFS 提供云原生 POSIX 兼容方案。
- 向量数据库——RAG 和语义搜索的基础设施。HNSW 索引是通用首选,Milvus/Qdrant 是生产级选择。
- Feature Store——消除训练-推理偏差的关键工具。Feast 提供了”一次定义,离线在线通用”的能力。
- 多级缓存——从 GPU HBM 到 S3 的五级存储层次,数据预热和 LRU 淘汰是核心策略。
存储架构设计的核心原则:让热数据离 GPU 尽可能近,让冷数据尽可能便宜。两者之间的缓存层和预热策略决定了你的 GPU 利用率和训练成本。
延伸阅读
- Amazon S3 Performance Guidelines — AWS 官方 S3 性能最佳实践
- Lustre Operations Manual —
https://doc.lustre.org/ - JuiceFS 文档 —
https://juicefs.com/docs/ - Milvus 教程 —
https://milvus.io/docs/ - Feast Documentation —
https://docs.feast.dev/ - FAISS: The Missing Manual (Chip Huyen) — 向量检索的工程实践