AI风险管理折腾手记

别急着给AI风险管理下定义,先看这次卡在哪。

我踩过最狠的坑是数据泄露——用户行为日志里混入了标注数据,导致模型在测试集上表现好得离谱,上线后直接翻车。

风险识别:先把坑挖出来再往下走

模型风险分类

数据风险

训练数据有问题,模型输出必然有问题。我踩过最狠的坑是数据泄露——用户行为日志里混入了标注数据,导致模型在测试集上表现好得离谱,上线后直接翻车。

# 数据泄露检测
def detect_data_leakage(train_data, test_data, similarity_threshold=0.95):
    """
    检测训练集和测试集之间的数据泄露

    Args:
        train_data: 训练数据 DataFrame
        test_data: 测试数据 DataFrame
        similarity_threshold: 相似度阈值

    Returns:
        泄露样本数量和相似度分布
    """
    # 计算相似度
    from sklearn.feature_extraction.text import TfidfVectorizer
    from sklearn.metrics.pairwise import cosine_similarity

    vectorizer = TfidfVectorizer()
    all_text = train_data['text'].tolist() + test_data['text'].tolist()
    tfidf_matrix = vectorizer.fit_transform(all_text)

    similarity_matrix = cosine_similarity(
        tfidf_matrix[:len(train_data)],
        tfidf_matrix[len(train_data):]
    )

    leaked_samples = (similarity_matrix > similarity_threshold).sum()
    print(f"检测到 {leaked_samples} 个潜在泄露样本")
    print(f"相似度分布: {np.histogram(similarity_matrix.flatten())}")

    return leaked_samples, similarity_matrix

模型风险

模型本身的风险主要在这几个方面:

# 模型风险评估框架
class ModelRiskAssessment:
    def __init__(self, model, validation_data):
        self.model = model
        self.validation_data = validation_data

    def assess_distribution_shift(self, reference_data, current_data):
        """
        评估分布漂移

        使用 KS 检验和 Wasserstein 距离检测特征分布变化
        """
        from scipy import stats
        from scipy.spatial.distance import wasserstein_distance

        drift_metrics = {}

        for feature in reference_data.columns:
            # KS 检验
            ks_stat, p_value = stats.ks_2samp(
                reference_data[feature].dropna(),
                current_data[feature].dropna()
            )

            # Wasserstein 距离
            wd = wasserstein_distance(
                reference_data[feature].dropna(),
                current_data[feature].dropna()
            )

            drift_metrics[feature] = {
                'ks_statistic': ks_stat,
                'p_value': p_value,
                'wasserstein_distance': wd,
                'is_drifted': p_value < 0.05 and wd > 0.1
            }

        return drift_metrics

    def assess_uncertainty(self, X):
        """
        评估模型不确定性

        使用蒙特卡洛 Dropout 或集成方法
        """
        predictions = []

        # 蒙特卡洛 Dropout
        if hasattr(self.model, 'predict_proba'):
            for _ in range(30):  # 采样 30 次
                pred = self.model.predict_proba(X)
                predictions.append(pred)

            predictions = np.array(predictions)
            mean_pred = predictions.mean(axis=0)
            uncertainty = predictions.std(axis=0)

            return {
                'mean_prediction': mean_pred,
                'uncertainty': uncertainty,
                'high_uncertainty_samples': np.where(uncertainty > 0.3)[0]
            }

        return {'error': '模型不支持不确定性估计'}

部署风险

模型上线后的环境、性能、资源问题。

# 模型部署前的性能基准测试
python benchmark_model.py \
  --model-path ./models/recommendation_v2.pt \
  --batch-size 32 \
  --concurrent-requests 10 \
  --duration 300 \
  --output ./reports/performance_baseline.json
# 性能基准测试脚本
import time
import psutil
import numpy as np
from concurrent.futures import ThreadPoolExecutor
import json

