大规模离线数据管道构建:样本获取、清洗、加工与合成

一、引言:数据是AI模型的根基

在大模型和深度学习的时代,算法模型的优劣很大程度上取决于训练数据的质量与规模。业界流传着一句话:"Garbage in, garbage out"------再先进的模型,如果喂给它的是低质量的数据,产出也必然有限。对于AI产品而言,模型训练数据的构建是一条完整的数据流水线,涉及样本获取、数据清洗、离线数据加工和数据合成等环节。

本文将围绕模型训练数据管道的全流程,深入讲解如何从原始日志、公开数据集、用户反馈等多源数据中,构建高质量的训练样本集,为模型训练提供可靠的数据支撑。

二、数据管道整体架构

数据管道采用分层架构,每一层承担特定的职责,层与层之间通过消息队列和分布式存储解耦。

层级 核心任务 技术选型
数据接入层 从多源采集原始数据 Flume / Kafka / 定时拉取
数据清洗层 去重、过滤、格式统一 Spark / Flink
数据加工层 特征工程、标注、格式化 Spark / Hive / Python
数据合成层 数据增强、负采样、混合 Python / Ray
数据管理层 版本管理、质量监控、发布 元数据中心 + 对象存储

三、样本获取:多源数据采集

3.1 数据源分类

模型训练样本的来源通常包括以下几类:

数据源 说明 获取方式
用户搜索日志 用户Query + 点击文档 实时流 + 离线导入
用户行为日志 曝光、点击、收藏、购买 埋点日志流
人工标注数据 相关性标注、质量评分 标注平台导出
公开数据集 开源语料、百科、社区问答 爬虫/API
运营配置数据 标准问答对、知识库 数据库导出

3.2 日志采集与接入

以用户搜索日志为例,其数据采集与接入流程如下。在服务端,搜索请求经过召回和排序后,将完整的请求参数、返回结果及用户后续行为一并写入Kafka,经Flink实时清洗后存入Hive/Doris供后续分析使用。

python 复制代码
# data_ingestion.py - 数据接入模块
import json
import logging
from kafka import KafkaProducer
from datetime import datetime
from typing import Dict, Any

class DataIngestionService:
    """数据接入服务:将各类数据源接入Kafka"""
    
    def __init__(self, bootstrap_servers: str = 'localhost:9092'):
        self.producer = KafkaProducer(
            bootstrap_servers=bootstrap_servers,
            value_serializer=lambda v: json.dumps(v).encode('utf-8'),
            compression_type='snappy',
            batch_size=16384,
            linger_ms=100
        )
        self.topics = {
            'search_log': 'raw.search_log',
            'click_log': 'raw.click_log',
            'user_profile': 'raw.user_profile',
            'external_corpus': 'raw.external_corpus'
        }
    
    def ingest_search_log(self, log_data: Dict[str, Any]):
        """接入搜索日志"""
        record = {
            'timestamp': datetime.now().isoformat(),
            'type': 'search_log',
            'data': log_data
        }
        self.producer.send(self.topics['search_log'], record)
    
    def ingest_click_log(self, click_data: Dict[str, Any]):
        """接入点击日志"""
        record = {
            'timestamp': datetime.now().isoformat(),
            'type': 'click_log',
            'data': click_data
        }
        self.producer.send(self.topics['click_log'], record)
    
    def ingest_external_data(self, data: Dict[str, Any], source: str):
        """接入外部数据(公开数据集、爬虫数据等)"""
        record = {
            'timestamp': datetime.now().isoformat(),
            'type': 'external',
            'source': source,
            'data': data
        }
        self.producer.send(self.topics['external_corpus'], record)
    
    def flush(self):
        """确保所有消息发送完成"""
        self.producer.flush()

# 使用示例
ingestion = DataIngestionService()

# 模拟搜索日志接入
search_log = {
    'query': '大模型应用开发',
    'user_id': 'user_12345',
    'session_id': 'sess_67890',
    'results': ['doc_001', 'doc_002', 'doc_003'],
    'clicked_doc': 'doc_002',
    'response_time_ms': 156
}
ingestion.ingest_search_log(search_log)
ingestion.flush()

