Python统计工具案例如何封装数据统计

wen python案例 30

本文目录导读:

Python统计工具案例如何封装数据统计

  1. 基础统计分析类
  2. 高级分组统计分析
  3. 时间序列统计器
  4. 数据统计工具类(整合版)
  5. 使用示例

我来分享几个Python数据统计工具的封装案例,从简单到复杂逐步演示。

基础统计分析类

import numpy as np
import pandas as pd
from typing import List, Dict, Union, Optional
import json
from datetime import datetime
class BasicStatistics:
    """基础统计分析器"""
    def __init__(self, data: Union[List, np.ndarray, pd.Series]):
        """
        初始化统计器
        Args:
            data: 输入数据
        """
        self.data = np.array(data) if not isinstance(data, np.ndarray) else data
        self._validate_data()
    def _validate_data(self):
        """验证数据有效性"""
        if len(self.data) == 0:
            raise ValueError("数据不能为空")
        if np.any(np.isnan(self.data)):
            raise ValueError("数据包含NaN值")
    def get_basic_stats(self) -> Dict:
        """获取基本统计量"""
        return {
            'count': len(self.data),
            'mean': np.mean(self.data),
            'std': np.std(self.data, ddof=1),
            'min': np.min(self.data),
            'max': np.max(self.data),
            'median': np.median(self.data),
            'q1': np.percentile(self.data, 25),
            'q3': np.percentile(self.data, 75),
            'iqr': np.percentile(self.data, 75) - np.percentile(self.data, 25),
            'skewness': self._skewness(),
            'kurtosis': self._kurtosis()
        }
    def _skewness(self) -> float:
        """计算偏度"""
        n = len(self.data)
        mean = np.mean(self.data)
        std = np.std(self.data, ddof=1)
        return np.sum((self.data - mean) ** 3) / (n * std ** 3)
    def _kurtosis(self) -> float:
        """计算峰度"""
        n = len(self.data)
        mean = np.mean(self.data)
        std = np.std(self.data, ddof=1)
        return np.sum((self.data - mean) ** 4) / (n * std ** 4) - 3
    def get_outliers(self, method='iqr', threshold=1.5) -> Dict:
        """
        检测异常值
        Args:
            method: 检测方法 ('iqr' 或 'zscore')
            threshold: 阈值
        Returns:
            异常值信息
        """
        if method == 'iqr':
            q1 = np.percentile(self.data, 25)
            q3 = np.percentile(self.data, 75)
            iqr = q3 - q1
            lower_bound = q1 - threshold * iqr
            upper_bound = q3 + threshold * iqr
            outliers = self.data[(self.data < lower_bound) | (self.data > upper_bound)]
        elif method == 'zscore':
            z_scores = np.abs((self.data - np.mean(self.data)) / np.std(self.data))
            outliers = self.data[z_scores > threshold]
        else:
            raise ValueError(f"不支持的检测方法: {method}")
        return {
            'count': len(outliers),
            'indices': np.where(np.isin(self.data, outliers))[0].tolist(),
            'values': outliers.tolist(),
            'method': method,
            'threshold': threshold
        }
    def summary(self) -> str:
        """生成统计摘要"""
        stats = self.get_basic_stats()
        outliers = self.get_outliers()
        summary = f"""
        ========== 统计摘要 ==========
        数据数量: {stats['count']}
        均值: {stats['mean']:.4f}
        标准差: {stats['std']:.4f}
        中位数: {stats['median']:.4f}
        最小值: {stats['min']:.4f}
        最大值: {stats['max']:.4f}
        四分位距: {stats['iqr']:.4f}
        偏度: {stats['skewness']:.4f}
        峰度: {stats['kurtosis']:.4f}
        --------------------------
        异常值统计:
        数量: {outliers['count']}
        阈值方法: {outliers['method']}
        ===========================
        """
        return summary

高级分组统计分析