def benchmark_model(model, input_shape, batch_size, concurrent_requests, duration):
    results = {
        'latencies': [],
        'throughput': [],
        'cpu_usage': [],
        'memory_usage': [],
        'errors': []
    }

    start_time = time.time()
    request_count = 0

    def make_request():
        try:
            # 模拟推理输入
            batch_input = np.random.randn(*input_shape[:1], batch_size, *input_shape[2:])

            req_start = time.time()
            output = model.predict(batch_input)
            req_end = time.time()

            latency = (req_end - req_start) * 1000  # 转换为毫秒
            results['latencies'].append(latency)
            results['throughput'].append(batch_size / latency * 1000)

            return True
        except Exception as e:
            results['errors'].append(str(e))
            return False

    while time.time() - start_time < duration:
        with ThreadPoolExecutor(max_workers=concurrent_requests) as executor:
            futures = [executor.submit(make_request) for _ in range(concurrent_requests)]
            for future in futures:
                future.result()
                request_count += 1

                # 记录资源使用
                results['cpu_usage'].append(psutil.cpu_percent())
                results['memory_usage'].append(psutil.virtual_memory().percent)

    # 统计结果
    summary = {
        'total_requests': request_count,
        'avg_latency_ms': np.mean(results['latencies']),
        'p95_latency_ms': np.percentile(results['latencies'], 95),
        'p99_latency_ms': np.percentile(results['latencies'], 99),
        'throughput_rps': np.mean(results['throughput']),
        'error_rate': len(results['errors']) / request_count,
        'avg_cpu_usage': np.mean(results['cpu_usage']),
        'max_memory_usage': np.max(results['memory_usage'])
    }

    return summary

风险识别清单

# AI 风险识别检查清单
risk_identification_checklist:
  data_risks:
    - name: "数据质量检查"
      checks:
        - 缺失值比例
        - 异常值检测
        - 数据分布一致性
        - 特征相关性分析
    - name: "数据安全性"
      checks:
        - PII 敏感信息识别
        - 数据加密存储
        - 访问权限控制
        - 审计日志记录

  model_risks:
    - name: "模型性能"
      checks:
        - 准确率/召回率/F1
        - ROC-AUC 曲线
        - 混淆矩阵分析
        - 跨类别表现差异
    - name: "模型公平性"
      checks:
        - 人口统计学均等
        - 机会均等
        - 校准误差分析
        - 反事实公平性
    - name: "模型可解释性"
      checks:
        - SHAP 值分析
        - 特征重要性排序
        - LIME 局部解释
        - 决策路径可视化

  deployment_risks:
    - name: "性能指标"
      checks:
        - 推理延迟(P50/P95/P99)
        - 吞吐量(QPS)
        - 资源占用(CPU/内存/GPU)
        - 扩缩容能力
    - name: "监控告警"
      checks:
        - 模型性能漂移监控
        - 输入数据分布监控
        - 错误率监控
        - 异常请求监控

风险应对:有预案总比没有强

防护机制

输入过滤

# 输入内容过滤系统
class InputContentFilter:
    def __init__(self, config_path):
        import yaml
        with open(config_path, 'r', encoding='utf-8') as f:
            self.config = yaml.safe_load(f)

        # 加载敏感词库
        self.sensitive_words = self._load_sensitive_words(
            self.config['sensitive_words_path']
        )

        # 加载正则规则
        self.regex_patterns = [
            re.compile(pattern) for pattern in self.config['regex_patterns']
        ]

    def _load_sensitive_words(self, path):
        """加载敏感词库"""
        with open(path, 'r', encoding='utf-8') as f:
            words = [line.strip() for line in f if line.strip()]
        return set(words)

    def filter(self, input_text):
        """
        过滤输入内容

        Returns:
            {
                'is_safe': bool,
                'reasons': list,
                'filtered_text': str
            }
        """
        reasons = []
        filtered_text = input_text

        # 敏感词过滤
        detected_words = []
        for word in self.sensitive_words:
            if word in input_text:
                detected_words.append(word)
                filtered_text = filtered_text.replace(word, '*' * len(word))

        if detected_words:
            reasons.append(f"检测到敏感词: {', '.join(detected_words)}")

        # 正则匹配
        for pattern in self.regex_patterns:
            matches = pattern.findall(input_text)
            if matches:
                reasons.append(f"检测到违规模式: {pattern.pattern}")

        # 长度检查
        if len(input_text) < self.config['min_length']:
            reasons.append(f"输入过短,最少 {self.config['min_length']} 字符")

        if len(input_text) > self.config['max_length']:
            reasons.append(f"输入过长,最多 {self.config['max_length']} 字符")

        return {
            'is_safe': len(reasons) == 0,
            'reasons': reasons,
            'filtered_text': filtered_text
        }

