You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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的类型推断,增加报错概率。

修复方法

  1. 不要在修改变量时切换类型,需要用numpy方法计算时做临时类型转换即可,不要把转换结果赋值回原变量
  2. 删除冗余的初始化逻辑,直接定义空结构减少类型推断干扰
  3. 列表切片后显式转回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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.29 22:24:24