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

使用Numba CUDA核加速光线追踪时自定义Ray类识别失败问题

解决Numba CUDA无法识别自定义Ray类的问题

Numba CUDA不支持原生Python类,必须使用@jitclass装饰器显式定义类的属性类型,才能让Numba识别并生成CUDA兼容的代码。以下是具体修复步骤和修改后的代码示例:

1. 定义Numba兼容的基础类型

首先确保你使用的vec3类型是Numba可识别的,先定义一个vec3的jitclass:

from numba import jitclass, float32, int32, boolean
import numba.cuda as cuda

# 定义vec3的jitclass
vec3_spec = [
    ('x', float32),
    ('y', float32),
    ('z', float32)
]

@jitclass(vec3_spec)
class vec3:
    def __init__(self, x, y, z):
        self.x = x
        self.y = y
        self.z = z
    
    def extract(self, hit_check):
        return vec3(self.x if hit_check else 0.0, 
                   self.y if hit_check else 0.0, 
                   self.z if hit_check else 0.0)
    
    def place(self, hit_check):
        return self if hit_check else vec3(0.0, 0.0, 0.0)

2. 将Ray类改为jitclass

修改你的Ray类,添加@jitclass装饰器并指定属性类型:

ray_spec = [
    ('origin', vec3.class_type.instance_type),
    ('dir', vec3.class_type.instance_type),
    ('depth', int32),
    ('n', vec3.class_type.instance_type),
    ('reflections', int32),
    ('transmissions', int32),
    ('diffuse_reflections', int32)
]

@jitclass(ray_spec)
class Ray:
    def __init__(self, origin, dir, depth, n, reflections, transmissions, diffuse_reflections):
        self.origin = origin   
        self.dir = dir
        self.depth = depth     
        self.n = n
        self.reflections = reflections
        self.transmissions = transmissions
        self.diffuse_reflections = diffuse_reflections
    
    def extract(self, hit_check):
        return Ray(self.origin.extract(hit_check), 
                   self.dir.extract(hit_check), 
                   self.depth,  
                   self.n.extract(hit_check), 
                   self.reflections, 
                   self.transmissions,
                   self.diffuse_reflections)

3. 处理Hit类(若在CUDA核中使用)

同样将Hit类转为jitclass,属性类型需匹配Numba支持的类型:

# 若Material、Collider等需在核中使用,需同步转为jitclass或用索引替代对象
hit_spec = [
    ('distance', float32),
    ('orientation', float32),
    ('material', int32),
    ('collider', int32),
    ('surface', int32),
    ('u', float32),
    ('v', float32),
    ('N', vec3.class_type.instance_type),
    ('point', vec3.class_type.instance_type)
]

@jitclass(hit_spec)
class Hit:
    def __init__(self, distance, orientation, material, collider, surface):
        self.distance = distance
        self.orientation = orientation
        self.material = material
        self.collider = collider
        self.surface = surface
        self.u = 0.0
        self.v = 0.0
        self.N = vec3(0.0, 0.0, 0.0)
        self.point = vec3(0.0, 0.0, 0.0)
    
    def get_uv(self):
        if self.u == 0.0 and self.v == 0.0:
            # 替换为Numba兼容的UV计算逻辑
            self.u, self.v = 0.0, 0.0
        return self.u, self.v
    
    def get_normal(self):
        if self.N.x == 0.0 and self.N.y == 0.0 and self.N.z == 0.0:
            # 替换为Numba兼容的法线计算逻辑
            self.N = vec3(0.0, 1.0, 0.0)
        return self.N

4. 重写CUDA核函数

将get_raycolor改为CUDA核,避免使用Numba CUDA不支持的Python特性(如zip、列表推导式):

@cuda.jit
def cuda_get_raycolor(ray, scene_data, result):
    distances = cuda.local.array(shape=(scene_data.collider_count,), dtype=float32)
    hit_orientation = cuda.local.array(shape=(scene_data.collider_count,), dtype=float32)
    
    # 显式循环计算每个碰撞体的相交结果
    for i in range(scene_data.collider_count):
        coll = scene_data.colliders[i]
        d, o = coll.intersect(ray.origin, ray.dir)
        distances[i] = d
        hit_orientation[i] = o
    
    # 计算最近距离
    nearest = float('inf')
    for d in distances:
        if d < nearest:
            nearest = d
    
    color = vec3(0.0, 0.0, 0.0)
    FARAWAY = 1e10
    
    for i in range(scene_data.collider_count):
        coll = scene_data.colliders[i]
        dis = distances[i]
        orient = hit_orientation[i]
        
        hit_check = (nearest != FARAWAY) & (dis == nearest)
        if hit_check:
            extracted_ray = ray.extract(hit_check)
            hit = Hit(dis, orient, coll.material_idx, coll.idx, coll.surface_idx)
            coll_color = coll.get_color(scene_data, extracted_ray, hit).place(hit_check)
            color.x += coll_color.x
            color.y += coll_color.y
            color.z += coll_color.z
    
    result[0] = color.x
    result[1] = color.y
    result[2] = color.z

关键注意事项

  • 所有在CUDA核中使用的类都必须是jitclass,属性类型只能是Numba支持的基本类型或其他jitclass实例。
  • 场景数据(scene)需转换为Numba可识别的结构,比如将碰撞体列表转为jitclass数组,而非普通Python列表。
  • 核函数参数需使用CUDA数组或jitclass实例,不能直接传递普通Python对象。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 23:48:11