# 使用示例
filter = InputContentFilter('./config/content_filter.yaml')
result = filter.filter("这是一段测试文本")

if not result['is_safe']:
    print(f"内容不安全: {result['reasons']}")
    print(f"过滤后: {result['filtered_text']}")
else:
    print("内容安全,可以处理")

输出限制

# 输出内容限制和脱敏
class OutputContentGuard:
    def __init__(self):
        # 定义限制规则
        self.restrictions = {
            'max_length': 2000,
            'forbidden_patterns': [
                r'\d{16,}',  # 长数字(可能是卡号)
                r'\b\d{3}-\d{2}-\d{4}\b',  # SSN 格式
                r'[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}',  # 邮箱
                r'password[:\s]+[^\s]+',  # 密码泄露
            ],
            'required_disclaimers': [
                "本内容由 AI 生成,请谨慎参考。"
            ]
        }

        # 脱敏规则
        self.masking_rules = {
            'phone': lambda x: x[:3] + '****' + x[-4:],
            'email': lambda x: x[:2] + '***@' + x.split('@')[1],
            'id_card': lambda x: x[:6] + '********' + x[-4:],
        }

    def validate(self, output_text):
        """
        验证输出内容

        Returns:
            {
                'is_valid': bool,
                'violations': list,
                'masked_output': str
            }
        """
        violations = []
        masked_output = output_text

        # 长度检查
        if len(output_text) > self.restrictions['max_length']:
            violations.append(f"输出过长,限制 {self.restrictions['max_length']} 字符")
            masked_output = masked_output[:self.restrictions['max_length']]

        # 禁止模式检查
        for pattern in self.restrictions['forbidden_patterns']:
            matches = re.findall(pattern, output_text)
            if matches:
                # 脱敏处理
                for match in matches:
                    if self._is_phone(match):
                        masked_output = masked_output.replace(
                            match,
                            self.masking_rules['phone'](match)
                        )
                    violations.append(f"检测到敏感信息: {match[:10]}...")

        return {
            'is_valid': len(violations) == 0,
            'violations': violations,
            'masked_output': masked_output
        }

    def add_disclaimer(self, output_text):
        """添加免责声明"""
        disclaimer = ' '.join(self.restrictions['required_disclaimers'])
        return f"{output_text}\n\n{disclaimer}"

    def _is_phone(self, text):
        """简单判断是否为手机号"""
        return re.match(r'^1[3-9]\d{9}$', text.strip()) is not None

降级策略