class GroupStatistics:
    """分组统计分析器"""
    def __init__(self, df: pd.DataFrame):
        """
        初始化分组统计器
        Args:
            df: 包含分组变量和数值变量的DataFrame
        """
        self.df = df.copy()
        self._validate_columns()
    def _validate_columns(self):
        """验证数据结构"""
        required_attrs = ['group_col', 'value_col']
        for attr in required_attrs:
            if not hasattr(self, attr):
                self.group_col = None
                self.value_col = None
    def set_variables(self, group_col: str, value_col: str):
        """
        设置分组变量和数值变量
        Args:
            group_col: 分组列名
            value_col: 数值列名
        """
        if group_col not in self.df.columns:
            raise ValueError(f"分组列 '{group_col}' 不存在")
        if value_col not in self.df.columns:
            raise ValueError(f"数值列 '{value_col}' 不存在")
        self.group_col = group_col
        self.value_col = value_col
    def get_group_stats(self, agg_funcs: List[str] = None) -> pd.DataFrame:
        """
        计算分组统计量
        Args:
            agg_funcs: 聚合函数列表,默认使用常用统计量
        Returns:
            分组统计结果DataFrame
        """
        if agg_funcs is None:
            agg_funcs = ['count', 'mean', 'std', 'min', 'max', 'median']
        # 自定义分位数函数
        def q25(x): return x.quantile(0.25)
        def q75(x): return x.quantile(0.75)
        agg_dict = {
            'count': 'count',
            'mean': 'mean',
            'std': 'std',
            'min': 'min',
            'max': 'max',
            'median': 'median',
            'q25': q25,
            'q75': q75
        }
        # 只选择需要的聚合函数
        selected_aggs = {func: agg_dict[func] for func in agg_funcs if func in agg_dict}
        return self.df.groupby(self.group_col)[self.value_col].agg(selected_aggs).round(3)
    def get_groups_summary(self) -> Dict:
        """
        获取分组摘要
        Returns:
            各组统计信息
        """
        group_stats = self.get_group_stats()
        summary = {}
        for group in group_stats.index:
            stats = group_stats.loc[group]
            group_data = self.df[self.df[self.group_col] == group][self.value_col]
            summary[group] = {
                'stats': stats.to_dict(),
                'outliers': BasicStatistics(group_data.values).get_outliers()
            }
        return summary
    def compare_groups(self, method='anova') -> Dict:
        """
        组间比较
        Args:
            method: 比较方法 ('anova' 或 'kruskal')
        Returns:
            比较结果
        """
        from scipy import stats as scipy_stats
        groups = [group[self.value_col].values 
                 for _, group in self.df.groupby(self.group_col)]
        if method == 'anova':
            stat, p_value = scipy_stats.f_oneway(*groups)
            test_name = '单因素方差分析'
        elif method == 'kruskal':
            stat, p_value = scipy_stats.kruskal(*groups)
            test_name = 'Kruskal-Wallis检验'
        else:
            raise ValueError(f"不支持的方法: {method}")
        return {
            'test': test_name,
            'statistic': stat,
            'p_value': p_value,
            'significant': p_value < 0.05,
            'groups': list(self.df[self.group_col].unique())
        }

时间序列统计器

class TimeSeriesStatistics:
    """时间序列统计分析器"""
    def __init__(self, dates: List, values: List):
        """
        初始化时间序列统计器
        Args:
            dates: 日期列表
            values: 数值列表
        """
        self.dates = pd.to_datetime(dates)
        self.values = np.array(values)
        self.ts = pd.Series(values, index=dates)
    def get_trend_stats(self) -> Dict:
        """
        计算趋势统计量
        """
        from scipy import stats
        x = np.arange(len(self.values))
        slope, intercept, r_value, p_value, std_err = stats.linregress(x, self.values)
        return {
            'slope': slope,
            'intercept': intercept,
            'r_squared': r_value ** 2,
            'p_value': p_value,
            'std_error': std_err,
            'trend': '上升' if slope > 0 else '下降',
            'trend_strength': '强' if abs(r_value) > 0.7 else '中等' if abs(r_value) > 0.3 else '弱'
        }
    def get_seasonal_stats(self, period: int = 12) -> Dict:
        """
        计算季节性统计
        Args:
            period: 周期长度
        Returns:
            季节性统计信息
        """
        # 移动平均
        ma = pd.Series(self.values).rolling(window=period).mean()
        # 季节性分解
        seasonal_component = self.values / ma.values if len(self.values) == len(ma) else None
        return {
            'period': period,
            'has_seasonality': len(self.values) >= 2 * period,
            'moving_average': ma.dropna().tolist(),
            'seasonal_strength': np.std(seasonal_component) if seasonal_component is not None else None
        }
    def forecast(self, steps: int = 5) -> Dict:
        """
        简单预测
        Args:
            steps: 预测步数
        Returns:
            预测结果
        """
        # 使用简单指数平滑
        from statsmodels.tsa.holtwinters import SimpleExpSmoothing
        model = SimpleExpSmoothing(self.values)
        fitted = model.fit()
        forecast_values = fitted.forecast(steps)
        # 生成未来日期
        last_date = self.dates[-1]
        future_dates = pd.date_range(start=last_date, periods=steps + 1, freq='D')[1:]
        return {
            'dates': future_dates.tolist(),
            'forecast': forecast_values.tolist(),
            'model_params': {
                'smoothing_level': fitted.params['smoothing_level']
            }
        }

数据统计工具类(整合版)