3.3 标注数据管理

人工标注数据是监督学习的黄金标准。标注数据的质量直接影响模型的上限。

python 复制代码
# annotation_manager.py - 标注数据管理
from typing import List, Dict, Optional
import uuid
import pandas as pd
from datetime import datetime

class AnnotationSample:
    """标注样本"""
    def __init__(self, sample_id: str, query: str, doc_id: str, 
                 doc_content: str, label: Optional[int] = None):
        self.id = sample_id
        self.query = query
        self.doc_id = doc_id
        self.doc_content = doc_content
        self.label = label  # 0: 不相关, 1: 部分相关, 2: 高度相关
        self.annotator = None
        self.annotate_time = None
    
    def to_dict(self) -> dict:
        return {
            'id': self.id,
            'query': self.query,
            'doc_id': self.doc_id,
            'doc_content': self.doc_content[:500],  # 只保留前500字符
            'label': self.label,
            'annotator': self.annotator,
            'annotate_time': self.annotate_time
        }

class AnnotationManager:
    """标注任务管理"""
    
    def __init__(self, storage_path: str):
        self.storage_path = storage_path
        self.samples: List[AnnotationSample] = []
        self.sample_queue: List[str] = []  # 待标注ID队列
    
    def create_task(self, queries: List[str], doc_pairs: List[tuple]) -> List[str]:
        """创建标注任务"""
        sample_ids = []
        for query, (doc_id, doc_content) in zip(queries, doc_pairs):
            sample_id = f"sample_{uuid.uuid4().hex[:8]}"
            sample = AnnotationSample(
                sample_id=sample_id,
                query=query,
                doc_id=doc_id,
                doc_content=doc_content
            )
            self.samples.append(sample)
            self.sample_queue.append(sample_id)
            sample_ids.append(sample_id)
        return sample_ids
    
    def submit_annotation(self, sample_id: str, label: int, annotator: str):
        """提交标注结果"""
        for sample in self.samples:
            if sample.id == sample_id:
                sample.label = label
                sample.annotator = annotator
                sample.annotate_time = datetime.now().isoformat()
                self.sample_queue.remove(sample_id)
                break
    
    def export_to_csv(self, filename: str):
        """导出标注数据为CSV"""
        data = [s.to_dict() for s in self.samples if s.label is not None]
        df = pd.DataFrame(data)
        df.to_csv(f"{self.storage_path}/{filename}", index=False)
        return len(data)
    
    def get_statistics(self) -> dict:
        """获取标注统计"""
        labeled = [s for s in self.samples if s.label is not None]
        if not labeled:
            return {'total': len(self.samples), 'labeled': 0}
        
        label_counts = {}
        for sample in labeled:
            label_counts[sample.label] = label_counts.get(sample.label, 0) + 1
        
        return {
            'total': len(self.samples),
            'labeled': len(labeled),
            'pending': len(self.sample_queue),
            'label_distribution': label_counts
        }

四、数据清洗:从原始数据到可用样本

4.1 清洗规则体系

数据清洗的目标是将"脏数据"转化为符合质量标准的数据。常见的"脏数据"包括空值异常、格式错误、编码问题和语义重复等,需要分层制定清洗规则。

python 复制代码
# data_cleaner.py - 数据清洗模块
import re
import hashlib
from typing import List, Dict, Any, Optional
from dataclasses import dataclass, field

@dataclass
class CleanedSample:
    """清洗后的样本"""
    id: str
    query: str
    query_tokens: List[str]
    doc_id: str
    doc_content: str
    doc_tokens: List[str]
    label: Optional[int]
    quality_score: float
    metadata: Dict[str, Any] = field(default_factory=dict)

