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

Numba中实现uint64乘法返回高低位/int128遇报错求解决方案

解决方案:在Numba中实现64位无符号整数乘法的高位/低位拆分,或返回128位整数

一、实现返回(uint64 higher, uint64 lower)的函数

你的原代码错误在于试图通过分配内存并返回指针构造Tuple,但Numba的intrinsic对于Tuple返回值,需要直接生成两个独立的IR值并打包成Tuple,而非返回内存指针。以下是修正后的代码:

import numpy as np
from llvmlite import ir
from numba import njit, types
from numba.extending import intrinsic

@intrinsic
def mul_uint64_split(ctx, a, b):
    # 定义函数签名:输入两个uint64,返回(uint64, uint64)的Tuple
    sig = types.Tuple((types.uint64, types.uint64))(types.uint64, types.uint64)
    
    def codegen(ctx, builder, sig, args):
        u64 = ir.IntType(64)
        u128 = ir.IntType(128)
        
        # 将两个uint64零扩展为uint128后相乘
        a_ext = builder.zext(args[0], u128)
        b_ext = builder.zext(args[1], u128)
        product = builder.mul(a_ext, b_ext)
        
        # 提取高位和低位:高位是右移64位后截断为uint64,低位直接截断
        high = builder.trunc(builder.lshr(product, u128(64)), u64)
        low = builder.trunc(product, u64)
        
        # 打包成Tuple返回(Numba会自动处理Tuple的IR构造)
        return ctx.make_tuple(builder, sig.return_type, (high, low))
    
    return sig, codegen

@njit
def mul(x, y):
    return mul_uint64_split(x, y)

# 测试
a = 2**63 - 1
high, low = mul(a, a)
print(f"高位: {high}, 低位: {low}")
# 验证:(2^63-1)^2 = 2^126 - 2^64 + 1,高位是2^62 -1,低位是2^64 -1
assert high == (2**62 - 1)
assert low == (2**64 - 1)

关键修正点:

  • 直接使用ctx.make_tuple构造返回的Tuple,无需手动分配内存
  • 移除错误的GEP操作,改为直接计算高位和低位的IR值
  • 统一IR操作的常量类型:所有常量与对应IR类型匹配(比如右移用u128(64)而非u8(64))

二、直接返回128位无符号整数并转换为Python可用格式

Numba原生支持types.uint128类型,你可以直接返回该类型,然后在Python环境中通过np.view拆分为两个uint64:

import numpy as np
from numba import njit, types
from numba.extending import intrinsic
from llvmlite import ir

@intrinsic
def mul_uint64_to_uint128(ctx, a, b):
    # 定义函数签名:输入两个uint64,返回uint128
    sig = types.uint128(types.uint64, types.uint64)
    
    def codegen(ctx, builder, sig, args):
        u128 = ir.IntType(128)
        a_ext = builder.zext(args[0], u128)
        b_ext = builder.zext(args[1], u128)
        return builder.mul(a_ext, b_ext)
    
    return sig, codegen

@njit
def mul_to_uint128(x, y):
    return mul_uint64_to_uint128(x, y)

# 测试并转换为两个uint64
a = 2**63 - 1
product_128 = mul_to_uint128(a, a)
# 转换为numpy的uint128数组,再view为两个uint64
arr = np.array([product_128], dtype=np.uint128)
high, low = arr.view(np.uint64)
print(f"高位: {high}, 低位: {low}")
assert high == (2**62 - 1)
assert low == (2**64 - 1)

说明:

  • Numba的uint128类型可直接返回至Python环境,作为Python整数对象(Python原生不支持uint128直接运算,转成numpy数组处理更方便)
  • 使用np.view可在不复制内存的情况下,将uint128拆分为两个uint64元素(注意字节序为系统端序,需固定端序可手动调整)

内容的提问来源于stack exchange,提问作者ZisIsNotZis

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 05:03:36