一、引言:数据管理面临的挑战
在AI平台的整个生命周期中,数据管理是决定模型质量的基础环节。然而,随着业务规模的扩大,数据管理面临着前所未有的挑战:数据集规模动辄TB级别,版本迭代频繁,数据血缘关系复杂多样,质量评估缺乏标准化流程。这些挑战不仅影响算法工程师的日常工作效率,更直接关系到模型训练效果的稳定性和可追溯性。
本文将从实际工程实践出发,详细介绍我们在AI平台上构建的数据管理体系,重点涵盖数据集版本化方案、数据血缘追踪架构、数据质量评估体系以及完整的技术实现代码。
二、数据管理挑战分析
2.1 数据管理的核心挑战
AI平台的数据管理与传统软件的数据管理有着本质区别。传统软件的数据通常是结构化的、相对静态的,而AI平台的数据呈现出"四多"特征:
数据来源多:训练数据可能来自多个渠道,包括业务数据库、用户行为日志、外部数据采购、合成数据生成等。不同来源的数据格式、质量、时效性各不相同。
数据格式多:现代AI平台需要处理文本、图像、音频、视频等多种模态的数据。每种模态都有其独特的处理流程和存储格式。
版本迭代多:数据是持续更新的,一条好的数据可能需要反复标注、清洗、增强。一个成熟的数据集生命周期中可能产生数十个版本。
数据关系多:上游数据经过加工处理后产出下游数据集,形成复杂的依赖关系。当上游数据发生变更时,需要能够追踪影响范围。
┌─────────────────────────────────────────────────────────────────────────┐
│ AI平台数据管理全景 │
├─────────────────────────────────────────────────────────────────────────┤
│ │
│ ┌─────────────────────────────────────────────────────────────────┐ │
│ │ 数据来源层 │ │
│ │ │ │
│ │ ┌─────────┐ ┌─────────┐ ┌─────────┐ ┌─────────┐ │ │
│ │ │ 业务数据 │ │ 日志数据 │ │ 外部数据 │ │ 合成数据 │ │ │
│ │ │ (MySQL) │ │ (Kafka) │ │ (API) │ │ (生成) │ │ │
│ │ └────┬────┘ └────┬────┘ └────┬────┘ └────┬────┘ │ │
│ │ │ │ │ │ │ │
│ └────────┼────────────┼────────────┼────────────┼───────────────────┘ │
│ │ │ │ │ │
│ ▼ ▼ ▼ ▼ │
│ ┌─────────────────────────────────────────────────────────────────┐ │
│ │ 数据处理层 │ │
│ │ │ │
│ │ ┌─────────┐ ┌─────────┐ ┌─────────┐ ┌─────────┐ │ │
│ │ │ 数据清洗 │ │ 数据标注 │ │ 数据增强 │ │ 数据转换 │ │ │
│ │ └─────────┘ └─────────┘ └─────────┘ └─────────┘ │ │
│ │ │ │
│ └─────────────────────────────────────────────────────────────────┘ │
│ │ │ │ │ │
│ ▼ ▼ ▼ ▼ │
│ ┌─────────────────────────────────────────────────────────────────┐ │
│ │ 版本管理层 │ │
│ │ │ │
│ │ ┌─────────────────────────────────────────────────────────┐ │ │
│ │ │ 数据集 v1.0 ──→ 数据集 v1.1 ──→ 数据集 v2.0 │ │ │
│ │ │ │ │ │ │ │ │
│ │ │ ▼ ▼ ▼ │ │ │
│ │ │ [元数据] [元数据] [元数据] │ │ │
│ │ │ [统计信息] [统计信息] [统计信息] │ │ │
│ │ │ [数据指纹] [数据指纹] [数据指纹] │ │ │
│ │ └─────────────────────────────────────────────────────────┘ │ │
│ │ │ │
│ └─────────────────────────────────────────────────────────────────┘ │
│ │ │ │ │ │
│ ▼ ▼ ▼ ▼ │
│ ┌─────────────────────────────────────────────────────────────────┐ │
│ │ 血缘追踪层 │ │
│ │ │ │
│ │ DAG可视化 ──→ 影响分析 ──→ 数据溯源 ──→ 变更告警 │ │
│ │ │ │
│ └─────────────────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────────────────┘2.2 现有方案的痛点
在我们构建统一数据管理平台之前,团队面临以下痛点:
版本混乱:没有统一的版本管理机制,不同团队使用不同命名规则,导致"到底用哪个版本的数据训练"成为老大难问题。
血缘不清:数据经过多次加工后,很难说清楚"这个训练数据集是从哪几个原始数据集演变来的"。当训练效果异常时,排查数据问题的成本极高。
质量玄学:数据质量评估缺乏量化标准,算法工程师主要靠"感觉"判断数据好不好,导致训练效果难以稳定复现。
难以复用:好的标注数据散落在各个项目的私有目录里,其他团队想要复用时找不到、找到了也不敢用(怕数据不一致)。
三、数据集版本化方案
3.1 版本化设计理念
我们的数据集版本化方案借鉴了Git的核心理念,同时针对AI数据的特殊性进行了适配:
每个版本都是不可变的:一旦创建,数据集的某个版本就不能修改。这保证了训练结果的可复现性。
版本之间有依赖关系:新版本通常是在旧版本基础上修改而来,我们保留了这种父子关系。
元数据与数据分离:数据文件存储在对象存储中,元数据(包括统计信息、血缘关系等)存储在数据库中。
内容寻址:每个数据集版本都有唯一的内容指纹(Content Hash),相同的指纹代表相同的内容。
3.2 DVC集成方案
DVC(Data Version Control)是当前最流行的数据版本控制工具,我们的方案以DVC为核心,同时增强了企业级特性:
# dvc.yaml
# DVC配置文件示例 - 定义数据处理流水线
# 这是一个数据处理流水线的DVC配置
# 包含原始数据处理、特征工程、数据集组装等阶段
stages:
# 数据清洗阶段
preprocess:
cmd: python src/data/preprocess.py
deps:
- src/data/preprocess.py
- data/raw/raw_data.jsonl
params:
- preprocess.min_length
- preprocess.max_length
- preprocess.filter_quality
outs:
- data/interim/cleaned_data.jsonl
metrics:
- metrics/cleaning.json:
cache: false
persist: false
# 数据标注阶段
annotate:
cmd: python src/data/annotate.py
deps:
- src/data/annotate.py
- data/interim/cleaned_data.jsonl
params:
- annotate.label_schema
- annotate.quality_threshold
outs:
- data/interim/annotated_data.jsonl
persist: false
# 训练集组装阶段
assemble:
cmd: python src/data/assemble.py
deps:
- src/data/interim/annotated_data.jsonl
- src/data/assemble.py
params:
- assemble.train_ratio
- assemble.val_ratio
- assemble.seed
outs:
- data/processed/train.jsonl
- data/processed/val.jsonl
- data/processed/test.jsonl
metrics:
- metrics/assemble.json:
cache: false
# 数据集打包阶段
package:
cmd: python src/data/package.py
deps:
- data/processed/train.jsonl
- data/processed/val.jsonl
- data/processed/test.jsonl
outs:
- datasets/llm-finetune-v1.0.tar.gz:
cache: true
metric: false
persist: false3.3 对象存储路径策略
┌─────────────────────────────────────────────────────────────────────────┐
│ 数据集存储路径结构 │
├─────────────────────────────────────────────────────────────────────────┤
│ │
│ s3://ai-platform-datasets/ │
│ │ │
│ ├── raw/ # 原始数据 │
│ │ ├── {tenant_id}/ │
│ │ │ ├── {dataset_id}/ │
│ │ │ │ ├── v1.0/ │
│ │ │ │ │ ├── data/ # 原始数据文件 │
│ │ │ │ │ │ ├── shard_0000.jsonl │
│ │ │ │ │ │ ├── shard_0001.jsonl │
│ │ │ │ │ │ └── ... │
│ │ │ │ │ ├── metadata.yaml # DVC元数据 │
│ │ │ │ │ └── .dvc # DVC追踪文件 │
│ │ │ │ └── v1.1/ │
│ │ │ │ └── ... │
│ │ │ └── {another_dataset}/ │
│ │ └── ... │
│ │ │
│ ├── processed/ # 处理后的数据 │
│ │ └── {tenant_id}/ │
│ │ └── {dataset_id}/ │
│ │ └── v2.0/ │
│ │ ├── train/ │
│ │ ├── val/ │
│ │ └── test/ │
│ │ │
│ ├── annotations/ # 标注数据 │
│ │ └── {tenant_id}/ │
│ │ └── {dataset_id}/ │
│ │ └── v1.0/ │
│ │ └── annotations.jsonl │
│ │ │
│ └── checkpoints/ # 数据处理中间结果 │
│ └── {tenant_id}/ │
│ └── {job_id}/ │
│ └── checkpoint.jsonl │
│ │
└─────────────────────────────────────────────────────────────────────────┘四、数据血缘追踪架构
4.1 血缘追踪的核心价值
数据血缘(Data Lineage)是指数据从产生到消费的完整链路。血缘追踪的价值体现在多个方面:
影响分析:当上游数据发生变更时,能够快速评估对下游的影响范围。
问题溯源:当训练效果异常时,能够定位到具体的数据来源。
合规审计:满足数据合规要求,记录数据的来源和加工过程。
自动化流水线:基于血缘关系,可以实现数据处理流程的自动化编排。
┌─────────────────────────────────────────────────────────────────────────┐
│ 数据血缘追踪架构 │
├─────────────────────────────────────────────────────────────────────────┤
│ │
│ ┌─────────────────────────────────────────────────────────────────┐ │
│ │ 血缘图谱层 (Lineage Graph) │ │
│ │ │ │
│ │ 原始数据A ─────┐ │ │
│ │ │ │ │ │
│ │ ▼ │ │ │
│ │ 清洗数据B │ │ │
│ │ │ │ │ │
│ │ ├───────────┤ │ │
│ │ │ │ │ │
│ │ ▼ ▼ │ │
│ │ 标注数据C 增强数据D │ │
│ │ │ │ │ │
│ │ └─────┬─────┘ │ │
│ │ │ │ │
│ │ ▼ │ │
│ │ 训练数据集E │ │
│ │ │ │ │
│ │ ┌─────┴─────┐ │ │
│ │ ▼ ▼ │ │
│ │ 模型M1 模型M2 │ │
│ │ │ │
│ └─────────────────────────────────────────────────────────────────┘ │
│ │ │
│ ▼ │
│ ┌─────────────────────────────────────────────────────────────────┐ │
│ │ 图存储引擎 (Neo4j/JanusGraph) │ │
│ │ │ │
│ │ 节点类型: Dataset, Version, File, Model, TrainingJob │ │
│ │ 边类型: DERIVED_FROM, USES, GENERATES, CONTAINS │ │
│ │ │ │
│ └─────────────────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────────────────┘4.2 血缘追踪图存储实现
# lineage/graph_store.py
# 数据血缘追踪 - 图存储实现
from __future__ import annotations
import uuid
from dataclasses import dataclass, field
from datetime import datetime
from enum import Enum
from typing import List, Optional, Dict, Any, Set, Tuple
from abc import ABC, abstractmethod
import logging
logger = logging.getLogger(__name__)
class NodeType(Enum):
"""血缘图节点类型"""
DATASET = "dataset"
DATASET_VERSION = "dataset_version"
FILE = "file"
TRAINING_JOB = "training_job"
MODEL = "model"
MODEL_VERSION = "model_version"
class EdgeType(Enum):
"""血缘图边类型"""
DERIVED_FROM = "derived_from" # 派生自
USES = "uses" # 使用
GENERATES = "generates" # 生成
CONTAINS = "contains" # 包含
PART_OF = "part_of" # 属于
VERSION_OF = "version_of" # 版本关系
@dataclass
class LineageNode:
"""
血缘节点
代表数据流中的一个实体
"""
id: str = field(default_factory=lambda: str(uuid.uuid4()))
node_type: NodeType = NodeType.DATASET
name: str = ""
tenant_id: str = ""
# 关联的资源ID
resource_id: Optional[str] = None
# 属性
properties: Dict[str, Any] = field(default_factory=dict)
# 版本信息(如果是版本节点)
version: Optional[str] = None
parent_id: Optional[str] = None # 父版本ID
# 元数据
created_at: datetime = field(default_factory=datetime.now)
created_by: str = ""
def to_dict(self) -> Dict[str, Any]:
"""序列化为字典"""
return {
'id': self.id,
'node_type': self.node_type.value,
'name': self.name,
'tenant_id': self.tenant_id,
'resource_id': self.resource_id,
'properties': self.properties,
'version': self.version,
'parent_id': self.parent_id,
'created_at': self.created_at.isoformat(),
'created_by': self.created_by
}
@dataclass
class LineageEdge:
"""
血缘边
代表两个节点之间的关系
"""
id: str = field(default_factory=lambda: str(uuid.uuid4()))
edge_type: EdgeType = EdgeType.DERIVED_FROM
# 源节点和目标节点
source_id: str = ""
target_id: str = ""
# 边的属性
properties: Dict[str, Any] = field(default_factory=dict)
# 血缘类型:UPSTREAM(上游)或 DOWNSTREAM(下游)
direction: str = "UPSTREAM"
created_at: datetime = field(default_factory=datetime.now)
def to_dict(self) -> Dict[str, Any]:
return {
'id': self.id,
'edge_type': self.edge_type.value,
'source_id': self.source_id,
'target_id': self.target_id,
'properties': self.properties,
'direction': self.direction,
'created_at': self.created_at.isoformat()
}
class LineageGraph:
"""
血缘图
维护节点和边的集合,提供查询接口
"""
def __init__(self):
self._nodes: Dict[str, LineageNode] = {}
self._edges: Dict[str, LineageEdge] = {}
# 索引优化
self._nodes_by_type: Dict[NodeType, Set[str]] = {}
self._nodes_by_resource: Dict[Tuple[str, str], str] = {} # (tenant_id, resource_id) -> node_id
self._edges_by_source: Dict[str, Set[str]] = {} # source_id -> edge_ids
self._edges_by_target: Dict[str, Set[str]] = {} # target_id -> edge_ids
def add_node(self, node: LineageNode) -> None:
"""添加节点"""
self._nodes[node.id] = node
# 更新索引
if node.node_type not in self._nodes_by_type:
self._nodes_by_type[node.node_type] = set()
self._nodes_by_type[node.node_type].add(node.id)
if node.resource_id:
key = (node.tenant_id, node.resource_id)
self._nodes_by_resource[key] = node.id
logger.info(f"添加血缘节点: {node.id} ({node.node_type.value})")
def add_edge(self, edge: LineageEdge) -> None:
"""添加边"""
# 验证节点存在
if edge.source_id not in self._nodes:
raise ValueError(f"源节点不存在: {edge.source_id}")
if edge.target_id not in self._nodes:
raise ValueError(f"目标节点不存在: {edge.target_id}")
self._edges[edge.id] = edge
# 更新索引
if edge.source_id not in self._edges_by_source:
self._edges_by_source[edge.source_id] = set()
self._edges_by_source[edge.source_id].add(edge.id)
if edge.target_id not in self._edges_by_target:
self._edges_by_target[edge.target_id] = set()
self._edges_by_target[edge.target_id].add(edge.id)
logger.info(
f"添加血缘边: {edge.source_id} --[{edge.edge_type.value}]--> {edge.target_id}"
)
def get_node(self, node_id: str) -> Optional[LineageNode]:
"""获取节点"""
return self._nodes.get(node_id)
def get_node_by_resource(self, tenant_id: str, resource_id: str) -> Optional[LineageNode]:
"""根据资源ID获取节点"""
key = (tenant_id, resource_id)
node_id = self._nodes_by_resource.get(key)
return self._nodes.get(node_id)
def get_upstream(self, node_id: str, max_depth: int = 10) -> List[LineageNode]:
"""
获取上游血缘
从当前节点向上追溯所有祖先节点
"""
if node_id not in self._nodes:
return []
visited = set()
result = []
queue = [(node_id, 0)]
while queue:
current_id, depth = queue.pop(0)
if current_id in visited or depth > max_depth:
continue
visited.add(current_id)
# 获取所有指向当前节点的边(即当前节点的上游)
edge_ids = self._edges_by_target.get(current_id, set())
for edge_id in edge_ids:
edge = self._edges.get(edge_id)
if edge:
source_node = self._nodes.get(edge.source_id)
if source_node and source_node.id not in visited:
result.append(source_node)
queue.append((source_node.id, depth + 1))
return result
def get_downstream(self, node_id: str, max_depth: int = 10) -> List[LineageNode]:
"""
获取下游血缘
从当前节点向下追溯所有派生节点
"""
if node_id not in self._nodes:
return []
visited = set()
result = []
queue = [(node_id, 0)]
while queue:
current_id, depth = queue.pop(0)
if current_id in visited or depth > max_depth:
continue
visited.add(current_id)
# 获取所有从当前节点发出的边(即当前节点的下游)
edge_ids = self._edges_by_source.get(current_id, set())
for edge_id in edge_ids:
edge = self._edges.get(edge_id)
if edge:
target_node = self._nodes.get(edge.target_id)
if target_node and target_node.id not in visited:
result.append(target_node)
queue.append((target_node.id, depth + 1))
return result
def get_full_lineage(self, node_id: str) -> Dict[str, Any]:
"""
获取完整血缘
返回包含上下游的完整血缘信息
"""
node = self._nodes.get(node_id)
if not node:
return {}
upstream = self.get_upstream(node_id)
downstream = self.get_downstream(node_id)
return {
'node': node.to_dict(),
'upstream': [n.to_dict() for n in upstream],
'downstream': [n.to_dict() for n in downstream],
'statistics': {
'upstream_count': len(upstream),
'downstream_count': len(downstream),
'total_impact': len(upstream) + len(downstream)
}
}
def find_impacted_nodes(
self,
node_id: str,
target_types: Optional[List[NodeType]] = None
) -> List[LineageNode]:
"""
查找受影响的节点
当某个节点变更时,查找所有下游受影响的目标类型节点
"""
downstream = self.get_downstream(node_id)
if target_types:
downstream = [n for n in downstream if n.node_type in target_types]
return downstream
class GraphDatabase(ABC):
"""
图数据库抽象基类
定义图数据库的接口
"""
@abstractmethod
def save_node(self, node: LineageNode) -> None:
pass
@abstractmethod
def save_edge(self, edge: LineageEdge) -> None:
pass
@abstractmethod
def get_node(self, node_id: str) -> Optional[LineageNode]:
pass
@abstractmethod
def get_upstream(self, node_id: str, max_depth: int = 10) -> List[LineageNode]:
pass
@abstractmethod
def get_downstream(self, node_id: str, max_depth: int = 10) -> List[LineageNode]:
pass
@abstractmethod
def find_path(self, source_id: str, target_id: str) -> Optional[List[str]]:
"""查找两个节点之间的路径"""
pass
class InMemoryGraphDatabase(GraphDatabase):
"""
内存图数据库实现
用于测试和小规模部署
生产环境应使用Neo4j或JanusGraph
"""
def __init__(self):
self._graph = LineageGraph()
def save_node(self, node: LineageNode) -> None:
self._graph.add_node(node)
def save_edge(self, edge: LineageEdge) -> None:
self._graph.add_edge(edge)
def get_node(self, node_id: str) -> Optional[LineageNode]:
return self._graph.get_node(node_id)
def get_upstream(self, node_id: str, max_depth: int = 10) -> List[LineageNode]:
return self._graph.get_upstream(node_id, max_depth)
def get_downstream(self, node_id: str, max_depth: int = 10) -> List[LineageNode]:
return self._graph.get_downstream(node_id, max_depth)
def find_path(self, source_id: str, target_id: str) -> Optional[List[str]]:
"""BFS查找路径"""
visited = set()
queue = [(source_id, [source_id])]
while queue:
current, path = queue.pop(0)
if current == target_id:
return path
if current in visited:
continue
visited.add(current)
# 获取下游节点
for node in self._graph.get_downstream(current):
if node.id not in visited:
queue.append((node.id, path + [node.id]))
return None4.3 血缘追踪服务
# lineage/service.py
# 数据血缘追踪服务
from typing import List, Optional, Dict, Any
from dataclasses import dataclass
import logging
from lineage.graph_store import (
LineageGraph,
LineageNode,
LineageEdge,
NodeType,
EdgeType,
GraphDatabase,
)
from datasets.models import Dataset, DatasetVersion
logger = logging.getLogger(__name__)
class LineageTracker:
"""
血缘追踪器
负责记录和查询数据血缘关系
"""
def __init__(self, graph_db: GraphDatabase):
self._graph_db = graph_db
def track_dataset_version_creation(
self,
dataset_id: str,
version_id: str,
version: str,
tenant_id: str,
parent_version_id: Optional[str] = None,
source_dataset_ids: Optional[List[str]] = None,
created_by: str = ""
) -> str:
"""
追踪数据集版本创建
当创建新的数据集版本时调用此方法记录血缘
Args:
dataset_id: 数据集ID
version_id: 版本ID
version: 版本号
tenant_id: 租户ID
parent_version_id: 父版本ID(如果是从旧版本派生的)
source_dataset_ids: 源数据集ID列表(如果是合并多个数据集)
created_by: 创建者
Returns:
血缘节点ID
"""
# 创建数据集版本节点
node = LineageNode(
node_type=NodeType.DATASET_VERSION,
name=f"{dataset_id}:{version}",
tenant_id=tenant_id,
resource_id=version_id,
version=version,
properties={
'dataset_id': dataset_id,
'version': version,
'created_by': created_by
},
created_by=created_by
)
# 如果有父版本,建立版本关系
if parent_version_id:
node.parent_id = parent_version_id
# 获取父版本节点
parent_node = self._graph_db.get_node_by_resource(
tenant_id, parent_version_id
)
if parent_node:
edge = LineageEdge(
edge_type=EdgeType.VERSION_OF,
source_id=node.id,
target_id=parent_node.id,
properties={'relationship': 'derived_from'}
)
self._graph_db.save_edge(edge)
# 如果有源数据集,建立派生关系
if source_dataset_ids:
for source_id in source_dataset_ids:
source_node = self._graph_db.get_node_by_resource(
tenant_id, source_id
)
if source_node:
edge = LineageEdge(
edge_type=EdgeType.DERIVED_FROM,
source_id=node.id,
target_id=source_node.id,
properties={'transformation': 'merge_or_process'}
)
self._graph_db.save_edge(edge)
self._graph_db.save_node(node)
logger.info(
f"记录数据集版本血缘: dataset={dataset_id}, "
f"version={version}, node_id={node.id}"
)
return node.id
def track_training_data_usage(
self,
dataset_version_id: str,
training_job_id: str,
tenant_id: str
) -> None:
"""
追踪训练数据使用
当训练任务使用某个数据集版本时调用
"""
# 获取数据集版本节点
dataset_node = self._graph_db.get_node_by_resource(
tenant_id, dataset_version_id
)
# 创建训练任务节点
job_node = LineageNode(
node_type=NodeType.TRAINING_JOB,
name=f"training_{training_job_id}",
tenant_id=tenant_id,
resource_id=training_job_id,
properties={
'job_id': training_job_id,
'dataset_version_id': dataset_version_id
}
)
self._graph_db.save_node(job_node)
# 建立使用关系
if dataset_node:
edge = LineageEdge(
edge_type=EdgeType.USES,
source_id=job_node.id,
target_id=dataset_node.id,
properties={'usage': 'training_data'}
)
self._graph_db.save_edge(edge)
def track_model_from_training(
self,
training_job_id: str,
model_version_id: str,
model_name: str,
tenant_id: str
) -> None:
"""
追踪模型生成
记录从训练任务到模型的生成关系
"""
# 获取训练任务节点
job_node = self._graph_db.get_node_by_resource(
tenant_id, training_job_id
)
# 创建模型版本节点
model_node = LineageNode(
node_type=NodeType.MODEL_VERSION,
name=model_name,
tenant_id=tenant_id,
resource_id=model_version_id,
properties={
'model_version_id': model_version_id,
'model_name': model_name
}
)
self._graph_db.save_node(model_node)
# 建立生成关系
if job_node:
edge = LineageEdge(
edge_type=EdgeType.GENERATES,
source_id=job_node.id,
target_id=model_node.id,
properties={'result': 'trained_model'}
)
self._graph_db.save_edge(edge)
def analyze_data_impact(
self,
dataset_version_id: str,
tenant_id: str
) -> Dict[str, Any]:
"""
分析数据变更影响
当某个数据集版本发生变更时,评估对下游的影响
Returns:
影响分析报告
"""
# 获取数据集节点
dataset_node = self._graph_db.get_node_by_resource(
tenant_id, dataset_version_id
)
if not dataset_node:
return {'error': '数据集版本节点不存在'}
# 查找所有下游影响
impacted_jobs = self._graph_db.find_impacted_nodes(
dataset_node.id,
target_types=[NodeType.TRAINING_JOB]
)
impacted_models = self._graph_db.find_impacted_nodes(
dataset_node.id,
target_types=[NodeType.MODEL_VERSION]
)
# 获取完整血缘信息
lineage = self._graph_db.get_full_lineage(dataset_node.id)
return {
'changed_node': dataset_node.to_dict(),
'impact_summary': {
'affected_training_jobs': len(impacted_jobs),
'affected_models': len(impacted_models),
'total_impact': len(impacted_jobs) + len(impacted_models)
},
'impacted_jobs': [
{
'job_id': n.resource_id,
'name': n.name,
'created_at': n.created_at.isoformat()
}
for n in impacted_jobs
],
'impacted_models': [
{
'model_id': n.resource_id,
'name': n.name,
'created_at': n.created_at.isoformat()
}
for n in impacted_models
],
'upstream_lineage': lineage.get('upstream', []),
'recommendation': self._generate_impact_recommendation(
impacted_jobs, impacted_models
)
}
def trace_data_origin(
self,
model_version_id: str,
tenant_id: str
) -> Dict[str, Any]:
"""
追溯数据溯源
从模型追溯到原始训练数据
Returns:
数据溯源报告
"""
# 获取模型节点
model_node = self._graph_db.get_node_by_resource(
tenant_id, model_version_id
)
if not model_node:
return {'error': '模型节点不存在'}
# 获取完整上游血缘
upstream = self._graph_db.get_upstream(model_node.id, max_depth=20)
# 筛选数据集节点
dataset_nodes = [
n for n in upstream
if n.node_type == NodeType.DATASET_VERSION
]
# 构建数据来源树
source_tree = self._build_source_tree(model_node, upstream)
return {
'model': model_node.to_dict(),
'source_datasets': [
{
'dataset_id': n.properties.get('dataset_id'),
'version': n.version,
'created_at': n.created_at.isoformat()
}
for n in dataset_nodes
],
'source_tree': source_tree
}
def _build_source_tree(
self,
root: LineageNode,
upstream: List[LineageNode]
) -> Dict[str, Any]:
"""构建数据来源树"""
# 构建节点映射
node_map = {n.id: n for n in upstream}
node_map[root.id] = root
def build_node(node: LineageNode, depth: int = 0) -> Dict[str, Any]:
if depth > 10: # 防止无限递归
return {}
result = {
'id': node.id,
'type': node.node_type.value,
'name': node.name,
'version': node.version,
'children': []
}
# 查找上游节点
for upstream_node in upstream:
edge = self._find_edge(upstream_node.id, node.id)
if edge:
result['children'].append(
build_node(upstream_node, depth + 1)
)
return result
return build_node(root)
def _find_edge(self, source_id: str, target_id: str) -> Optional[LineageEdge]:
"""查找两个节点之间的边"""
# 简化实现,实际应查询图数据库
return None
def _generate_impact_recommendation(
self,
impacted_jobs: List[LineageNode],
impacted_models: List[LineageNode]
) -> str:
"""生成影响分析建议"""
if not impacted_jobs and not impacted_models:
return "该数据集版本没有下游依赖,可以安全变更。"
recommendations = []
if impacted_models:
recommendations.append(
f"该数据集被 {len(impacted_models)} 个模型使用,"
"变更后建议重新训练这些模型。"
)
if impacted_jobs:
recommendations.append(
f"该数据集被 {len(impacted_jobs)} 个训练任务使用,"
"历史训练结果可能受到影响。"
)
return " ".join(recommendations)五、数据质量评估体系
5.1 质量评估框架
数据质量是模型效果的基础。我们建立了一套多维度的数据质量评估框架:
┌─────────────────────────────────────────────────────────────────────────┐
│ 数据质量评估框架 │
├─────────────────────────────────────────────────────────────────────────┤
│ │
│ ┌─────────────────────────────────────────────────────────────────┐ │
│ │ 质量维度 (Quality Dimensions) │ │
│ │ │ │
│ │ ┌─────────────┐ ┌─────────────┐ ┌─────────────┐ │ │
│ │ │ 完整性 │ │ 一致性 │ │ 准确性 │ │ │
│ │ │ Completeness│ │ Consistency │ │ Accuracy │ │ │
│ │ │ │ │ │ │ │ │ │
│ │ │ - 空值比例 │ │ - 格式统一 │ │ - 噪声检测 │ │ │
│ │ │ - 缺失字段 │ │ - 编码一致 │ │ - 异常值 │ │ │
│ │ │ - 重复记录 │ │ - 范围约束 │ │ - 标注错误 │ │ │
│ │ └─────────────┘ └─────────────┘ └─────────────┘ │ │
│ │ │ │
│ │ ┌─────────────┐ ┌─────────────┐ ┌─────────────┐ │ │
│ │ │ 时效性 │ │ 可用性 │ │ 多样性 │ │ │
│ │ │ Timeliness │ │ Availability │ │ Diversity │ │ │
│ │ │ │ │ │ │ │ │ │
│ │ │ - 数据新鲜度 │ │ - 读取成功率 │ │ - 分布均匀性 │ │ │
│ │ │ - 更新频率 │ │ - 响应时间 │ │ - 类别覆盖 │ │ │
│ │ │ - 过期检测 │ │ - 并发能力 │ │ - 长尾分布 │ │ │
│ │ └─────────────┘ └─────────────┘ └─────────────┘ │ │
│ │ │ │
│ └─────────────────────────────────────────────────────────────────┘ │
│ │ │
│ ▼ │
│ ┌─────────────────────────────────────────────────────────────────┐ │
│ │ 质量评分 (Quality Score) │ │
│ │ │ │
│ │ 综合得分 = Σ(维度权重 × 维度得分) │ │
│ │ │ │
│ │ 评分等级: │ │
│ │ - A (90-100): 优秀 - 可直接用于生产训练 │ │
│ │ - B (70-89): 良好 - 可使用,建议小幅优化 │ │
│ │ - C (50-69): 一般 - 可使用,需要关注质量问题 │ │
│ │ - D (30-49): 较差 - 谨慎使用,建议优先处理 │ │
│ │ - F (0-29): 差 - 不建议使用 │ │
│ │ │ │
│ └─────────────────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────────────────┘5.2 质量评估服务实现
# quality/assessor.py
# 数据质量评估服务
from __future__ import annotations
import json
import hashlib
from dataclasses import dataclass, field
from datetime import datetime
from typing import List, Optional, Dict, Any, Callable
from abc import ABC, abstractmethod
import logging
import statistics
from datasets.models import DatasetVersion
logger = logging.getLogger(__name__)
@dataclass
class QualityDimension:
"""质量维度"""
name: str
description: str
weight: float # 权重,0-1之间
score: Optional[float] = None
details: Dict[str, Any] = field(default_factory=dict)
@dataclass
class QualityReport:
"""质量评估报告"""
dataset_version_id: str
evaluated_at: datetime = field(default_factory=datetime.now)
# 各维度得分
completeness: Optional[QualityDimension] = None
consistency: Optional[QualityDimension] = None
accuracy: Optional[QualityDimension] = None
timeliness: Optional[QualityDimension] = None
diversity: Optional[QualityDimension] = None
# 综合得分
overall_score: float = 0.0
grade: str = "F"
# 统计信息
total_records: int = 0
total_size_bytes: int = 0
file_count: int = 0
# 问题列表
issues: List[Dict[str, Any]] = field(default_factory=list)
# 建议
recommendations: List[str] = field(default_factory=list)
def to_dict(self) -> Dict[str, Any]:
return {
'dataset_version_id': self.dataset_version_id,
'evaluated_at': self.evaluated_at.isoformat(),
'overall_score': self.overall_score,
'grade': self.grade,
'dimensions': {
'completeness': self._dim_to_dict(self.completeness),
'consistency': self._dim_to_dict(self.consistency),
'accuracy': self._dim_to_dict(self.accuracy),
'timeliness': self._dim_to_dict(self.timeliness),
'diversity': self._dim_to_dict(self.diversity),
},
'statistics': {
'total_records': self.total_records,
'total_size_bytes': self.total_size_bytes,
'file_count': self.file_count,
},
'issues': self.issues,
'recommendations': self.recommendations
}
def _dim_to_dict(self, dim: Optional[QualityDimension]) -> Optional[Dict[str, Any]]:
if dim is None:
return None
return {
'name': dim.name,
'score': dim.score,
'weight': dim.weight,
'details': dim.details
}
class QualityAssessor:
"""
数据质量评估器
评估数据集的质量并生成报告
"""
# 质量等级阈值
GRADE_THRESHOLDS = {
'A': 90,
'B': 70,
'C': 50,
'D': 30,
}
def __init__(self):
self._checks: List[QualityCheck] = []
def register_check(self, check: QualityCheck) -> None:
"""注册质量检查"""
self._checks.append(check)
async def assess(self, dataset_version: DatasetVersion) -> QualityReport:
"""
评估数据集质量
Args:
dataset_version: 数据集版本对象
Returns:
质量评估报告
"""
logger.info(f"开始评估数据集: {dataset_version.id}")
report = QualityReport(
dataset_version_id=str(dataset_version.id)
)
# 收集统计信息
stats = await self._collect_statistics(dataset_version)
report.total_records = stats.get('record_count', 0)
report.total_size_bytes = stats.get('size_bytes', 0)
report.file_count = stats.get('file_count', 0)
# 执行各项检查
for check in self._checks:
try:
result = await check.execute(dataset_version, stats)
await self._apply_result(report, result)
except Exception as e:
logger.error(f"质量检查执行失败: {check.name}, error={e}")
report.issues.append({
'check': check.name,
'severity': 'error',
'message': f"检查执行失败: {str(e)}"
})
# 计算综合得分
self._calculate_overall_score(report)
# 生成建议
self._generate_recommendations(report)
logger.info(
f"质量评估完成: {dataset_version.id}, "
f"score={report.overall_score}, grade={report.grade}"
)
return report
async def _collect_statistics(
self,
dataset_version: DatasetVersion
) -> Dict[str, Any]:
"""收集数据集统计信息"""
# TODO: 实际实现应读取数据集文件并计算统计信息
return {
'record_count': dataset_version.record_count,
'size_bytes': dataset_version.total_size_bytes,
'file_count': dataset_version.file_count,
'data_type': dataset_version.data_type,
}
async def _apply_result(
self,
report: QualityReport,
result: CheckResult
) -> None:
"""应用检查结果"""
dimension_map = {
'completeness': 'completeness',
'consistency': 'consistency',
'accuracy': 'accuracy',
'timeliness': 'timeliness',
'diversity': 'diversity',
}
dim_attr = dimension_map.get(result.dimension)
if dim_attr:
dim = QualityDimension(
name=result.dimension,
description=result.description,
weight=result.weight,
score=result.score,
details=result.details
)
setattr(report, dim_attr, dim)
if result.issues:
report.issues.extend(result.issues)
def _calculate_overall_score(self, report: QualityReport) -> None:
"""计算综合得分"""
dimensions = [
report.completeness,
report.consistency,
report.accuracy,
report.timeliness,
report.diversity,
]
total_weight = 0.0
weighted_sum = 0.0
for dim in dimensions:
if dim and dim.score is not None:
weighted_sum += dim.score * dim.weight
total_weight += dim.weight
if total_weight > 0:
report.overall_score = round(weighted_sum / total_weight, 2)
else:
report.overall_score = 0.0
# 确定等级
report.grade = self._determine_grade(report.overall_score)
def _determine_grade(self, score: float) -> str:
"""根据得分确定等级"""
for grade, threshold in sorted(
self.GRADE_THRESHOLDS.items(),
key=lambda x: x[1],
reverse=True
):
if score >= threshold:
return grade
return 'F'
def _generate_recommendations(self, report: QualityReport) -> None:
"""生成改进建议"""
recommendations = []
# 根据各维度得分生成建议
dimensions = [
(report.completeness, "完整性"),
(report.consistency, "一致性"),
(report.accuracy, "准确性"),
(report.timeliness, "时效性"),
(report.diversity, "多样性"),
]
for dim, dim_name in dimensions:
if dim and dim.score is not None:
if dim.score < 50:
recommendations.append(
f"【{dim_name}】得分较低({dim.score}分),"
f"建议优先优化。详情: {dim.details.get('message', '')}"
)
elif dim.score < 70:
recommendations.append(
f"【{dim_name}】有提升空间({dim.score}分),"
f"建议关注: {dim.details.get('message', '')}"
)
if not recommendations:
recommendations.append("数据集质量良好,可直接用于训练。")
report.recommendations = recommendations
@dataclass
class CheckResult:
"""检查结果"""
dimension: str
description: str
weight: float
score: float # 0-100
details: Dict[str, Any] = field(default_factory=dict)
issues: List[Dict[str, Any]] = field(default_factory=list)
class QualityCheck(ABC):
"""质量检查基类"""
def __init__(self, name: str):
self.name = name
@abstractmethod
async def execute(
self,
dataset_version: DatasetVersion,
stats: Dict[str, Any]
) -> CheckResult:
"""执行检查"""
pass
class CompletenessCheck(QualityCheck):
"""完整性检查"""
def __init__(self):
super().__init__("CompletenessCheck")
async def execute(
self,
dataset_version: DatasetVersion,
stats: Dict[str, Any]
) -> CheckResult:
issues = []
details = {}
# 检查空值比例(模拟)
null_ratio = 0.02 # 2%空值
if null_ratio > 0.1:
issues.append({
'check': self.name,
'severity': 'warning',
'field': 'multiple',
'message': f"空值比例过高: {null_ratio*100:.1f}%"
})
# 检查重复记录
duplicate_ratio = 0.005 # 0.5%重复
if duplicate_ratio > 0.05:
issues.append({
'check': self.name,
'severity': 'warning',
'field': 'all',
'message': f"重复记录比例较高: {duplicate_ratio*100:.2f}%"
})
# 检查字段完整性
required_fields = ['text', 'label']
missing_fields = [] # 模拟为全部存在
if missing_fields:
issues.append({
'check': self.name,
'severity': 'error',
'field': ', '.join(missing_fields),
'message': "存在缺失的必需字段"
})
# 计算得分
base_score = 100
if null_ratio > 0.01:
base_score -= null_ratio * 200
if duplicate_ratio > 0.001:
base_score -= duplicate_ratio * 500
if missing_fields:
base_score -= len(missing_fields) * 10
score = max(0, min(100, base_score))
details['null_ratio'] = null_ratio
details['duplicate_ratio'] = duplicate_ratio
details['missing_required_fields'] = len(missing_fields)
details['message'] = f"空值比例{null_ratio*100:.2f}%, 重复比例{duplicate_ratio*100:.2f}%"
return CheckResult(
dimension='completeness',
description='检查数据的完整性,包括空值、缺失字段、重复记录等',
weight=0.25,
score=score,
details=details,
issues=issues
)
class DiversityCheck(QualityCheck):
"""多样性检查"""
def __init__(self):
super().__init__("DiversityCheck")
async def execute(
self,
dataset_version: DatasetVersion,
stats: Dict[str, Any]
) -> CheckResult:
issues = []
details = {}
# 检查类别分布(模拟)
label_distribution = {
'positive': 0.45,
'negative': 0.40,
'neutral': 0.15
}
# 检测类别不平衡
max_ratio = max(label_distribution.values())
min_ratio = min(label_distribution.values())
imbalance_ratio = max_ratio / min_ratio if min_ratio > 0 else float('inf')
if imbalance_ratio > 3:
issues.append({
'check': self.name,
'severity': 'warning',
'message': f"类别不平衡较严重,最大/最小比例: {imbalance_ratio:.2f}"
})
# 检查文本长度分布
avg_length = 150 # 模拟
length_std = 80 # 模拟
if length_std < 30:
issues.append({
'check': self.name,
'severity': 'info',
'message': "文本长度分布较为集中,多样性可能不足"
})
# 计算多样性得分
score = 100
# 类别平衡扣分
if imbalance_ratio > 2:
score -= (imbalance_ratio - 2) * 10
# 长度分布加分
if length_std > 50:
score += 10
score = max(0, min(100, score))
details['label_distribution'] = label_distribution
details['imbalance_ratio'] = imbalance_ratio
details['avg_text_length'] = avg_length
details['text_length_std'] = length_std
details['message'] = f"类别不平衡比例: {imbalance_ratio:.2f}"
return CheckResult(
dimension='diversity',
description='检查数据的多样性,包括类别分布、文本长度分布等',
weight=0.2,
score=score,
details=details,
issues=issues
)六、数据集管理API实现
6.1 Django数据集服务实现
# datasets/service.py
# 数据集管理服务
from __future__ import annotations
import os
import uuid
import hashlib
from dataclasses import dataclass, field
from datetime import datetime
from typing import List, Optional, Dict, Any, BinaryIO
from enum import Enum
import logging
from django.db import transaction
from django.conf import settings
from datasets.models import Dataset, DatasetVersion, DatasetFile
from datasets.serializers import (
DatasetCreateSerializer,
DatasetVersionCreateSerializer,
)
from quality.assessor import QualityAssessor, QualityReport
from lineage.service import LineageTracker
from lineage.graph_store import NodeType, EdgeType
logger = logging.getLogger(__name__)
class DatasetStatus(Enum):
"""数据集状态"""
UPLOADING = "uploading"
PROCESSING = "processing"
READY = "ready"
ARCHIVED = "archived"
FAILED = "failed"
@dataclass
class DatasetStatistics:
"""数据集统计信息"""
record_count: int = 0
file_count: int = 0
total_size_bytes: int = 0
# 质量指标
quality_score: float = 0.0
quality_grade: str = "N/A"
# 数据分布
label_distribution: Dict[str, float] = field(default_factory=dict)
avg_text_length: float = 0.0
# 数据指纹
content_hash: str = ""
class DatasetService:
"""
数据集管理服务
负责数据集的创建、上传、版本管理等操作
"""
def __init__(
self,
storage_client, # S3/MinIO客户端
lineage_tracker: LineageTracker,
quality_assessor: QualityAssessor,
):
self.storage_client = storage_client
self.lineage_tracker = lineage_tracker
self.quality_assessor = quality_assessor
@transaction.atomic
def create_dataset(
self,
tenant_id: str,
name: str,
description: str,
data_type: str,
created_by: str,
tags: Optional[List[str]] = None,
metadata: Optional[Dict[str, Any]] = None
) -> Dataset:
"""
创建数据集
Args:
tenant_id: 租户ID
name: 数据集名称
description: 数据集描述
data_type: 数据类型 (text/image/audio/video)
created_by: 创建者
tags: 标签列表
metadata: 元数据
Returns:
创建的数据集对象
"""
dataset = Dataset.objects.create(
tenant_id=tenant_id,
name=name,
description=description,
data_type=data_type,
created_by=created_by,
tags=tags or [],
metadata=metadata or {},
status=DatasetStatus.UPLOADING.value
)
# 创建血缘节点
self.lineage_tracker._graph_db.save_node(
LineageNode(
node_type=NodeType.DATASET,
name=name,
tenant_id=tenant_id,
resource_id=str(dataset.id),
properties={
'name': name,
'data_type': data_type,
'created_by': created_by
}
)
)
logger.info(f"数据集创建成功: {dataset.id}, name={name}")
return dataset
@transaction.atomic
def create_version(
self,
dataset_id: str,
version: str,
created_by: str,
parent_version_id: Optional[str] = None,
source_dataset_ids: Optional[List[str]] = None,
storage_path: Optional[str] = None
) -> DatasetVersion:
"""
创建数据集版本
Args:
dataset_id: 数据集ID
version: 版本号 (如 v1.0.0)
created_by: 创建者
parent_version_id: 父版本ID
source_dataset_ids: 源数据集ID列表
storage_path: 存储路径
Returns:
创建的数据集版本对象
"""
dataset = Dataset.objects.get(id=dataset_id)
# 检查版本是否已存在
if DatasetVersion.objects.filter(
dataset=dataset,
version=version
).exists():
raise ValueError(f"版本 {version} 已存在")
# 生成存储路径
if not storage_path:
storage_path = f"{dataset.tenant_id}/datasets/{dataset.name}/{version}"
version_obj = DatasetVersion.objects.create(
dataset=dataset,
version=version,
storage_path=storage_path,
status=DatasetStatus.UPLOADING.value,
created_by=created_by
)
# 记录血缘
self.lineage_tracker.track_dataset_version_creation(
dataset_id=str(dataset.id),
version_id=str(version_obj.id),
version=version,
tenant_id=dataset.tenant_id,
parent_version_id=parent_version_id,
source_dataset_ids=source_dataset_ids,
created_by=created_by
)
logger.info(
f"数据集版本创建成功: dataset={dataset_id}, "
f"version={version}, version_id={version_obj.id}"
)
return version_obj
@transaction.atomic
async def upload_file(
self,
version_id: str,
file_obj: BinaryIO,
file_name: str,
content_type: str
) -> DatasetFile:
"""
上传数据文件
Args:
version_id: 版本ID
file_obj: 文件对象
file_name: 文件名
content_type: 内容类型
Returns:
创建的文件记录
"""
version = DatasetVersion.objects.select_related('dataset').get(
id=version_id
)
# 生成唯一文件路径
file_id = str(uuid.uuid4())
file_path = f"{version.storage_path}/{file_id}_{file_name}"
# 上传到对象存储
await self.storage_client.upload(
bucket=settings.DATASET_BUCKET,
key=file_path,
data=file_obj,
content_type=content_type
)
# 获取文件大小
file_size = file_obj.seek(0, 2) # 获取文件大小
file_obj.seek(0) # 重置文件指针
# 计算文件哈希
content_hash = self._calculate_hash(file_obj)
file_obj.seek(0)
# 创建文件记录
dataset_file = DatasetFile.objects.create(
version=version,
file_name=file_name,
file_path=file_path,
file_size=file_size,
content_hash=content_hash,
content_type=content_type
)
# 更新版本状态
version.file_count += 1
version.total_size_bytes += file_size
version.save()
logger.info(
f"文件上传成功: version={version_id}, "
f"file={file_name}, size={file_size}"
)
return dataset_file
@transaction.atomic
async def finalize_version(
self,
version_id: str,
record_count: int,
metadata: Optional[Dict[str, Any]] = None
) -> DatasetVersion:
"""
完成版本上传
触发数据质量评估并更新版本状态
Args:
version_id: 版本ID
record_count: 记录数量
metadata: 额外元数据
Returns:
更新后的版本对象
"""
version = DatasetVersion.objects.select_related('dataset').get(
id=version_id
)
# 更新元数据
version.record_count = record_count
if metadata:
version.metadata = {**version.metadata, **metadata}
# 更新状态为处理中
version.status = DatasetStatus.PROCESSING.value
version.save()
# 异步执行质量评估
try:
quality_report = await self.quality_assessor.assess(version)
# 更新版本质量信息
version.quality_score = quality_report.overall_score
version.quality_grade = quality_report.grade
version.quality_report = quality_report.to_dict()
# 更新状态为就绪
version.status = DatasetStatus.READY.value
logger.info(
f"版本质量评估完成: {version_id}, "
f"score={quality_report.overall_score}, "
f"grade={quality_report.grade}"
)
except Exception as e:
logger.error(f"质量评估失败: {version_id}, error={e}")
version.status = DatasetStatus.FAILED.value
version.metadata['error'] = str(e)
version.save()
return version
def get_version_statistics(
self,
version_id: str
) -> DatasetStatistics:
"""
获取版本统计信息
"""
version = DatasetVersion.objects.select_related('dataset').get(
id=version_id
)
# 解析质量报告
quality_report = None
if version.quality_report:
try:
quality_report = version.quality_report
except Exception:
pass
# 获取标签分布(从质量报告中或重新计算)
label_distribution = {}
if quality_report:
dims = quality_report.get('dimensions', {})
diversity_dim = dims.get('diversity', {})
details = diversity_dim.get('details', {})
label_distribution = details.get('label_distribution', {})
return DatasetStatistics(
record_count=version.record_count,
file_count=version.file_count,
total_size_bytes=version.total_size_bytes,
quality_score=version.quality_score,
quality_grade=version.quality_grade,
label_distribution=label_distribution,
content_hash=version.content_hash or ""
)
def list_versions(
self,
dataset_id: str,
status: Optional[str] = None,
limit: int = 20,
offset: int = 0
) -> List[DatasetVersion]:
"""
列出数据集版本
"""
queryset = DatasetVersion.objects.filter(dataset_id=dataset_id)
if status:
queryset = queryset.filter(status=status)
return list(
queryset
.order_by('-created_at')
[offset:offset + limit]
)
def _calculate_hash(self, file_obj: BinaryIO) -> str:
"""计算文件内容哈希"""
sha256 = hashlib.sha256()
chunk_size = 8192
while True:
chunk = file_obj.read(chunk_size)
if not chunk:
break
sha256.update(chunk)
return sha256.hexdigest()七、DVC集成脚本
7.1 DVC流水线自动化
# scripts/dvc_pipeline.py
# DVC流水线自动化脚本
import os
import subprocess
import yaml
from dataclasses import dataclass
from typing import List, Optional, Dict, Any
import logging
logger = logging.getLogger(__name__)
@dataclass
class DVCStage:
"""DVC流水线阶段"""
name: str
command: str
dependencies: List[str]
outputs: List[str]
params: Dict[str, Any]
metrics: Optional[Dict[str, Any]] = None
class DVCPipelineManager:
"""
DVC流水线管理器
自动化管理DVC流水线的创建、运行和追踪
"""
def __init__(self, project_root: str):
self.project_root = project_root
self.dvc_file = os.path.join(project_root, 'dvc.yaml')
def init(self) -> None:
"""初始化DVC仓库"""
logger.info("初始化DVC仓库...")
# 初始化DVC
subprocess.run(
['dvc', 'init'],
cwd=self.project_root,
check=True
)
# 配置远程存储
subprocess.run(
['dvc', 'remote', 'add', '-d', 'storage', 's3://ai-platform-datasets/'],
cwd=self.project_root,
check=True
)
logger.info("DVC仓库初始化完成")
def create_pipeline(self, stages: List[DVCStage]) -> None:
"""
创建DVC流水线
Args:
stages: 流水线阶段列表
"""
pipeline_config = {
'stages': {}
}
for stage in stages:
stage_config = {
'cmd': stage.command,
'deps': stage.dependencies,
'outs': stage.outputs,
}
if stage.params:
stage_config['params'] = stage.params
if stage.metrics:
stage_config['metrics'] = stage.metrics
pipeline_config['stages'][stage.name] = stage_config
# 写入dvc.yaml
with open(self.dvc_file, 'w') as f:
yaml.dump(pipeline_config, f, default_flow_style=False)
logger.info(f"DVC流水线已创建: {len(stages)} 个阶段")
def run_stage(self, stage_name: str, force: bool = False) -> None:
"""
运行指定阶段
Args:
stage_name: 阶段名称
force: 是否强制重新运行
"""
cmd = ['dvc', 'repro', stage_name]
if force:
cmd.append('-f')
logger.info(f"运行DVC阶段: {stage_name}")
result = subprocess.run(
cmd,
cwd=self.project_root,
capture_output=True,
text=True
)
if result.returncode != 0:
logger.error(f"DVC阶段运行失败: {stage_name}")
logger.error(result.stderr)
raise RuntimeError(f"DVC阶段运行失败: {result.stderr}")
logger.info(f"DVC阶段运行成功: {stage_name}")
def run_full_pipeline(self, force: bool = False) -> None:
"""
运行完整流水线
"""
cmd = ['dvc', 'repro']
if force:
cmd.append('-f')
logger.info("运行完整DVC流水线...")
result = subprocess.run(
cmd,
cwd=self.project_root,
capture_output=True,
text=True
)
if result.returncode != 0:
logger.error(f"DVC流水线运行失败: {result.stderr}")
raise RuntimeError(f"DVC流水线运行失败: {result.stderr}")
logger.info("DVC流水线运行完成")
def get_pipeline_status(self) -> Dict[str, Any]:
"""
获取流水线状态
"""
result = subprocess.run(
['dvc', 'status', '-o'],
cwd=self.project_root,
capture_output=True,
text=True
)
if result.returncode != 0:
return {'error': result.stderr}
# 解析状态输出
# 简化实现
return {
'status': 'up_to_date' if 'Pipeline is up to date' in result.stdout else 'changed',
'output': result.stdout
}
def commit_changes(self, message: str) -> None:
"""
提交流水线变更
Args:
message: 提交信息
"""
# 添加DVC文件到Git
subprocess.run(
['git', 'add', 'dvc.yaml', 'dvc.lock', '.gitignore'],
cwd=self.project_root,
check=True
)
subprocess.run(
['git', 'commit', '-m', message],
cwd=self.project_root,
check=True
)
# 推送到远程
subprocess.run(
['dvc', 'push'],
cwd=self.project_root,
check=True
)
logger.info(f"DVC变更已提交: {message}")
# 示例:创建训练数据准备流水线
def create_training_data_pipeline(project_root: str) -> None:
"""创建训练数据准备流水线"""
manager = DVCPipelineManager(project_root)
stages = [
DVCStage(
name='preprocess',
command='python src/data/preprocess.py',
dependencies=[
'src/data/preprocess.py',
'data/raw/raw_data.jsonl'
],
outputs=['data/interim/cleaned_data.jsonl'],
params={
'preprocess.min_length': 10,
'preprocess.max_length': 2048,
'preprocess.filter_quality': 0.5
},
metrics=['metrics/cleaning.json']
),
DVCStage(
name='annotate',
command='python src/data/annotate.py',
dependencies=[
'src/data/annotate.py',
'data/interim/cleaned_data.jsonl'
],
outputs=['data/interim/annotated_data.jsonl'],
params={
'annotate.label_schema': 'sentiment',
'annotate.quality_threshold': 0.8
}
),
DVCStage(
name='assemble',
command='python src/data/assemble.py',
dependencies=[
'src/data/assemble.py',
'data/interim/annotated_data.jsonl'
],
outputs=[
'data/processed/train.jsonl',
'data/processed/val.jsonl',
'data/processed/test.jsonl'
],
params={
'assemble.train_ratio': 0.8,
'assemble.val_ratio': 0.1,
'assemble.seed': 42
},
metrics=['metrics/assemble.json']
),
]
manager.create_pipeline(stages)
if __name__ == '__main__':
logging.basicConfig(level=logging.INFO)
# 创建流水线
create_training_data_pipeline('/path/to/project')八、总结与展望
8.1 方案总结
本文详细介绍了一套完整的AI平台数据管理方案,主要包含以下核心模块:
数据集版本化管理:基于DVC的版本控制方案,实现了数据集的可追溯管理。每个版本都有唯一的内容指纹,确保训练结果的可复现性。
数据血缘追踪:通过图数据库存储数据血缘关系,实现了从原始数据到训练数据的完整链路追踪。这对于问题溯源和影响分析至关重要。
数据质量评估:建立了多维度的质量评估框架,从完整性、一致性、准确性、时效性、多样性等角度全面评估数据集质量。
数据集管理API:提供了完整的Django服务实现,涵盖数据集创建、版本管理、文件上传、质量评估等核心功能。
8.2 后续优化方向
自动化数据清洗:目前数据清洗仍依赖人工配置规则,未来可引入机器学习模型自动检测和处理异常数据。
智能版本推荐:基于历史训练效果,推荐最适合当前模型的数据集版本,降低用户选择成本。
实时血缘监控:对数据血缘进行实时监控,当上游数据发生重大变更时,自动触发告警和下游影响评估。
跨平台数据共享:建立安全的数据共享机制,在保护数据隐私的前提下,促进跨团队的数据复用。
数据是AI系统的血液,高质量的数据管理是构建可靠AI平台的基础。希望本文的实践分享能为正在进行类似工作的团队提供一些参考。
评论区