class DataCleaner:
    """数据清洗器"""
    
    def __init__(self):
        self.stopwords = set(['的', '了', '是', '在', '和', '与', '或', '等'])
        self.min_content_length = 20
        self.max_content_length = 100000
        self.quality_threshold = 0.3
    
    def clean_search_log(self, raw_log: Dict[str, Any]) -> Optional[Dict[str, Any]]:
        """清洗搜索日志"""
        # 1. 校验必要字段
        required_fields = ['query', 'user_id', 'results']
        for field in required_fields:
            if field not in raw_log or not raw_log[field]:
                return None
        
        # 2. Query清洗
        query = self._clean_text(raw_log['query'])
        if len(query) < 2:  # Query太短,无效
            return None
        
        # 3. 去重:基于query+user_id去重
        dup_key = f"{query}_{raw_log['user_id']}"
        
        # 4. 格式统一
        cleaned = {
            'query': query,
            'user_id': raw_log['user_id'][:32],  # 脱敏
            'session_id': raw_log.get('session_id', ''),
            'results': raw_log.get('results', [])[:20],  # 截断
            'clicked_doc': raw_log.get('clicked_doc', ''),
            'timestamp': raw_log.get('timestamp', ''),
            'quality_score': self._calculate_quality(raw_log)
        }
        
        if cleaned['quality_score'] < self.quality_threshold:
            return None
        
        return cleaned
    
    def clean_text_corpus(self, raw_text: str, doc_id: str) -> Optional[Dict[str, Any]]:
        """清洗文本语料"""
        # 1. 去除HTML标签
        text = re.sub(r'<[^>]+>', '', raw_text)
        
        # 2. 去除多余空白
        text = re.sub(r'\s+', ' ', text).strip()
        
        # 3. 去除特殊字符(保留中文、英文、数字、基本标点)
        text = re.sub(r'[^\u4e00-\u9fa5a-zA-Z0-9,。!?、;:""''()\s]', '', text)
        
        # 4. 长度过滤
        if len(text) < self.min_content_length:
            return None
        if len(text) > self.max_content_length:
            text = text[:self.max_content_length]
        
        # 5. 去重(基于内容哈希)
        content_hash = hashlib.md5(text.encode()).hexdigest()
        
        return {
            'doc_id': doc_id,
            'content': text,
            'content_hash': content_hash,
            'length': len(text)
        }
    
    def _clean_text(self, text: str) -> str:
        """通用文本清洗"""
        # 去除首尾空白
        text = text.strip()
        # 全角转半角
        text = text.replace(',', ',').replace('。', '.').replace(';', ';')
        # 去除多余空格
        text = re.sub(r'\s+', ' ', text)
        return text
    
    def _calculate_quality(self, log: Dict[str, Any]) -> float:
        """计算数据质量得分"""
        score = 0.0
        # Query长度得分
        query_len = len(log.get('query', ''))
        if 3 <= query_len <= 50:
            score += 0.3
        elif query_len > 50:
            score += 0.1
        
        # 是否有点击(表明用户满意)
        if log.get('clicked_doc'):
            score += 0.3
        
        # 结果数量
        results_count = len(log.get('results', []))
        if results_count > 0:
            score += min(results_count / 20, 0.2)
        
        # 响应时间(越快越好)
        response_time = log.get('response_time_ms', 1000)
        if response_time < 200:
            score += 0.2
        elif response_time < 500:
            score += 0.1
        
        return min(score, 1.0)

4.2 大规模清洗:Spark作业实现

对于TB级别的数据,需要使用Spark进行分布式清洗。以下是一个基于PySpark的清洗作业示例:

python 复制代码
# spark_cleaner.py - Spark分布式清洗
from pyspark.sql import SparkSession
from pyspark.sql.functions import udf, col, when, length, regexp_replace
from pyspark.sql.types import StringType, DoubleType
import re