class StatisticsTool:
    """综合性数据统计工具"""
    def __init__(self, data_source: Union[str, pd.DataFrame, List, Dict]):
        """
        初始化统计工具
        Args:
            data_source: 数据源(文件路径、DataFrame、列表或字典)
        """
        self.data = self._load_data(data_source)
        self.statistics = {}
        self.timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
    def _load_data(self, data_source):
        """加载数据"""
        if isinstance(data_source, str):
            # 假设是CSV文件
            return pd.read_csv(data_source)
        elif isinstance(data_source, pd.DataFrame):
            return data_source
        elif isinstance(data_source, list):
            return pd.DataFrame({'values': data_source})
        elif isinstance(data_source, dict):
            return pd.DataFrame(data_source)
        else:
            raise TypeError("不支持的数据源类型")
    def analyze_numeric_column(self, column: str) -> Dict:
        """
        分析数值列
        Args:
            column: 列名
        Returns:
            分析结果
        """
        if column not in self.data.columns:
            raise ValueError(f"列 '{column}' 不存在")
        values = self.data[column].dropna()
        basic_stats = BasicStatistics(values)
        result = {
            'column': column,
            'basic_stats': basic_stats.get_basic_stats(),
            'outliers': basic_stats.get_outliers(),
            'data_type': str(self.data[column].dtype),
            'missing_count': int(self.data[column].isnull().sum()),
            'missing_pct': round(self.data[column].isnull().sum() / len(self.data) * 100, 2)
        }
        self.statistics[column] = result
        return result
    def analyze_categorical_column(self, column: str) -> Dict:
        """
        分析分类列
        Args:
            column: 列名
        Returns:
            分析结果
        """
        if column not in self.data.columns:
            raise ValueError(f"列 '{column}' 不存在")
        value_counts = self.data[column].value_counts()
        return {
            'column': column,
            'unique_values': len(value_counts),
            'top_values': value_counts.head(5).to_dict(),
            'missing_count': int(self.data[column].isnull().sum()),
            'missing_pct': round(self.data[column].isnull().sum() / len(self.data) * 100, 2)
        }
    def correlation_analysis(self, columns: List[str] = None) -> pd.DataFrame:
        """
        相关性分析
        Args:
            columns: 要分析的列列表
        Returns:
            相关性矩阵
        """
        if columns:
            corr_data = self.data[columns]
        else:
            corr_data = self.data.select_dtypes(include=[np.number])
        return corr_data.corr().round(4)
    def generate_report(self) -> Dict:
        """
        生成完整统计报告
        """
        report = {
            'generated_at': self.timestamp,
            'dataset_info': {
                'rows': len(self.data),
                'columns': len(self.data.columns),
                'column_names': list(self.data.columns),
                'memory_usage': f"{self.data.memory_usage(deep=True).sum() / 1024 / 1024:.2f} MB"
            },
            'column_analysis': {},
            'correlations': None
        }
        # 分析每列
        for col in self.data.columns:
            if pd.api.types.is_numeric_dtype(self.data[col]):
                report['column_analysis'][col] = self.analyze_numeric_column(col)
            else:
                report['column_analysis'][col] = self.analyze_categorical_column(col)
        # 相关性分析(仅数值列)
        numeric_cols = self.data.select_dtypes(include=[np.number]).columns
        if len(numeric_cols) > 1:
            corr_matrix = self.correlation_analysis(numeric_cols.tolist())
            report['correlations'] = corr_matrix.to_dict()
        return report
    def export_to_json(self, filepath: str = 'statistics_report.json'):
        """
        导出统计报告到JSON文件
        Args:
            filepath: 输出文件路径
        """
        report = self.generate_report()
        # 处理DataFrame对象
        def json_serializable(obj):
            if isinstance(obj, (np.integer, np.floating)):
                return obj.item()
            elif isinstance(obj, np.ndarray):
                return obj.tolist()
            elif isinstance(obj, (pd.Series, pd.DataFrame)):
                return obj.to_dict()
            return obj
        # 递归转换
        def convert_to_serializable(d):
            for key, value in d.items():
                if isinstance(value, dict):
                    d[key] = convert_to_serializable(value)
                else:
                    d[key] = json_serializable(value)
            return d
        report = convert_to_serializable(report)
        with open(filepath, 'w', encoding='utf-8') as f:
            json.dump(report, f, ensure_ascii=False, indent=2)
        print(f"统计报告已保存到: {filepath}")

使用示例

# 示例1:基础统计
data = [12, 15, 14, 10, 18, 22, 16, 11, 9, 20]
stats = BasicStatistics(data)
print(stats.summary())
# 示例2:分组统计
df = pd.DataFrame({
    'group': ['A', 'B', 'A', 'B', 'A', 'B'] * 5,
    'value': np.random.normal(100, 15, 30)
})
group_stats = GroupStatistics(df)
group_stats.set_variables('group', 'value')
print(group_stats.get_group_stats())
# 示例3:综合统计工具
tool = StatisticsTool(df)
report = tool.generate_report()
print(json.dumps(report, indent=2, ensure_ascii=False))
  1. 面向对象设计:使用类封装数据和相关方法
  2. 类型安全:添加类型注解和验证
  3. 异常处理:完善的数据验证和错误处理
  4. 灵活性:支持多种数据源和参数配置
  5. 可扩展性:模块化设计,易于添加新功能
  6. 输出格式化:提供多种输出格式(字典、DataFrame、JSON)

这样的封装既保证了代码的可复用性,又提供了清晰的接口供其他模块调用。

抱歉,评论功能暂时关闭!