本文目录导读:

我来分享几个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))
- 面向对象设计:使用类封装数据和相关方法
- 类型安全:添加类型注解和验证
- 异常处理:完善的数据验证和错误处理
- 灵活性:支持多种数据源和参数配置
- 可扩展性:模块化设计,易于添加新功能
- 输出格式化:提供多种输出格式(字典、DataFrame、JSON)
这样的封装既保证了代码的可复用性,又提供了清晰的接口供其他模块调用。