class SparkDataCleaner:
    """基于Spark的大规模数据清洗"""
    
    def __init__(self, app_name: str = "DataCleaner"):
        self.spark = SparkSession.builder \
            .appName(app_name) \
            .config("spark.sql.adaptive.enabled", "true") \
            .config("spark.sql.adaptive.coalescePartitions.enabled", "true") \
            .getOrCreate()
    
    def clean_search_logs(self, input_path: str, output_path: str):
        """清洗大规模搜索日志"""
        df = self.spark.read.parquet(input_path)
        
        # 定义清洗UDF
        def clean_query(query: str) -> str:
            if not query:
                return ""
            query = re.sub(r'[^\w\u4e00-\u9fa5]', ' ', query)
            query = re.sub(r'\s+', ' ', query).strip()
            return query
        
        def is_valid_query(query: str) -> bool:
            return query and len(query) >= 2
        
        clean_query_udf = udf(clean_query, StringType())
        is_valid_udf = udf(is_valid_query, StringType())
        
        # 执行清洗
        cleaned_df = df \
            .filter(col("query").isNotNull()) \
            .filter(col("user_id").isNotNull()) \
            .withColumn("query_cleaned", clean_query_udf(col("query"))) \
            .filter(is_valid_udf(col("query_cleaned"))) \
            .withColumn("user_id_hashed", 
                       when(length(col("user_id")) > 32, 
                            regexp_replace(col("user_id"), r'(.{8}).*', r'$1****'))) \
            .dropDuplicates(["query_cleaned", "user_id"]) \
            .select("query_cleaned", "user_id", "session_id", 
                    "clicked_doc", "timestamp")
        
        cleaned_df.write \
            .mode("overwrite") \
            .parquet(output_path)
        
        return cleaned_df.count()
    
    def deduplicate_corpus(self, input_path: str, output_path: str, 
                           key_field: str = "content_hash"):
        """基于MinHash去重大规模文本语料"""
        df = self.spark.read.parquet(input_path)
        
        # 按哈希去重(保留每个哈希的第一条)
        deduped_df = df.dropDuplicates([key_field])
        
        # 统计去重效果
        total_count = df.count()
        unique_count = deduped_df.count()
        print(f"去重前: {total_count}, 去重后: {unique_count}, "
              f"去重率: {(1 - unique_count/total_count)*100:.2f}%")
        
        deduped_df.write.mode("overwrite").parquet(output_path)
        return unique_count
    
    def filter_by_quality(self, input_path: str, output_path: str, 
                          threshold: float = 0.3):
        """按质量阈值过滤"""
        df = self.spark.read.parquet(input_path)
        
        # 计算质量得分
        df_with_quality = df.withColumn(
            "quality_score",
            when(length(col("content")) > 100, 0.5) +
            when(length(col("content")) > 500, 0.3) +
            when(col("has_image") == True, 0.2)
        )
        
        filtered_df = df_with_quality.filter(
            col("quality_score") >= threshold
        )
        
        filtered_df.write.mode("overwrite").parquet(output_path)
        return filtered_df.count()
    
    def stop(self):
        self.spark.stop()

五、离线数据加工:从样本到训练集

5.1 特征工程与标签生成

经过清洗的原始数据,还需要加工成模型可直接使用的格式,包括特征提取、标签生成和格式转换。

python 复制代码
# data_processor.py - 数据加工模块
import numpy as np
import pandas as pd
from typing import List, Dict, Any, Tuple
from sklearn.model_selection import train_test_split