# 模型降级策略
class ModelFallbackStrategy:
    def __init__(self, primary_model, fallback_model, rules_config):
        self.primary_model = primary_model
        self.fallback_model = fallback_model
        self.rules_config = rules_config
        self.fallback_counter = 0

    def predict(self, input_data):
        """
        智能预测,支持降级

        Returns:
            {
                'prediction': result,
                'model_used': 'primary' | 'fallback',
                'fallback_reason': str | None,
                'confidence': float
            }
        """
        # 检查主模型可用性
        primary_available = self._check_primary_availability()

        if primary_available:
            try:
                result = self.primary_model.predict(input_data)
                confidence = self._estimate_confidence(result, input_data)

                # 检查是否需要降级
                if confidence < self.rules_config['min_confidence_threshold']:
                    return self._use_fallback(
                        input_data,
                        reason=f"主模型置信度过低: {confidence:.2f}"
                    )

                return {
                    'prediction': result,
                    'model_used': 'primary',
                    'fallback_reason': None,
                    'confidence': confidence
                }

            except Exception as e:
                return self._use_fallback(input_data, reason=f"主模型异常: {str(e)}")

        else:
            return self._use_fallback(
                input_data,
                reason="主模型不可用"
            )

    def _check_primary_availability(self):
        """检查主模型可用性"""
        # 检查模型服务是否正常
        # 检查资源占用是否过高
        # 检查是否有健康检查失败
        return True  # 简化示例

    def _estimate_confidence(self, result, input_data):
        """估算模型置信度"""
        # 实现置信度估算逻辑
        # 可以基于预测概率、模型不确定性等
        return 0.85  # 简化示例

    def _use_fallback(self, input_data, reason):
        """使用降级模型"""
        self.fallback_counter += 1

        try:
            result = self.fallback_model.predict(input_data)
            return {
                'prediction': result,
                'model_used': 'fallback',
                'fallback_reason': reason,
                'confidence': 0.5  # 降级模型置信度通常较低
            }
        except Exception as e:
            # 如果降级模型也失败,返回兜底响应
            return {
                'prediction': self._get_default_response(),
                'model_used': 'default',
                'fallback_reason': f"主模型和降级模型都失败: {str(e)}",
                'confidence': 0.0
            }

    def _get_default_response(self):
        """获取默认响应"""
        return {
            'error': '服务暂时不可用,请稍后重试',
            'timestamp': int(time.time())
        }

    def get_fallback_stats(self):
        """获取降级统计"""
        return {
            'total_fallbacks': self.fallback_counter,
            'fallback_rate': self.fallback_counter / max(1, self._get_total_predictions())
        }

监控告警

实时监控

# 模型实时监控系统
class ModelMonitor:
    def __init__(self, model_name, alert_config):
        self.model_name = model_name
        self.alert_config = alert_config
        self.metrics_buffer = deque(maxlen=1000)

        # 启动监控线程
        self.monitor_thread = threading.Thread(
            target=self._monitor_loop,
            daemon=True
        )
        self.monitor_thread.start()

    def record_prediction(self, prediction, ground_truth=None):
        """记录预测结果"""
        metric = {
            'timestamp': time.time(),
            'model': self.model_name,
            'prediction': prediction,
            'ground_truth': ground_truth
        }

        if ground_truth is not None:
            metric['accuracy'] = int(prediction == ground_truth)

        self.metrics_buffer.append(metric)

    def _monitor_loop(self):
        """监控循环"""
        while True:
            time.sleep(self.alert_config['check_interval'])

            # 检查各种指标
            self._check_accuracy_drift()
            self._check_latency()
            self._check_error_rate()
            self._check_data_drift()

    def _check_accuracy_drift(self):
        """检查准确率漂移"""
        recent_metrics = list(self.metrics_buffer)[-100:]  # 最近 100 个

        if not recent_metrics:
            return

        accuracies = [m.get('accuracy', 0) for m in recent_metrics]
        current_accuracy = np.mean(accuracies)

        if current_accuracy < self.alert_config['min_accuracy']:
            self._send_alert(
                alert_type='accuracy_drift',
                message=f"模型准确率下降: {current_accuracy:.2%}",
                severity='high',
                metrics={
                    'current_accuracy': current_accuracy,
                    'threshold': self.alert_config['min_accuracy']
                }
            )

    def _check_latency(self):
        """检查延迟"""
        # 实现 P95/P99 延迟检查
        pass

    def _check_error_rate(self):
        """检查错误率"""
        # 实现错误率检查
        pass

    def _check_data_drift(self):
        """检查数据漂移"""
        # 实现输入数据分布漂移检查
        pass

    def _send_alert(self, alert_type, message, severity, metrics):
        """发送告警"""
        # 发送到告警系统(Prometheus Alertmanager、钉钉、企业微信等)
        print(f"[{severity.upper()}] {alert_type}: {message}")
        print(f"Metrics: {json.dumps(metrics, indent=2)}")

# 监控配置
alert_config = {
    'check_interval': 60,  # 每 60 秒检查一次
    'min_accuracy': 0.85,  # 最低准确率
    'max_p95_latency': 500,  # P95 延迟阈值(毫秒)
    'max_error_rate': 0.05  # 最大错误率
}

