Numba嵌套列表类型定义问题求助
解决Numba嵌套列表类型识别问题
错误原因
你遇到的类型不兼容问题核心是两点:
permutations初始化为普通Python列表[],但你往里面添加的是Numba typed List,两种类型无法统一;- Numba的typed嵌套列表需要明确声明内部元素类型,不能依赖自动推断。
方案1:使用Numba typed嵌套List
通过预先创建类型模板,让Numba明确嵌套List的类型,同时统一内外层列表为typed List:
from numba.typed import List from numba import njit import numpy as np @njit() def build_permutations(position_length, starting_leverage, ending_leverage, starting_amount=100000): # 指定numpy数组类型,避免类型模糊 current_permutation = np.full(position_length, starting_leverage, dtype=np.int64) # 创建子列表类型模板,告诉Numba内部是int64类型的List sublist_template = List() sublist_template.append(List(current_permutation)) sublist_template.pop() # 初始化外层列表为嵌套List类型 permutations = List.empty_list(sublist_template._type_) sublist = List() while True: total = (np.sum(current_permutation) * 440000) / starting_amount if 10 < total < 100: # 将numpy数组转为typed List并添加 sublist.append(List(current_permutation)) if len(sublist) == 100_000_000: permutations.append(sublist) sublist = List() # 生成下一个排列 i = position_length - 1 while i >= 0 and current_permutation[i] == ending_leverage: current_permutation[i] = starting_leverage i -= 1 if i < 0: break current_permutation[i] += 1 if sublist: permutations.append(sublist) return permutations
方案2:改用元组列表(更简洁)
Numba对不可变元组的类型推断更友好,无需显式声明类型,代码更简洁:
from numba.typed import List from numba import njit import numpy as np @njit() def build_permutations(position_length, starting_leverage, ending_leverage, starting_amount=100000): current_permutation = np.full(position_length, starting_leverage, dtype=np.int64) # 初始化外层和内层为存元组的typed List permutation_type = tuple(current_permutation).__class__ permutations = List.empty_list(permutation_type) sublist = List.empty_list(permutation_type) while True: total = (np.sum(current_permutation) * 440000) / starting_amount if 10 < total < 100: # 将numpy数组转为元组添加 sublist.append(tuple(current_permutation)) if len(sublist) == 100_000_000: permutations.append(sublist) sublist = List.empty_list(permutation_type) # 生成下一个排列 i = position_length - 1 while i >= 0 and current_permutation[i] == ending_leverage: current_permutation[i] = starting_leverage i -= 1 if i < 0: break current_permutation[i] += 1 if sublist: permutations.append(sublist) return permutations
额外提示
- 如果数据规模不大,也可以用普通Python列表存元组,但typed List在Numba JIT函数中的性能更优;
- 确保
starting_leverage和ending_leverage是整数类型,避免numpy数组出现类型不一致问题。
内容的提问来源于stack exchange,提问作者sf8193
相关产品推荐
相关产品推荐