class DataProcessor:
    """数据加工:特征工程、标签生成、数据集划分"""
    
    def __init__(self, config: dict):
        self.config = config
        self.label_mapping = {
            'irrelevant': 0,
            'partially_relevant': 1,
            'highly_relevant': 2,
            'click': 1,
            'no_click': 0
        }
    
    def generate_ranking_samples(self, search_logs: List[Dict]) -> pd.DataFrame:
        """
        从搜索日志生成排序训练样本
        正样本:被点击的文档
        负样本:被展示但未被点击的文档
        """
        samples = []
        
        for log in search_logs:
            query = log.get('query', '')
            results = log.get('results', [])
            clicked = log.get('clicked_doc', '')
            
            if not query or not results:
                continue
            
            for doc_id in results:
                sample = {
                    'query': query,
                    'doc_id': doc_id,
                    'label': 1 if doc_id == clicked else 0,
                    'is_clicked': doc_id == clicked
                }
                samples.append(sample)
        
        # 负采样:保持正负样本平衡
        df = pd.DataFrame(samples)
        positive = df[df['label'] == 1]
        negative = df[df['label'] == 0]
        
        # 下采样负样本(控制比例1:3)
        if len(negative) > len(positive) * 3:
            negative = negative.sample(n=len(positive) * 3, random_state=42)
        
        balanced_df = pd.concat([positive, negative]).sample(frac=1, random_state=42)
        
        return balanced_df
    
    def extract_query_doc_features(self, df: pd.DataFrame, 
                                   doc_content_dict: Dict[str, str]) -> pd.DataFrame:
        """提取Query-Doc匹配特征"""
        features = []
        
        for _, row in df.iterrows():
            query = row['query']
            doc_id = row['doc_id']
            content = doc_content_dict.get(doc_id, '')
            
            feature = {
                'query': query,
                'doc_id': doc_id,
                'label': row['label']
            }
            
            # 计算基础特征
            q_words = set(self._tokenize(query))
            d_words = set(self._tokenize(content))
            
            overlap = q_words & d_words
            feature['term_overlap'] = len(overlap)
            feature['term_overlap_ratio'] = len(overlap) / max(len(q_words), 1)
            feature['query_length'] = len(query)
            feature['doc_length'] = len(content)
            feature['doc_contains_all_terms'] = 1.0 if all(w in content for w in q_words) else 0.0
            
            features.append(feature)
        
        return pd.DataFrame(features)
    
    def _tokenize(self, text: str) -> List[str]:
        """简单分词"""
        import jieba
        return [w for w in jieba.lcut(text) if len(w) > 1]
    
    def split_dataset(self, df: pd.DataFrame, test_size: float = 0.2, 
                      valid_size: float = 0.1, random_state: int = 42):
        """划分训练/验证/测试集"""
        # 先分测试集
        train_valid, test = train_test_split(
            df, test_size=test_size, random_state=random_state, stratify=df['label']
        )
        # 再分验证集
        valid_ratio = valid_size / (1 - test_size)
        train, valid = train_test_split(
            train_valid, test_size=valid_ratio, random_state=random_state, 
            stratify=train_valid['label']
        )
        
        return {
            'train': train,
            'valid': valid,
            'test': test
        }
    
    def convert_to_training_format(self, df: pd.DataFrame, 
                                   feature_columns: List[str],
                                   label_column: str = 'label',
                                   output_format: str = 'tfrecord') -> str:
        """
        转换为模型训练格式
        支持:tfrecord、parquet、csv、jsonl
        """
        import json
        
        # 提取特征矩阵和标签
        X = df[feature_columns].values
        y = df[label_column].values
        
        # 保存为JSONL格式(示例)
        output_path = f"data/training_samples.{output_format}"
        
        if output_format == 'jsonl':
            with open(output_path, 'w') as f:
                for i, row in df.iterrows():
                    sample = {
                        'features': {col: row[col] for col in feature_columns},
                        'label': int(row[label_column])
                    }
                    f.write(json.dumps(sample, ensure_ascii=False) + '\n')
        
        print(f"数据集已保存到 {output_path},样本数: {len(df)}")
        return output_path

5.2 离线数据管道编排

使用Apache Airflow编排完整的数据处理流水线:

python 复制代码
# data_pipeline_dag.py - Airflow DAG定义
from airflow import DAG
from airflow.operators.python_operator import PythonOperator
from airflow.operators.bash_operator import BashOperator
from datetime import datetime, timedelta

default_args = {
    'owner': 'data_team',
    'depends_on_past': False,
    'start_date': datetime(2024, 1, 1),
    'email_on_failure': True,
    'retries': 3,
    'retry_delay': timedelta(minutes=5)
}

dag = DAG(
    'model_training_data_pipeline',
    default_args=default_args,
    description='模型训练数据流水线',
    schedule_interval='0 2 * * *',  # 每天凌晨2点执行
    catchup=False,
    max_active_runs=1
)

# 任务1: 数据采集
task_collect_logs = BashOperator(
    task_id='collect_logs',
    bash_command='python /data/pipeline/collect_logs.py --date {{ ds }}',
    dag=dag
)