# 使用示例
monitor = ModelMonitor('recommendation_v2', alert_config)

# 在预测时记录
for user_id, features in test_data:
    prediction = model.predict(features)
    ground_truth = get_ground_truth(user_id)

    monitor.record_prediction(prediction, ground_truth)

告警规则配置

# Prometheus 告警规则
groups:
  - name: model_performance
    interval: 1m
    rules:
      # 准确率下降告警
      - alert: ModelAccuracyDrift
        expr: |
          (
            avg_over_time(model_accuracy[5m])
          ) < 0.85
        for: 2m
        labels:
          severity: critical
          team: ml-team
        annotations:
          summary: "模型准确率下降告警"
          description: "模型 {{ $labels.model_name }} 准确率低于 85%,当前值:{{ $value }}"

      # 延迟过高告警
      - alert: HighModelLatency
        expr: |
          (
            histogram_quantile(0.95,
              rate(model_inference_latency_bucket[5m])
            )
          ) > 500
        for: 1m
        labels:
          severity: warning
          team: ml-team
        annotations:
          summary: "模型推理延迟过高"
          description: "模型 {{ $labels.model_name }} P95 延迟超过 500ms,当前值:{{ $value }}ms"

      # 错误率过高告警
      - alert: HighModelErrorRate
        expr: |
          (
            rate(model_errors_total[5m])
            /
            rate(model_predictions_total[5m])
          ) > 0.05
        for: 2m
        labels:
          severity: warning
          team: ml-team
        annotations:
          summary: "模型错误率过高"
          description: "模型 {{ $labels.model_name }} 错误率超过 5%,当前值:{{ $value }}"

  - name: data_quality
    interval: 5m
    rules:
      # 数据分布漂移告警
      - alert: DataDistributionDrift
        expr: |
          (
            data_drift_score
          ) > 0.3
        for: 10m
        labels:
          severity: warning
          team: data-team
        annotations:
          summary: "数据分布发生漂移"
          description: "特征 {{ $labels.feature_name }} 分布漂移分数超过 0.3,当前值:{{ $value }}"

      # 异常值比例过高告警
      - alert: HighOutlierRatio
        expr: |
          (
            outlier_count
            /
            total_input_count
          ) > 0.1
        for: 5m
        labels:
          severity: info
          team: data-team
        annotations:
          summary: "异常值比例过高"
          description: "输入数据异常值比例超过 10%,当前值:{{ $value }}%"

踩过的坑

坑一:数据漂移导致模型失效

推荐模型上线初期表现不错,两个月后 CTR 突然下降 30%。查了半天才发现,用户群体发生了变化——平台做了一个新活动,引入了大量年轻用户,但模型训练数据主要基于老用户群体。

解决:上线后持续监控输入数据分布,发现显著漂移时触发模型重训练。

