使用np.array_equal()嵌套循环聚类股票的代码错误排查
股票特征聚类错误排查与修正
常见错误原因
- 未标记已聚类股票:嵌套循环中没跳过已归入聚类的股票,导致同一组股票被多次创建新聚类。比如处理完stock1后,处理stock3时没检查它是否已在聚类里,又新建一个聚类放stock3,最终聚类数翻倍。
- 误包含stock_id列:提取特征时没排除stock_id列,每个股票的stock_id唯一,
np.array_equal会直接判定特征不同,每个股票单独成聚类。 - 时序索引不一致:两个股票的特征DataFrame索引(比如日期)不同,哪怕特征值完全一致,
np.array_equal也会返回False,因为数组形状/索引不匹配。 - 浮点数精度问题:如果特征是浮点型,微小计算误差(比如0.1和0.1000000001)会让
np.array_equal判定不相等,不过你说特征完全一致,这个可能性较低,但可以用np.allclose替代排查。
修正后的代码示例
假设df2结构是stock_id列 + 特征列,按stock_id分组后处理:
import pandas as pd import numpy as np # 示例数据 data = { 'stock_id': [1,1,3,3,4,4,5,5], 'feature1': [10,20,10,20,30,40,30,40], 'feature2': [5,6,5,6,7,8,7,8] } df2 = pd.DataFrame(data) # 按stock_id分组,提取特征数组(排除stock_id列) stock_groups = {sid: group.drop('stock_id', axis=1).values for sid, group in df2.groupby('stock_id')} clusters = [] processed_stocks = set() for sid1, feat1 in stock_groups.items(): if sid1 in processed_stocks: continue # 新建聚类并加入当前股票 cluster = [sid1] processed_stocks.add(sid1) for sid2, feat2 in stock_groups.items(): if sid2 == sid1 or sid2 in processed_stocks: continue # 检查特征是否完全一致 if np.array_equal(feat1, feat2): cluster.append(sid2) processed_stocks.add(sid2) clusters.append(cluster) print(clusters) # 输出:[[1, 3], [4, 5]]
关键修正点
- 用
processed_stocks集合记录已处理的股票,避免重复创建聚类。 - 分组时明确排除
stock_id列,只提取特征数据。 - 先检查当前股票是否已处理,跳过已处理项,减少无效循环。
如果时序数据存在索引不匹配问题,可以先重置索引再提取特征数组:
group.drop('stock_id', axis=1).reset_index(drop=True).values
内容的提问来源于stack exchange,提问作者Jose_Peeterson
相关产品推荐
相关产品推荐