# 任务2: 数据清洗(Spark)
task_clean_data = BashOperator(
    task_id='clean_data',
    bash_command='spark-submit /data/pipeline/spark_cleaner.py --input /raw/logs/{{ ds }} --output /clean/{{ ds }}',
    dag=dag
)

# 任务3: 去重
task_deduplicate = BashOperator(
    task_id='deduplicate',
    bash_command='spark-submit /data/pipeline/deduplicate.py --input /clean/{{ ds }} --output /dedup/{{ ds }}',
    dag=dag
)

# 任务4: 特征提取
task_extract_features = PythonOperator(
    task_id='extract_features',
    python_callable=extract_features_func,
    provide_context=True,
    dag=dag
)

# 任务5: 数据集划分与导出
task_split_dataset = PythonOperator(
    task_id='split_dataset',
    python_callable=split_dataset_func,
    provide_context=True,
    dag=dag
)

# 任务6: 数据验证与质量报告
task_validate = PythonOperator(
    task_id='validate_data',
    python_callable=validate_data_func,
    provide_context=True,
    dag=dag
)

# 定义任务依赖
task_collect_logs >> task_clean_data >> task_deduplicate >> \
    task_extract_features >> task_split_dataset >> task_validate

六、数据合成:扩充与增强

6.1 数据增强策略

当标注数据不足时,数据合成技术可以有效扩充训练集。

python 复制代码
# data_augmentation.py - 数据增强模块
import random
import jieba
from typing import List, Dict, Any
import numpy as np

class DataAugmenter:
    """数据增强:扩充训练样本"""
    
    def __init__(self):
        self.synonym_dict = self._load_synonyms()
        self.back_translation_models = self._load_translation_models()
    
    def synonym_replacement(self, text: str, ratio: float = 0.2) -> str:
        """同义词替换"""
        words = jieba.lcut(text)
        if not words:
            return text
        
        n_replace = max(1, int(len(words) * ratio))
        replace_indices = random.sample(range(len(words)), min(n_replace, len(words)))
        
        for idx in replace_indices:
            word = words[idx]
            if word in self.synonym_dict:
                synonyms = self.synonym_dict[word]
                if synonyms:
                    words[idx] = random.choice(synonyms)
        
        return ''.join(words)
    
    def random_insertion(self, text: str, ratio: float = 0.1) -> str:
        """随机插入"""
        words = jieba.lcut(text)
        if not words:
            return text
        
        n_insert = max(1, int(len(words) * ratio))
        for _ in range(n_insert):
            # 从文本中随机选一个词,插入它的同义词
            word = random.choice(words)
            if word in self.synonym_dict and self.synonym_dict[word]:
                pos = random.randint(0, len(words))
                words.insert(pos, random.choice(self.synonym_dict[word]))
        
        return ''.join(words)
    
    def random_deletion(self, text: str, ratio: float = 0.1) -> str:
        """随机删除"""
        words = jieba.lcut(text)
        if not words:
            return text
        
        if len(words) <= 3:
            return text
        
        n_delete = max(1, int(len(words) * ratio))
        delete_indices = random.sample(range(len(words)), min(n_delete, len(words) - 1))
        
        kept_words = [w for i, w in enumerate(words) if i not in delete_indices]
        return ''.join(kept_words)
    
    def back_translation(self, text: str, target_lang: str = 'en') -> str:
        """回译:中文 -> 英文 -> 中文"""
        # 使用翻译API
        # 这里用占位实现
        return text  # 实际调用翻译服务
    
    def mixup(self, sample1: Dict, sample2: Dict, alpha: float = 0.4) -> Dict:
        """Mixup数据增强(适用于特征向量)"""
        lambda_val = np.random.beta(alpha, alpha)
        
        # 混合特征
        mixed_features = {}
        for key in sample1['features']:
            val1 = sample1['features'][key]
            val2 = sample2['features'][key]
            mixed_features[key] = lambda_val * val1 + (1 - lambda_val) * val2
        
        # 混合标签(软标签)
        mixed_label = lambda_val * sample1['label'] + (1 - lambda_val) * sample2['label']
        
        return {
            'features': mixed_features,
            'label': mixed_label
        }
    
    def _load_synonyms(self) -> Dict[str, List[str]]:
        """加载同义词词典"""
        # 实际从文件加载
        return {
            '好': ['优秀', '良好', '出色'],
            '大': ['巨大', '庞大', '宏大'],
            '快': ['迅速', '快速', '飞快']
        }