# 数据漂移监控系统
class DataDriftMonitor:
    def __init__(self, reference_data, drift_threshold=0.3):
        self.reference_data = reference_data
        self.drift_threshold = drift_threshold
        self.drift_detector = DriftDetector()

    def check_drift(self, current_data):
        """
        检查数据漂移

        Returns:
            {
                'has_drift': bool,
                'drift_score': float,
                'drifted_features': list
            }
        """
        drifted_features = []

        for feature in self.reference_data.columns:
            # 使用 KS 检验检测数值型特征漂移
            if self.reference_data[feature].dtype in [np.int64, np.float64]:
                ks_stat, p_value = stats.ks_2samp(
                    self.reference_data[feature].dropna(),
                    current_data[feature].dropna()
                )

                # 使用 Population Stability Index (PSI) 检测类别型特征漂移
                psi_score = self._calculate_psi(
                    self.reference_data[feature],
                    current_data[feature]
                )

                if p_value < 0.05 or psi_score > self.drift_threshold:
                    drifted_features.append({
                        'feature': feature,
                        'ks_statistic': ks_stat,
                        'p_value': p_value,
                        'psi_score': psi_score
                    })

        overall_drift_score = len(drifted_features) / len(self.reference_data.columns)

        return {
            'has_drift': len(drifted_features) > 0,
            'drift_score': overall_drift_score,
            'drifted_features': drifted_features
        }

    def _calculate_psi(self, expected, actual, bins=10):
        """
        计算群体稳定性指数(PSI)

        PSI < 0.1: 无显著漂移
        0.1 <= PSI < 0.25: 轻微漂移
        PSI >= 0.25: 显著漂移
        """
        def calculate_bins(data, bins):
            # 计算分箱边界
            _, bin_edges = np.histogram(data, bins=bins)
            return bin_edges

        def calculate_psi_values(expected, actual, bin_edges):
            # 基于 expected 数据的分箱计算 PSI
            expected_counts, _ = np.histogram(expected, bins=bin_edges)
            actual_counts, _ = np.histogram(actual, bins=bin_edges)

            # 归一化为比例
            expected_percents = expected_counts / expected_counts.sum()
            actual_percents = actual_counts / actual_counts.sum()

            # 添加小值避免除零
            expected_percents = np.maximum(expected_percents, 0.0001)
            actual_percents = np.maximum(actual_percents, 0.0001)

            # 计算 PSI
            psi_values = (actual_percents - expected_percents) * np.log(
                actual_percents / expected_percents
            )

            return psi_values.sum()

        bin_edges = calculate_bins(expected, bins)
        psi = calculate_psi_values(expected, actual, bin_edges)

        return psi

# 使用示例
# 训练时保存参考数据分布
reference_data = train_data[['age', 'gender', 'city_tier', 'user_level']]

# 运行时定期检查漂移
monitor = DataDriftMonitor(reference_data, drift_threshold=0.25)

# 每天检查一次
daily_data = collect_daily_user_data()
drift_result = monitor.check_drift(daily_data)

if drift_result['has_drift']:
    print(f"检测到数据漂移,分数: {drift_result['drift_score']:.2f}")
    print(f"漂移特征: {drift_result['drifted_features']}")

    # 触发模型重训练
    trigger_model_retraining()

坑二:模型过拟合特定子群体

风控模型上线后,发现对某些年龄段的拒绝率异常高。排查发现模型对训练数据中某些子群体过拟合,导致泛化能力差。

解决:训练时进行子群体分析,确保模型在各个子群体上的表现均衡。

