Python多播统计客户端脚本性能优化咨询:提速analyze方法、内存优化及Numba应用方案
首先,咱们先拆解你的核心问题:当前analyze方法每次都要对全量原始数据调用scipy.stats的方法,随着数据量增长,计算耗时和内存占用会线性飙升。下面分模块给出具体可落地的优化方案:
一、通过预处理/缓存统计量提速(彻底摆脱全量数据依赖)
你的当前实现里,self.package存储了所有原始数据,但实际上绝大多数统计量都可以通过增量维护中间值来计算,完全不需要每次遍历全量数据:
1. 几何标准差的增量计算
几何标准差的本质是 gstd(x) = exp(std(log(x))),所以我们只需要维护三个变量:
log_sum:所有数据的对数之和log_sq_sum:所有数据的对数平方之和- 已有的
package_count(数据总数)
每次收到数据时,计算log_val = math.log(float(response[1])),然后更新log_sum += log_val、log_sq_sum += log_val**2。计算几何标准差时直接用公式推导:
import math log_mean = self.log_sum / self.package_count log_var = (self.log_sq_sum / self.package_count) - (log_mean ** 2) log_std = math.sqrt(log_var) gstd_val = math.exp(log_std)
这一步能把几何标准差的计算耗时从O(n)降到O(1)。
2. 众数的增量统计
从你的server代码看,数据范围是固定的(1-10),我们可以用一个固定大小的计数器数组实时统计每个值的出现次数:
- 初始化
self.value_counts = [0] * 10(对应1-10的数值) - 每次收到数据
val = float(response[1]),执行self.value_counts[int(val)-1] += 1(数值转数组索引) - 求众数时,直接找
self.value_counts中最大值对应的索引+1即可,完全不用调用stats.mode
如果数据范围不固定,也可以用collections.defaultdict(int)来计数,同样是增量维护,比每次遍历全量数据快得多。
3. 数据清理:移除不必要的原始数据存储
如果业务不需要保留所有原始数据,直接删除self.package列表,只维护上述中间统计量。这样内存占用会从O(n)降到O(1)(固定内存),彻底解决内存持续升高的问题。
二、内存优化细节(如果必须保留原始数据)
如果确实需要保留部分原始数据用于回溯分析,可以做这些优化:
- 用
numpy.ndarray替代Python列表:numpy数组的内存效率远高于Python列表,2000万条float数据用numpy存储只需要约160MB,而Python列表每个元素还要额外存储指针和对象头,内存占用会翻倍以上。 - 滑动窗口存储:设置一个固定大小的窗口(比如只保留最近100万条数据),当数据量超过窗口时,移除最早的部分数据,同时同步更新中间统计量(比如重新计算
log_sum、log_sq_sum)。
三、使用Numba加速计算(针对必须全量计算的场景)
如果因为某些原因必须保留全量原始数据,Numba可以有效加速这些CPU密集型操作,具体实现如下:
1. 安装Numba
pip install numba
2. 用Numba加速自定义统计函数
实现自定义的几何标准差和众数计算函数,并用@numba.jit装饰(nopython=True会编译成机器码,最大化性能):
import numba import numpy as np @numba.jit(nopython=True) def numba_gstd(arr): log_arr = np.log(arr) mean_log = np.mean(log_arr) std_log = np.std(log_arr) return np.exp(std_log) @numba.jit(nopython=True) def numba_mode(arr): counts = np.bincount(arr.astype(np.int64)) return np.argmax(counts)
然后在analyze方法中,把self.package转换成numpy数组后调用这些函数:
def analyze(self): while True: _ = input('按下回车键以输出统计信息\n') start = time.time() arr = np.array(self.package) print( f'数据包数量: {self.package_count}\n' f'丢失数据包数量: {self.lose_package_count}\n' f'算术平均值: {self.package_sum / self.package_count:.4f}\n' f'几何标准差: {numba_gstd(arr):.4f}\n' f'众数: {numba_mode(arr)}\n' f'数据处理耗时: {time.time() - start:.4f}秒\n' )
注意:第一次调用Numba装饰的函数会有编译开销,后续调用速度会比scipy.stats的实现快数倍。
四、优化后的客户端代码示例
这里给出一个整合了增量统计和内存优化的client.py版本,彻底解决性能和内存问题:
import socket import struct import threading import time import math from collections import defaultdict class Client: def __init__(self, mcast_grp, mcast_port, data_min=1, data_max=10): self.package_count = 0 self.lose_package_count = 0 # 原代码初始化1可能有误,修正为0 self.package_sum = 0.0 # 几何标准差增量统计变量 self.log_sum = 0.0 self.log_sq_sum = 0.0 # 众数统计计数器(固定范围用数组更高效) self.value_counts = [0] * (data_max - data_min + 1) self.socket = self.create_socket(mcast_grp, mcast_port) self.prev_number = 1 def create_socket(self, mcast_grp, mcast_port): """ 创建用于统计数据收集的套接字 """ sock = socket.socket( socket.AF_INET, socket.SOCK_DGRAM, socket.IPPROTO_UDP ) sock.setsockopt( socket.SOL_SOCKET, socket.SO_REUSEADDR, 1 ) sock.bind(('', mcast_port)) mreq = struct.pack( "4sl", socket.inet_aton(mcast_grp), socket.INADDR_ANY ) sock.setsockopt(socket.IPPROTO_IP, socket.IP_ADD_MEMBERSHIP, mreq) return sock def accept(self): """ 接收多播数据包并统计基础数据 """ while True: try: response = self.socket.recv(10240).decode().split(':') pkg_num = int(response[0]) val = float(response[1]) # 更新丢包计数 self.lose_package_count += (pkg_num - self.prev_number) - 1 self.prev_number = pkg_num # 更新基础统计量 self.package_count += 1 self.package_sum += val # 更新几何标准差中间值 log_val = math.log(val) self.log_sum += log_val self.log_sq_sum += log_val ** 2 # 更新众数计数器 self.value_counts[int(val)-1] += 1 except Exception: continue def analyze(self): """ 循环接收输入并输出统计结果 """ while True: _ = input('按下回车键以输出统计信息\n') start = time.time() # 计算算术平均值 mean_val = self.package_sum / self.package_count if self.package_count > 0 else 0 # 计算几何标准差 if self.package_count > 0: log_mean = self.log_sum / self.package_count log_var = (self.log_sq_sum / self.package_count) - (log_mean ** 2) log_std = math.sqrt(log_var) gstd_val = math.exp(log_std) else: gstd_val = 0 # 计算众数 mode_val = self.value_counts.index(max(self.value_counts)) + 1 if any(self.value_counts) else 0 # 输出结果 print( f'数据包数量: {self.package_count}\n' f'丢失数据包数量: {self.lose_package_count}\n' f'算术平均值: {mean_val:.4f}\n' f'几何标准差: {gstd_val:.4f}\n' f'众数: {mode_val}\n' f'数据处理耗时: {time.time() - start:.4f}秒\n' ) if __name__ == '__main__': client = Client('224.1.1.1', 5004, data_min=1, data_max=10) threading.Thread(target=client.accept, daemon=True).start() client.analyze()
这个版本彻底移除了self.package列表,所有统计量都是增量维护,内存占用固定,analyze方法的耗时几乎可以忽略不计。
内容的提问来源于stack exchange,提问作者Maksim Rumyantsev