6.2 负样本合成

在搜索排序中,高质量的负样本对模型训练至关重要。

python 复制代码
# negative_sampling.py - 负样本合成
import random
import numpy as np
from typing import List, Dict, Tuple

class NegativeSampler:
    """负样本合成策略"""
    
    def __init__(self, corpus: List[str], 
                 negative_sampling_ratio: int = 3):
        self.corpus = corpus  # 全量文档集合
        self.negative_ratio = negative_sampling_ratio
    
    def random_negative(self, query: str, positive_docs: List[str], 
                        exclude_ids: set) -> List[Tuple[str, int]]:
        """随机负采样:从全量语料中随机抽取"""
        negative_samples = []
        available = [doc for doc in self.corpus if doc not in exclude_ids]
        
        if not available:
            return negative_samples
        
        n_needed = len(positive_docs) * self.negative_ratio
        sampled = random.sample(available, min(n_needed, len(available)))
        
        for doc in sampled:
            negative_samples.append((doc, 0))
        
        return negative_samples
    
    def hard_negative(self, query: str, positive_docs: List[str],
                      ranker, top_k: int = 50) -> List[Tuple[str, int]]:
        """困难负采样:检索结果中排名靠前但未被点击的文档"""
        # 使用当前模型对全量文档进行排序
        all_scores = []
        for doc_id in self.corpus:
            score = ranker.score(query, doc_id)
            all_scores.append((doc_id, score))
        
        # 按得分排序
        sorted_docs = sorted(all_scores, key=lambda x: x[1], reverse=True)
        
        # 排除正样本
        positive_set = set(positive_docs)
        candidates = [(doc, score) for doc, score in sorted_docs 
                     if doc not in positive_set]
        
        # 取Top-K作为困难负样本
        hard_negatives = []
        for doc, score in candidates[:top_k]:
            hard_negatives.append((doc, 0))
        
        return hard_negatives
    
    def batch_negative_sampling(self, queries: List[str], 
                               positive_dict: Dict[str, List[str]],
                               strategy: str = 'mix') -> List[Dict]:
        """批量负采样"""
        samples = []
        
        for query in queries:
            positives = positive_dict.get(query, [])
            if not positives:
                continue
            
            if strategy == 'random':
                negatives = self.random_negative(query, positives, set(positives))
            elif strategy == 'hard':
                negatives = self.hard_negative(query, positives, None)
            else:  # mix: 混合策略
                random_neg = self.random_negative(query, positives, set(positives))
                hard_neg = self.hard_negative(query, positives, None)
                # 取随机负样本的70% + 困难负样本的30%
                n_random = int(self.negative_ratio * len(positives) * 0.7)
                n_hard = self.negative_ratio * len(positives) - n_random
                negatives = random_neg[:n_random] + hard_neg[:n_hard]
            
            # 构建样本
            for doc_id, label in negatives:
                samples.append({
                    'query': query,
                    'doc_id': doc_id,
                    'label': label,
                    'type': 'negative'
                })
            
            for doc_id in positives:
                samples.append({
                    'query': query,
                    'doc_id': doc_id,
                    'label': 1,
                    'type': 'positive'
                })
        
        return samples

七、数据质量监控

python 复制代码
# data_quality_monitor.py - 数据质量监控
import pandas as pd
import numpy as np
from typing import Dict, Any
from datetime import datetime

