Python中如何批量获取SimulationPoint类多实例的统计指标并分析
优化方案
一、批量提取统计指标的最简实现
你完全可以结合列表推导和pandas的内置统计能力实现需求,不用自己循环累加计算。300个数据量很小,该方案足够高效,代码可读性也更高:
import os import pandas as pd import numpy as np from scipy.stats import kurtosis, skew # 1. 批量实例化,保留文件名与实例的对应关系 root_path = '../path_to_data/' sim_points = {} for dp_name in os.listdir(root_path): dp_path = os.path.join(root_path, dp_name) # 过滤非目录文件,避免读取到无关文件报错 if not os.path.isdir(dp_path): continue # 注意:你需要补全scalar1、scalar2的获取逻辑,可从文件名/目录下的配置文件读取 scalar1 = parse_scalar1_from_path(dp_path) scalar2 = parse_scalar2_from_path(dp_path) sim_points[dp_name] = SimulationPoint(dp_path, scalar1, scalar2) # 2. 一行代码提取所有需要的指标为DataFrame stats_df = pd.DataFrame([ { "filename": fname, "normalized_input": sp.normalize_input(), "scalar1": sp.scalar1, "scalar2": sp.scalar2, "mesh_value_mean": sp.mesh["value"].mean(), "mesh_value_std": sp.mesh["value"].std(), "spot_centers": sp.spot_centers(nbr_spots=5) } for fname, sp in sim_points.items() ])
后续获取统计量、绘图直接调用pandas内置方法即可:
- 基础统计(均值、标准差、分位数等):直接执行
stats_df.describe() - 自定义统计(峰度、偏度等):
stats_df["normalized_input"].agg([np.mean, np.std, kurtosis, skew]) - 绘制直方图:
stats_df["normalized_input"].plot.hist(bins=20)
二、自定义类的优化建议
针对你当前的SimulationPoint实现,可以做以下优化提升稳定性和性能:
- 路径拼接规范:不要直接用字符串拼接路径,改用
os.path.join避免不同操作系统路径分隔符不兼容的问题:
# 原代码:pd.read_csv(path+'mesh.csv') pd.read_csv(os.path.join(path, 'mesh.csv'))
- 懒加载+缓存避免重复计算:网格数据读取、方法计算结果可以加缓存,多次调用不会重复执行IO或计算逻辑,能大幅节省资源:
from functools import lru_cache class SimulationPoint: def __init__(self, path, scalar1, scalar2): self.scalar1 = scalar1 self.scalar2 = scalar2 self.mesh_path = os.path.join(path, 'mesh.csv') self._mesh = None # 私有变量,存储懒加载的网格数据 @property def mesh(self): # 只有第一次调用mesh属性时才会读取csv if self._mesh is None: self._mesh = pd.read_csv(self.mesh_path) return self._mesh @lru_cache(maxsize=None) def normalize_input(self): # 缓存计算结果,多次调用只算一次 return func1(self.scalar1) @lru_cache(maxsize=None) def spot_centers(self, nbr_spots=5): points_mesh = self.mesh[self.mesh.value >= self.mesh.value.quantile(0.9)].copy() xyz_mesh = points_mesh.drop(['value'], axis=1).to_numpy() return func2(xyz_mesh, n_clusters=nbr_spots)
- 封装数据集类简化批量操作:如果后续经常需要做批量处理,可以把实例化、统计提取的逻辑封装成
SimulationDataset类,调用更简洁:
class SimulationDataset: def __init__(self, root_path): self.root_path = root_path self.points = self._load_all_points() def _load_all_points(self): points = {} for dp_name in os.listdir(self.root_path): dp_path = os.path.join(self.root_path, dp_name) if not os.path.isdir(dp_path): continue # 补全scalar1、scalar2的读取逻辑 s1 = parse_scalar1_from_path(dp_path) s2 = parse_scalar2_from_path(dp_path) points[dp_name] = SimulationPoint(dp_path, s1, s2) return points def get_stats_df(self): return pd.DataFrame([ {"filename": fname, "normalized_input": sp.normalize_input(), "scalar1": sp.scalar1} for fname, sp in self.points.items() ]) # 调用示例 dataset = SimulationDataset('../path_to_data/') stats_df = dataset.get_stats_df() print(stats_df.describe())
三、更大数据量下的性能优化
如果后续数据量上涨到数千个以上,可以用多线程加速IO密集型的实例化过程:
from concurrent.futures import ThreadPoolExecutor def load_single_point(dp_name, root_path): dp_path = os.path.join(root_path, dp_name) if not os.path.isdir(dp_path): return None s1 = parse_scalar1_from_path(dp_path) s2 = parse_scalar2_from_path(dp_path) return (dp_name, SimulationPoint(dp_path, s1, s2)) root_path = '../path_to_data/' with ThreadPoolExecutor(max_workers=8) as executor: tasks = [executor.submit(load_single_point, dp, root_path) for dp in os.listdir(root_path)] sim_points = {} for task in tasks: res = task.result() if res: sim_points[res[0]] = res[1]
内容的提问来源于stack exchange,提问作者coyote
相关产品推荐
相关产品推荐

