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

如何将两个非尾递归函数从递归转换为迭代(适配CUDA)

将递归C函数转换为CUDA兼容的迭代实现

CUDA对递归支持有限,且迭代实现通常效率更高,以下针对两个非尾递归函数提供具体的迭代转换方案:


1. get_resulting_node 函数转换

该函数在for循环中发起递归调用,递归返回后需完成距离比较与变量更新,核心思路是用栈保存循环上下文+递归返回后待执行的逻辑。

迭代实现思路

  • 定义栈帧结构,存储每次递归调用前的循环状态、参数及中间变量;
  • 通过栈模拟递归流程:首次处理时完成前置计算,若需递归则压入当前上下文(标记为待后续处理),再压入新的递归参数;递归返回后,取出上下文完成后续的距离判断与变量更新。

迭代版代码

#include <stdlib.h>

// 栈帧结构:保存递归调用的上下文及处理阶段
typedef struct {
    // 函数参数
    const uint64_t* np;
    const BiomeTree* bt;
    int idx;
    int alt;
    uint64_t ds;
    int depth;
    // 循环状态
    uint32_t i;
    uint32_t n;
    uint16_t inner;
    uint32_t step;
    uint64_t ds_inner;
    // 处理阶段:0=首次进入,1=循环处理,2=等待递归返回
    int phase;
    // 中间变量
    int leaf;
    // 递归返回值
    int leaf2;
} NodeStackFrame;

int get_resulting_node_iter(const uint64_t np[6], const BiomeTree *bt, int idx,
    int alt, uint64_t ds, int depth)
{
    // CUDA中建议使用固定大小栈(避免动态分配开销),此处假设最大深度足够
    NodeStackFrame stack[64];
    int stack_ptr = -1;
    int final_leaf = alt;

    // 初始压入根调用上下文
    stack_ptr++;
    NodeStackFrame* frame = &stack[stack_ptr];
    frame->np = np;
    frame->bt = bt;
    frame->idx = idx;
    frame->alt = alt;
    frame->ds = ds;
    frame->depth = depth;
    frame->phase = 0;

    while (stack_ptr >= 0) {
        frame = &stack[stack_ptr];
        const BiomeTree* curr_bt = frame->bt;

        if (frame->phase == 0) {
            // 首次处理:执行原函数前置逻辑
            if (curr_bt->steps[frame->depth] == 0) {
                final_leaf = frame->idx;
                stack_ptr--;
                continue;
            }

            uint32_t step;
            int curr_depth = frame->depth;
            do {
                step = curr_bt->steps[curr_depth];
                curr_depth++;
            } while (frame->idx + step >= curr_bt->len);

            uint64_t node = curr_bt->nodes[frame->idx];
            uint16_t inner = node >> 48;

            // 保存上下文状态,进入循环处理阶段
            frame->step = step;
            frame->inner = inner;
            frame->leaf = frame->alt;
            frame->i = 0;
            frame->n = curr_bt->order;
            frame->depth = curr_depth;
            frame->phase = 1;
            continue;
        } else if (frame->phase == 1) {
            // 处理for循环迭代
            if (frame->i >= frame->n) {
                final_leaf = frame->leaf;
                stack_ptr--;
                continue;
            }

            uint64_t ds_inner = get_np_dist(frame->np, curr_bt, frame->inner);
            if (ds_inner < frame->ds) {
                // 需要递归,保存当前循环状态,压入递归调用
                frame->ds_inner = ds_inner;
                frame->phase = 2;

                stack_ptr++;
                NodeStackFrame* rec_frame = &stack[stack_ptr];
                rec_frame->np = frame->np;
                rec_frame->bt = curr_bt;
                rec_frame->idx = frame->inner;
                rec_frame->alt = frame->leaf;
                rec_frame->ds = frame->ds;
                rec_frame->depth = frame->depth;
                rec_frame->phase = 0;
                continue;
            } else {
                // 无需递归,进入下一次循环
                frame->inner += frame->step;
                if (frame->inner >= curr_bt->len) {
                    final_leaf = frame->leaf;
                    stack_ptr--;
                    continue;
                }
                frame->i++;
                continue;
            }
        } else if (frame->phase == 2) {
            // 递归返回,处理后续逻辑
            int leaf2 = final_leaf;
            uint64_t ds_leaf2 = (frame->inner == leaf2) ? frame->ds_inner : get_np_dist(frame->np, curr_bt, leaf2);

            if (ds_leaf2 < frame->ds) {
                frame->ds = ds_leaf2;
                frame->leaf = leaf2;
            }

            // 进入下一次循环
            frame->inner += frame->step;
            if (frame->inner >= curr_bt->len) {
                final_leaf = frame->leaf;
                stack_ptr--;
                continue;
            }
            frame->i++;
            frame->phase = 1;
        }
    }

    return final_leaf;
}

2. getSpline 函数转换