class DataQualityMonitor:
    """数据质量监控:统计指标、异常检测、报告生成"""
    
    def __init__(self, slack_webhook: str = None):
        self.slack_webhook = slack_webhook
        self.metrics_history = []
    
    def compute_metrics(self, df: pd.DataFrame) -> Dict[str, Any]:
        """计算数据质量指标"""
        metrics = {
            'timestamp': datetime.now().isoformat(),
            'total_samples': len(df),
            'positive_ratio': df['label'].sum() / len(df) if 'label' in df else None,
            'missing_rate': df.isnull().sum().sum() / (df.shape[0] * df.shape[1]),
            'duplicate_rate': df.duplicated().sum() / len(df),
            'avg_query_length': df['query'].str.len().mean() if 'query' in df else None,
            'avg_doc_length': df['doc_id'].nunique() / len(df) if 'doc_id' in df else None,
        }
        return metrics
    
    def detect_anomalies(self, current_metrics: Dict, 
                         baseline: Dict) -> List[str]:
        """检测异常指标"""
        anomalies = []
        thresholds = {
            'total_samples': (0.5, 2.0),  # 相对于基线
            'positive_ratio': (0.05, 0.95),
            'missing_rate': (0, 0.1),
            'duplicate_rate': (0, 0.3)
        }
        
        for key, (lower, upper) in thresholds.items():
            if key in current_metrics:
                current_val = current_metrics[key]
                if key in baseline:
                    ratio = current_val / (baseline[key] + 1e-6)
                    if ratio < lower or ratio > upper:
                        anomalies.append(f"{key}: {current_val:.3f} (baseline: {baseline[key]:.3f})")
                else:
                    if current_val < lower or current_val > upper:
                        anomalies.append(f"{key}: {current_val:.3f}")
        
        return anomalies
    
    def generate_report(self, metrics: Dict) -> str:
        """生成质量报告"""
        report = f"""
        ===== 数据质量报告 =====
        时间: {metrics['timestamp']}
        样本总数: {metrics['total_samples']:,}
        正样本比例: {metrics['positive_ratio']:.2%}
        缺失率: {metrics['missing_rate']:.2%}
        重复率: {metrics['duplicate_rate']:.2%}
        平均Query长度: {metrics['avg_query_length']:.1f}
        =========================
        """
        return report

八、总结

本文系统阐述了大规模离线数据管道的完整构建流程:

阶段 核心任务 关键技术
样本获取 多源日志采集、标注数据管理 Kafka、Flume、数据埋点
数据清洗 去重、过滤、格式统一、质量评分 Spark、正则表达式、布隆过滤器
数据加工 特征提取、标签生成、数据集划分 Pandas、Scikit-learn、Hive
数据合成 同义词替换、回译、困难负采样 NLTK、翻译API、Mixup
质量监控 指标统计、异常检测、报告生成 Pandas、可视化报表

构建高质量的训练数据管道是AI产品从实验室走向生产环境的关键基础设施。数据管道的核心原则可以归纳为三点:

  1. 自动化:将数据采集、清洗、加工
相关推荐
意图共鸣1 小时前
意图共鸣科技:人文驱动AI——从“能不能“到“好不好“
人工智能·microsoft
tang777891 小时前
分布式爬虫优化指南:如何用代理IP把采集效率提升300%
分布式·爬虫·python·tcp/ip·分布式爬虫·爬虫代理·代理ip
海兰1 小时前
【思考】银行落地业务Agent的思考
人工智能·机器学习·agent
梦想三三1 小时前
YOLOv5口罩检测完整实战(二):从现成数据集到训练、评估与摄像头识别
人工智能·yolo·计算机视觉·目标跟踪
gwf2161 小时前
Soft-RoCE与Soft-iWARP深度解析:无硬件RDMA学习环境搭建(零基础必知必会)
人工智能·python·tcp/ip·tcp·tcpdump
云端漫步19872 小时前
HarmonyOS NEXT AI 智能生活助手:会话管理与聊天记录保存
人工智能·生活·harmonyos
AI行业学习2 小时前
Claude Code + cc-switch + Git + Node.js 一站式完整安装配置教程【8.3】
git·python·安全·前端框架·node.js·html·notepad++
AI行业学习2 小时前
Claude Code + cc-switch + Git + Node.js 一站式完整安装配置教程(2026最新·国内可用版)
人工智能·git·python·安全·node.js·html·notepad++