# 子群体公平性分析
class SubgroupFairnessAnalyzer:
    def __init__(self, sensitive_attributes):
        self.sensitive_attributes = sensitive_attributes

    def analyze(self, X, y_true, y_pred):
        """
        分析模型在各子群体上的表现

        Returns:
            {
                'overall_metrics': dict,
                'subgroup_metrics': dict,
                'fairness_gaps': dict
            }
        """
        # 整体指标
        overall_metrics = self._calculate_metrics(y_true, y_pred)

        # 子群体指标
        subgroup_metrics = {}
        for attr in self.sensitive_attributes:
            attr_values = X[attr].unique()

            for value in attr_values:
                mask = X[attr] == value
                subgroup_y_true = y_true[mask]
                subgroup_y_pred = y_pred[mask]

                if len(subgroup_y_true) > 0:
                    subgroup_metrics[f"{attr}_{value}"] = self._calculate_metrics(
                        subgroup_y_true,
                        subgroup_y_pred
                    )

        # 公平性差距
        fairness_gaps = self._calculate_fairness_gaps(subgroup_metrics)

        return {
            'overall_metrics': overall_metrics,
            'subgroup_metrics': subgroup_metrics,
            'fairness_gaps': fairness_gaps
        }

    def _calculate_metrics(self, y_true, y_pred):
        """计算各项指标"""
        from sklearn.metrics import (
            accuracy_score, precision_score, recall_score,
            f1_score, roc_auc_score, confusion_matrix
        )

        metrics = {
            'accuracy': accuracy_score(y_true, y_pred),
            'precision': precision_score(y_true, y_pred, average='weighted'),
            'recall': recall_score(y_true, y_pred, average='weighted'),
            'f1': f1_score(y_true, y_pred, average='weighted')
        }

        # 如果是二分类,额外计算 ROC-AUC
        if len(set(y_true)) == 2:
            try:
                metrics['roc_auc'] = roc_auc_score(y_true, y_pred)
            except:
                pass

        # 混淆矩阵
        cm = confusion_matrix(y_true, y_pred)
        metrics['confusion_matrix'] = cm.tolist()

        # 各类别的 TPR/FPR
        for i, label in enumerate(sorted(set(y_true))):
            tp = cm[i, i]
            fn = cm[i, :].sum() - tp
            fp = cm[:, i].sum() - tp
            tn = cm.sum() - tp - fn - fp

            metrics[f'tpr_class_{label}'] = tp / (tp + fn) if (tp + fn) > 0 else 0
            metrics[f'fpr_class_{label}'] = fp / (fp + tn) if (fp + tn) > 0 else 0

        return metrics

    def _calculate_fairness_gaps(self, subgroup_metrics):
        """计算公平性差距"""
        gaps = {}

        # 找出同一敏感属性下的不同子群体
        attr_groups = {}
        for key in subgroup_metrics.keys():
            parts = key.split('_')
            if len(parts) >= 2:
                attr = parts[0]
                if attr not in attr_groups:
                    attr_groups[attr] = []
                attr_groups[attr].append(key)

        # 计算差距
        for attr, groups in attr_groups.items():
            if len(groups) >= 2:
                # 计算 TPR 差距
                tpr_values = [
                    subgroup_metrics[g].get('tpr_class_1', 0)
                    for g in groups
                    if 'tpr_class_1' in subgroup_metrics[g]
                ]

                if tpr_values:
                    gaps[f'tpr_gap_{attr}'] = max(tpr_values) - min(tpr_values)

                # 计算 FPR 差距
                fpr_values = [
                    subgroup_metrics[g].get('fpr_class_1', 0)
                    for g in groups
                    if 'fpr_class_1' in subgroup_metrics[g]
                ]

                if fpr_values:
                    gaps[f'fpr_gap_{attr}'] = max(fpr_values) - min(fpr_values)

        return gaps

# 使用示例
sensitive_attributes = ['age_group', 'gender', 'region']
analyzer = SubgroupFairnessAnalyzer(sensitive_attributes)

# 分析模型表现
fairness_result = analyzer.analyze(
    test_data,
    test_data['true_label'],
    test_data['predicted_label']
)

print("整体指标:")
print(json.dumps(fairness_result['overall_metrics'], indent=2))

print("\n子群体指标:")
for subgroup, metrics in fairness_result['subgroup_metrics'].items():
    print(f"{subgroup}: Accuracy={metrics['accuracy']:.2%}")

print("\n公平性差距:")
print(json.dumps(fairness_result['fairness_gaps'], indent=2))

# 如果差距过大,需要调整模型或添加约束
if any(gap > 0.1 for gap in fairness_result['fairness_gaps'].values()):
    print("警告:检测到显著的公平性差距")
    # 考虑使用公平性约束训练或后处理调整

坑三:监控告警配置不当

第一次上监控时,告警阈值设得太敏感,导致半夜频繁收到误报,团队最后把告警全关了。没过几天,一个真正的问题发生时没收到告警,直到用户投诉才发现。

解决:先收集基线数据,再根据业务需求合理设置阈值,逐步调优。