该函数存在多分支递归(单递归调用、双递归调用),递归返回后需完成插值计算,核心思路是用栈保存计算阶段+中间变量,区分不同的处理步骤。

迭代实现思路

  • 定义栈帧结构,存储当前Spline指针、输入值、中间计算变量及处理阶段(初始、等待子递归结果、最终计算);
  • 按阶段处理栈帧:初始阶段完成参数检查与分支判断;需要递归时压入当前上下文,再压入子Spline的初始帧;递归返回后更新中间变量,进入下一阶段直到完成最终计算。

迭代版代码

#include <stdlib.h>
#include <stdio.h>

// 假设lerp函数已实现:float lerp(float t, float a, float b) { return a + t*(b-a); }

// 处理阶段枚举
typedef enum {
    SPLINE_INIT,
    SPLINE_BOUND_WAIT, // 等待边界分支的子递归结果
    SPLINE_NEED_SP1,   // 等待sp1的递归结果
    SPLINE_NEED_SP2    // 等待sp2的递归结果
} SplinePhase;

// 栈帧结构
typedef struct {
    const Spline* sp;
    const float* vals;
    // 中间变量
    float f;
    int i;
    float g;
    float h;
    float k;
    float l;
    float m;
    float n; // sp1的结果
    float v; // 边界分支的子结果
    SplinePhase phase;
} SplineStackFrame;

float getSpline_iter(const Spline *sp, const float *vals)
{
    // CUDA中使用固定大小栈
    SplineStackFrame stack[64];
    int stack_ptr = -1;
    float final_result = 0.0f;

    // 初始压入根调用
    stack_ptr++;
    SplineStackFrame* frame = &stack[stack_ptr];
    frame->sp = sp;
    frame->vals = vals;
    frame->phase = SPLINE_INIT;

    while (stack_ptr >= 0) {
        frame = &stack[stack_ptr];
        const Spline* curr_sp = frame->sp;
        const float* curr_vals = frame->vals;

        switch (frame->phase) {
            case SPLINE_INIT:
                if (!curr_sp || curr_sp->len <= 0 || curr_sp->len >= 12) {
                    printf("getSpline(): bad parameters\n");
                    exit(1);
                }

                if (curr_sp->len == 1) {
                    final_result = ((FixSpline*)curr_sp)->val;
                    stack_ptr--;
                    continue;
                }

                float f = curr_vals[curr_sp->typ];
                int i;
                for (i = 0; i < curr_sp->len; i++)
                    if (curr_sp->loc[i] >= f)
                        break;

                if (i == 0 || i == curr_sp->len) {
                    if (i) i--;
                    // 边界分支:压入子Spline调用
                    frame->f = f;
                    frame->i = i;
                    frame->phase = SPLINE_BOUND_WAIT;

                    stack_ptr++;
                    SplineStackFrame* child_frame = &stack[stack_ptr];
                    child_frame->sp = curr_sp->val[i];
                    child_frame->vals = curr_vals;
                    child_frame->phase = SPLINE_INIT;
                    continue;
                } else {
                    // 中间分支:先调用sp1
                    frame->f = f;
                    frame->i = i;
                    frame->g = curr_sp->loc[i-1];
                    frame->h = curr_sp->loc[i];
                    frame->k = (f - frame->g) / (frame->h - frame->g);
                    frame->l = curr_sp->der[i-1];
                    frame->m = curr_sp->der[i];
                    frame->phase = SPLINE_NEED_SP1;

                    stack_ptr++;
                    SplineStackFrame* sp1_frame = &stack[stack_ptr];
                    sp1_frame->sp = curr_sp->val[i-1];
                    sp1_frame->vals = curr_vals;
                    sp1_frame->phase = SPLINE_INIT;
                    continue;
                }
                break;

            case SPLINE_BOUND_WAIT:
                // 计算边界分支结果
                frame->v = final_result;
                final_result = frame->v + curr_sp->der[frame->i] * (frame->f - curr_sp->loc[frame->i]);
                stack_ptr--;
                break;

            case SPLINE_NEED_SP1:
                // 保存sp1结果,调用sp2
                frame->n = final_result;
                frame->phase = SPLINE_NEED_SP2;

                stack_ptr++;
                SplineStackFrame* sp2_frame = &stack[stack_ptr];
                sp2_frame->sp = curr_sp->val[frame->i];
                sp2_frame->vals = curr_vals;
                sp2_frame->phase = SPLINE_INIT;
                break;

            case SPLINE_NEED_SP2:
                // 计算最终插值结果
                float o = final_result;
                float k = frame->k;
                float l = frame->l;
                float m = frame->m;
                float n = frame->n;
                float h_g = frame->h - frame->g;
                float p = l * h_g - (o - n);
                float q = -m * h_g + (o - n);
                final_result = lerp(k, n, o) + k * (1.0F - k) * lerp(k, p, q);
                stack_ptr--;
                break;

            default:
                stack_ptr--;
                break;
        }
    }

    return final_result;
}

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 14:05:53