第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 场景下的典型用法

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[容灾恢复]

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 keys

S3 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',
        }
    }
)
Tip

性能对比:标准 S3 的 GET 延迟约 30-100ms(p99),S3 Express One Zone 约 5-10ms(p99)。对于训练数据加载,这意味着你可以在不复制到本地的情况下直接从 S3 读取 shard。

Warning

成本提醒: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/jfs
Tip

JuiceFS 的杀手锏:和 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 极低 超大规模(内存不够)
Tip

实践经验:对于大多数应用(< 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,
)
Warning

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_task

Checkpoint 缓存策略

模型 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()
Tip

生产环境关键指标:监控 数据加载吞吐量(GB/s)和 GPU 利用率(nvidia-smi)。如果 GPU 利用率低于 90% 且数据加载是瓶颈,说明你的存储层次或缓存策略需要优化。一个简单的诊断方法:把数据集完全加载到内存(如果内存够大),看 GPU 利用率是否提升。如果是,问题就在存储/加载层。

小结

存储架构是 AI 基础设施中组件最多、权衡最复杂的一层。本章覆盖了五个核心主题:

  1. 对象存储——AI 数据的默认存储层。S3 Express One Zone 将延迟降到了亚毫秒级,让”训练数据直接存 S3”成为可能。
  2. 分布式文件系统——大规模训练的标准配置。Lustre 追求极致吞吐,JuiceFS 提供云原生 POSIX 兼容方案。
  3. 向量数据库——RAG 和语义搜索的基础设施。HNSW 索引是通用首选,Milvus/Qdrant 是生产级选择。
  4. Feature Store——消除训练-推理偏差的关键工具。Feast 提供了”一次定义,离线在线通用”的能力。
  5. 多级缓存——从 GPU HBM 到 S3 的五级存储层次,数据预热和 LRU 淘汰是核心策略。

存储架构设计的核心原则:让热数据离 GPU 尽可能近,让冷数据尽可能便宜。两者之间的缓存层和预热策略决定了你的 GPU 利用率和训练成本。

延伸阅读

  • Amazon S3 Performance Guidelines — AWS 官方 S3 性能最佳实践
  • Lustre Operations Manualhttps://doc.lustre.org/
  • JuiceFS 文档https://juicefs.com/docs/
  • Milvus 教程https://milvus.io/docs/
  • Feast Documentationhttps://docs.feast.dev/
  • FAISS: The Missing Manual (Chip Huyen) — 向量检索的工程实践