# 智能告警阈值设置
class SmartAlertThreshold:
    def __init__(self, historical_data_window=30):
        self.historical_data_window = historical_data_window
        self.historical_metrics = deque(maxlen=historical_data_window)

    def update_baseline(self, metrics):
        """更新基线数据"""
        self.historical_metrics.append({
            'timestamp': time.time(),
            'metrics': metrics
        })

    def calculate_thresholds(self, sensitivity='medium'):
        """
        计算动态阈值

        sensitivity:
        - low: 宽松阈值(P99)
        - medium: 适中阈值(P95)
        - high: 严格阈值(P90)
        """
        if len(self.historical_metrics) < 7:
            # 数据不足时返回默认阈值
            return self._get_default_thresholds()

        thresholds = {}
        sensitivity_map = {'low': 0.99, 'medium': 0.95, 'high': 0.90}
        percentile = sensitivity_map.get(sensitivity, 0.95)

        # 收集各项指标的历史值
        metric_history = defaultdict(list)
        for record in self.historical_metrics:
            for key, value in record['metrics'].items():
                if isinstance(value, (int, float)):
                    metric_history[key].append(value)

        # 计算阈值
        for metric, values in metric_history.items():
            if len(values) >= 10:  # 至少有 10 个数据点
                lower_threshold = np.percentile(values, (1 - percentile) * 100)
                upper_threshold = np.percentile(values, percentile * 100)

                # 添加一些缓冲
                buffer = (upper_threshold - lower_threshold) * 0.1
                upper_threshold += buffer
                lower_threshold -= buffer

                thresholds[metric] = {
                    'lower': lower_threshold,
                    'upper': upper_threshold,
                    'percentile': percentile
                }

        return thresholds

    def check_alert(self, current_metrics, sensitivity='medium'):
        """
        检查是否需要告警

        Returns:
            {
                'should_alert': bool,
                'triggered_metrics': list,
                'details': list
            }
        """
        thresholds = self.calculate_thresholds(sensitivity)
        triggered_metrics = []
        details = []

        for metric, current_value in current_metrics.items():
            if isinstance(current_value, (int, float)) and metric in thresholds:
                lower = thresholds[metric]['lower']
                upper = thresholds[metric]['upper']

                if current_value < lower or current_value > upper:
                    triggered_metrics.append(metric)

                    if current_value < lower:
                        direction = 'below'
                        threshold_value = lower
                    else:
                        direction = 'above'
                        threshold_value = upper

                    details.append({
                        'metric': metric,
                        'current_value': current_value,
                        'threshold': threshold_value,
                        'direction': direction,
                        'deviation': abs(current_value - threshold_value) / threshold_value
                    })

        return {
            'should_alert': len(triggered_metrics) > 0,
            'triggered_metrics': triggered_metrics,
            'details': details
        }

    def _get_default_thresholds(self):
        """获取默认阈值"""
        return {
            'accuracy': {'lower': 0.80, 'upper': 1.0},
            'latency_p95': {'lower': 0, 'upper': 1000},
            'error_rate': {'lower': 0, 'upper': 0.1}
        }

# 使用示例
# 1. 收集历史数据(至少 7 天)
threshold_manager = SmartAlertThreshold(historical_data_window=30)

# 2. 每天更新基线
for day_data in historical_metrics_data:
    threshold_manager.update_baseline(day_data)

# 3. 实时检查
current_metrics = {
    'accuracy': 0.82,
    'latency_p95': 850,
    'error_rate': 0.08,
    'throughput': 1200
}

alert_result = threshold_manager.check_alert(current_metrics, sensitivity='medium')

if alert_result['should_alert']:
    print("触发告警:")
    for detail in alert_result['details']:
        print(f"  - {detail['metric']}: {detail['current_value']} "
              f"{detail['direction']} 阈值 {detail['threshold']} "
              f"(偏差: {detail['deviation']:.1%})")

写在最后

AI 风险管理不是一次性工作,是持续的过程。

解决了

  • 模型失效能及时发现
  • 数据漂移能自动触发重训练
  • 子群体公平性问题能提前识别
  • 监控告警不会再半夜误报

留下了

  • 模型解释能力还是不够强
  • 复杂场景的降级策略还需要更多实践
  • 跨模型的系统性风险治理还不够成熟

不是说有了这些工具就万事大吉。AI 系统再智能,也还是需要人去设计、去监控、去兜底。风险管理做得到位,AI 才能真正发挥价值,而不是变成定时炸弹。


这次 AI 风险管理体系搭建花了三个月,从识别到应对再到监控。落地后模型稳定性明显提升,线上事故减少了 70%,团队对模型上线的信心也强了很多。

版权声明: 本文首发于 指尖魔法屋-AI风险管理折腾手记https://blog.thinkmoon.cn/post/196-ai-risk-management-identification-response-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!