Numba nopython模式实现聚类算法时出现TypingError报错如何解决
Numba nopython模式下类型统一错误修复
问题背景
使用Numba的@jit(nopython=True)装饰器复现论文算法时触发TypingError,涉及函数为max_ssdbscan,参数定义如下:
mst:float类型numpy二维数组,存储距离矩阵labels:一维numpy数组,存储样本标签num_normal_cluster:整数类型,代表正常簇的数量
问题代码
from numba import jit from heapq import heappush, heappop import numpy as np @jit(nopython=True) def max_ssdbscan(mst,labels,num_normal_cluster): normal_indices=np.where(labels>1)[0] cluster_container=[] max_reachability=[] for i in range(num_normal_cluster): cluster_container.append([]) max_reachability.append(1) max_reachability.append(1) max_reachability.append(1) reachability_matrix=np.ones((len(labels),num_normal_cluster+2))*np.inf for i in normal_indices: visited=set() label=labels[i] reach=np.zeros(len(labels)) h=[(5.9,1.0,5)] _=h.pop() heappush(h,(0.0,0.0,int(i))) reached_different=0 cluster_indices=[1] _=cluster_indices.pop() cluster_reach=[0.0] _=cluster_reach.pop() while(len(visited)<len(labels)): distance,pre_dist,current_index=heappop(h) max_dist=max(distance,pre_dist) visited.add(current_index) reach[current_index]=max_dist if((labels[current_index] ==0 or labels[current_index]==label) and reached_different==0): cluster_indices.append(current_index) max_reach=max_dist cluster_reach.append(pre_dist) elif(reached_different==0 and labels[current_index]!=label): reached_different=1 cluster_reach=np.array(cluster_reach) max_i=np.argmax(cluster_reach) max_i+=1 max_i=int(max_i) cluster_indices=cluster_indices[0:max_i] else: reached_different=1 next_indices=np.where(mst[current_index]>0)[0] for next_ind in next_indices: if(next_ind not in visited): heappush(h,(mst[current_index,next_ind],max_dist,next_ind)) cluster_container[label-2]+=cluster_indices cluster_container[label-2]=list(set(cluster_container[label-2])) prev_reach=reachability_matrix[:,label].reshape(1,-1) reach=reach.reshape(1,-1) combined=np.concatenate((prev_reach,reach),axis=0) reachability_matrix[:,label]=np.amin(combined,axis=0) max_reachability[label]=max(max_reachability[label],max_reach) return reachability_matrix,cluster_container,max_reachability
报错信息
TypingError: Cannot unify list(float64)<iv=None> and array(float64, 1d, C) for 'cluster_reach.3', defined at /home/jiahao/Desktop/cluster_with_outlier/fast_ssdbscan.py (93) File "fast_ssdbscan.py", line 93: def max_ssdbscan(mst,labels,num_normal_cluster): <source elided> reach[current_index]=max_dist #print(2) ^ During: typing of assignment at /home/jiahao/Desktop/cluster_with_outlier/fast_ssdbscan.py (93)
错误原因
Numba的nopython模式要求变量在整个生命周期内类型固定,不允许中途切换类型。代码中cluster_reach初始是存储float64的Python列表,但在分支逻辑中执行cluster_reach=np.array(cluster_reach)将其直接转为numpy一维数组,Numba无法统一同一个变量的两种不同类型,因此抛出类型错误。
另外代码中存在多处冗余初始化逻辑(先给空结构塞无关值再pop掉),也会干扰Numba的类型推断,增加报错概率。
修复方法
- 不要在修改变量时切换类型,需要用numpy方法计算时做临时类型转换即可,不要把转换结果赋值回原变量
- 删除冗余的初始化逻辑,直接定义空结构减少类型推断干扰
- 列表切片后显式转回list,保持类型一致
对应修改的核心代码段如下:
for i in normal_indices: visited=set() label=labels[i] reach=np.zeros(len(labels)) # 删除冗余堆初始化 h = [] heappush(h,(0.0,0.0,int(i))) reached_different=0 # 删除冗余列表初始化 cluster_indices=[] cluster_reach=[] while(len(visited)<len(labels)): distance,pre_dist,current_index=heappop(h) max_dist=max(distance,pre_dist) visited.add(current_index) reach[current_index]=max_dist if((labels[current_index] ==0 or labels[current_index]==label) and reached_different==0): cluster_indices.append(current_index) max_reach=max_dist cluster_reach.append(pre_dist) elif(reached_different==0 and labels[current_index]!=label): reached_different=1 # 临时转数组计算索引,不修改原cluster_reach的列表类型 max_i=np.argmax(np.array(cluster_reach)) max_i+=1 max_i=int(max_i) # 切片后显式转list,保持类型一致 cluster_indices=list(cluster_indices[0:max_i]) else: reached_different=1 next_indices=np.where(mst[current_index]>0)[0] for next_ind in next_indices: if(next_ind not in visited): heappush(h,(mst[current_index,next_ind],max_dist,next_ind)) # 其余逻辑保持不变
额外提示:Numba对Python set的支持有限,如果后续运行仍有类型相关报错,可以把
visited替换为布尔类型numpy数组标记访问状态,兼容性更好。
内容的提问来源于stack exchange,提问作者Peter Deng
相关产品推荐
相关产品推荐

