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

如何将穷举式枢轴选择的递归函数转为迭代函数及优化咨询

穷举式枢轴选择的优化需求

我正在开展穷举式枢轴选择的优化工作,背景如下:

  • 数据集 S = {xi | i = 1,2,...,n} 包含n个数据点
  • k个枢轴标记为 P = {pj | j = 1,2,...,k}
  • 每个数据点xi与枢轴集形成向量 Xp = (d(x,p1),...,d(x,pk))
  • 我们通过最大化或最小化所有向量Xp的切比雪夫距离和来筛选最优枢轴集

现有代码通过递归函数Combination生成所有枢轴组合,再调用SumDistance计算距离和。我需要将该递归函数改为迭代实现,同时寻求其他可行的代码优化方案。


原代码实现

#include<stdio.h>
#include<stdlib.h>
#include<math.h>
#include<sys/time.h>

// Calculate sum of distance while combining different pivots. Complexity : O( n^2 )
double SumDistance(const int k, const int n, const int dim, double* coord, int* pivots){
    double* rebuiltCoord = (double*)malloc(sizeof(double) * n * k);
    int i;
    for(i=0; i<n*k; i++){
        rebuiltCoord[i] = 0;
    }

    // Rebuild coordinates. New coordinate of one point is its distance to each pivot.
    for(i=0; i<n; i++){
        int ki;
        for(ki=0; ki<k; ki++){
            double distance = 0;
            int pivoti = pivots[ki];
            int j;
            for(j=0; j<dim; j++){
                distance += pow(coord[pivoti*dim + j] - coord[i*dim + j] ,2);
            }
            rebuiltCoord[i*k + ki] = sqrt(distance);
        }
    }

    // Calculate the sum of Chebyshev distance with rebuilt coordinates between every points
    double chebyshevSum = 0;
    for(i=0; i<n; i++){
        int j;
        for(j=0; j<n; j++){
            double chebyshev = 0;
            int ki;
            for(ki=0; ki<k; ki++){
                double dis = fabs(rebuiltCoord[i*k + ki] - rebuiltCoord[j*k + ki]);
                chebyshev = dis>chebyshev ? dis : chebyshev;
            }
            chebyshevSum += chebyshev;
        }
    }

    free(rebuiltCoord);

    return chebyshevSum;
}

// Recursive function Combination() : combine pivots and calculate the sum of distance while combining different pivots.
// ki  : current depth of the recursion
// k   : number of pivots
// n   : number of points
// dim : dimension of metric space
// M   : number of combinations to store
// coord  : coordinates of points
// pivots : indexes of pivots
// maxDistanceSum  : the largest M distance sum
// maxDisSumPivots : the top M pivots combinations
// minDistanceSum  : the smallest M distance sum
// minDisSumPivots : the bottom M pivots combinations
void Combination(int ki, const int k, const int n, const int dim, const int M, double* coord, int* pivots,
                 double* maxDistanceSum, int* maxDisSumPivots, double* minDistanceSum, int* minDisSumPivots){
    if(ki==k-1){
        int i;
        for(i=pivots[ki-1]+1; i<n; i++){
            pivots[ki] = i;

            // Calculate sum of distance while combining different pivots.
            double distanceSum = SumDistance(k, n, dim, coord, pivots);

            // put data at the end of array
            maxDistanceSum[M] = distanceSum;
            minDistanceSum[M] = distanceSum;
            int kj;
            for(kj=0; kj<k; kj++){
                maxDisSumPivots[M*k + kj] = pivots[kj];
            }
            for(kj=0; kj<k; kj++){
                minDisSumPivots[M*k + kj] = pivots[kj];
            }
            // sort
            int a;
            for(a=M; a>0; a--){
                if(maxDistanceSum[a] > maxDistanceSum[a-1]){
                    double temp = maxDistanceSum[a];
                    maxDistanceSum[a] = maxDistanceSum[a-1];
                    maxDistanceSum[a-1] = temp;
                    int kj;
                    for(kj=0; kj<k; kj++){
                        int temp = maxDisSumPivots[a*k + kj];
                        maxDisSumPivots[a*k + kj] = maxDisSumPivots[(a-1)*k + kj];
                        maxDisSumPivots[(a-1)*k + kj] = temp;
                    }
                }
            }
            for(a=M; a>0; a--){
                if(minDistanceSum[a] < minDistanceSum[a-1]){
                    double temp = minDistanceSum[a];
                    minDistanceSum[a] = minDistanceSum[a-1];
                    minDistanceSum[a-1] = temp;
                    int kj;
                    for(kj=0; kj<k; kj++){
                        int temp = minDisSumPivots[a*k + kj];
                        minDisSumPivots[a*k + kj] = minDisSumPivots[(a-1)*k + kj];
                        minDisSumPivots[(a-1)*k + kj] = temp;
                    }
                }
            }
        }
        return;
    }

    // Recursively call Combination() to combine pivots
    int i;
    for(i=pivots[ki-1]+1; i<n; i++) {
        pivots[ki] = i;
        Combination(ki+1, k, n, dim, M, coord, pivots, maxDistanceSum, maxDisSumPivots, minDistanceSum, minDisSumPivots);

        /** Iteration Log : pivots computed, best pivots, max distance sum, min distance sum pivots, min distance sum
        *** You can delete the logging code. **/
        if(ki==k-2){
            int kj;
            for(kj=0; kj<k; kj++){
                printf("%d ", pivots[kj]);
            }
            putchar('\t');
            for(kj=0; kj<k; kj++){
                printf("%d ", maxDisSumPivots[kj]);
            }
            printf("%lf\t", maxDistanceSum[0]);
            for(kj=0; kj<k; kj++){
                printf("%d ", minDisSumPivots[kj]);
            }
            printf("%lf\n", minDistanceSum[0]);
        }
    }
}

int main(int argc, char* argv[]){
    // filename : input file namespace
    char* filename = (char*)"uniformvector-2dim-5h.txt";
    if( argc==2 ) {
        filename = argv[1];
    }  else if(argc != 1) {
        printf("Usage: ./pivot <filename>\n");
        return -1;
    }
    // M : number of combinations to store
    const int M = 1000;
    // dim : dimension of metric space
    int dim;
    // n : number of points
    int n;
    // k : number of pivots
    int k;

    // Read parameter
    FILE* file = fopen(filename, "r");
    if( file == NULL ) {
        printf("%s file not found.\n", filename);
        return -1;
    }
    fscanf(file, "%d", &dim);
    fscanf(file, "%d", &n);
    fscanf(file, "%d", &k);
    printf("dim = %d, n = %d, k = %d\n", dim, n, k);

    // Start timing
    struct timeval start;

    // Read Data
    double* coord = (double*)malloc(sizeof(double) * dim * n);
    int i;
    for(i=0; i<n; i++){
        int j;
        for(j=0; j<dim; j++){
            fscanf(file, "%lf", &coord[i*dim + j]);
        }
    }
    fclose(file);
    gettimeofday(&start, NULL);

    // maxDistanceSum : the largest M distance sum
    double* maxDistanceSum = (double*)malloc(sizeof(double) * (M+1));
    for(i=0; i<M; i++){
        maxDistanceSum[i] = 0;
    }
    // maxDisSumPivots : the top M pivots combinations
    int* maxDisSumPivots = (int*)malloc(sizeof(int) * k * (M+1));
    for(i=0; i<M; i++){
        int ki;
        for(ki=0; ki<k; ki++){
            maxDisSumPivots[i*k + ki] = 0;
        }
    }
    // minDistanceSum : the smallest M distance sum
    double* minDistanceSum = (double*)malloc(sizeof(double) * (M+1));
    for(i=0; i<M; i++){
        minDistanceSum[i] = __DBL_MAX__;
    }
    // minDisSumPivots : the bottom M pivots combinations
    int* minDisSumPivots = (int*)malloc(sizeof(int) * k * (M+1));
    for(i=0; i<M; i++){
        int ki;
        for(ki=0; ki<k; ki++){
            minDisSumPivots[i*k + ki] = 0;
        }
    }

    // temp : indexes of pivots with dummy array head
    int* temp = (int*)malloc(sizeof(int) * (k+1));
    temp[0] = -1;

    // Main loop. Combine different pivots with recursive function and evaluate them. Complexity : O( n^(k+2) )
    Combination(0, k, n, dim, M, coord, &temp[1], maxDistanceSum, maxDisSumPivots, minDistanceSum, minDisSumPivots);

    // End timing
    struct timeval end;
    gettimeofday (&end, NULL);
    printf("Using time : %f ms\n", (end.tv_sec-start.tv_sec)*1000.0+(end.tv_usec-start.tv_usec)/1000.0);

    // Store the result
    FILE* out = fopen("result.txt", "w");
    for(i=0; i<M; i++){
        int ki;
        for(ki=0; ki<k-1; ki++){
            fprintf(out, "%d ", maxDisSumPivots[i*k + ki]);
        }
        fprintf(out, "%d\n", maxDisSumPivots[i*k + k-1]);
    }
    for(i=0; i<M; i++){
        int ki;
        for(ki=0; ki<k-1; ki++){
            fprintf(out, "%d ", minDisSumPivots[i*k + ki]);
        }
        fprintf(out, "%d\n", minDisSumPivots[i*k + k-1]);
    }
    fclose(out);

    // Log
    int ki;
    printf("max : ");
    for(ki=0; ki<k; ki++){
        printf("%d ", maxDisSumPivots[ki]);
    }
    printf("%lf\n", maxDistanceSum[0]);
    printf("min : ");
    for(ki=0; ki<k; ki++){
        printf("%d ", minDisSumPivots[ki]);
    }
    printf("%lf\n", minDistanceSum[0]);

    return 0;
}

优化方案

一、递归转迭代实现Combination函数

用数组模拟递归栈,记录当前构建的枢轴位置和起始索引,按递增顺序生成所有不重复的枢轴组合:

void CombinationIterative(const int k, const int n, const int dim, const int M, double* coord, int* pivots,
                          double* maxDistanceSum, int* maxDisSumPivots, double* minDistanceSum, int* minDisSumPivots) {
    // 栈结构:每个元素存储[当前枢轴位置ki, 当前层起始索引]
    int* stack = (int*)malloc(sizeof(int) * k * 2);
    int stackPtr = 0;
    // 初始状态:第0层,起始索引为0(对应原递归的temp[0]=-1)
    stack[stackPtr++] = 0;
    stack[stackPtr++] = 0;

    while (stackPtr > 0) {
        int currentKi = stack[--stackPtr];
        int startI = stack[--stackPtr];

        if (currentKi == k-1) {
            // 处理最后一层,生成所有可能的最终枢轴
            for (int i = startI; i < n; i++) {
                pivots[currentKi] = i;
                double distanceSum = SumDistance(k, n, dim, coord, pivots);

                // 更新最大/最小距离集合
                maxDistanceSum[M] = distanceSum;
                minDistanceSum[M] = distanceSum;
                for (int kj = 0; kj < k; kj++) {
                    maxDisSumPivots[M*k + kj] = pivots[kj];
                    minDisSumPivots[M*k + kj] = pivots[kj];
                }

                // 维护最大距离Top M(提前终止有序部分)
                for (int a = M; a > 0; a--) {
                    if (maxDistanceSum[a] > maxDistanceSum[a-1]) {
                        double tempD = maxDistanceSum[a];
                        maxDistanceSum[a] = maxDistanceSum[a-1];
                        maxDistanceSum[a-1] = tempD;
                        for (int kj = 0; kj < k; kj++) {
                            int tempI = maxDisSumPivots[a*k + kj];
                            maxDisSumPivots[a*k + kj] = maxDisSumPivots[(a-1)*k + kj];
                            maxDisSumPivots[(a-1)*k + kj] = tempI;
                        }
                    } else {
                        break;
                    }
                }

                // 维护最小距离Top M(提前终止有序部分)
                for (int a = M; a > 0; a--) {
                    if (minDistanceSum[a] < minDistanceSum[a-1]) {
                        double tempD = minDistanceSum[a];
                        minDistanceSum[a] = minDistanceSum[a-1];
                        minDistanceSum[a-1] = tempD;
                        for (int kj = 0; kj < k; kj++) {
                            int tempI = minDisSumPivots[a*k + kj];
                            minDisSumPivots[a*k + kj] = minDisSumPivots[(a-1)*k + kj];
                            minDisSumPivots[(a-1)*k + kj] = tempI;
                        }
                    } else {
                        break;
                    }
                }

                // 原日志输出逻辑
                int kj;
                for(kj=0; kj<k; kj++){
                    printf("%d ", pivots[kj]);
                }
                putchar('\t');
                for(kj=0; kj<k; kj++){
                    printf("%d ", maxDisSumPivots[kj]);
                }
                printf("%lf\t", maxDistanceSum[0]);
                for(kj=0; kj<k; kj++){
                    printf("%d ", minDisSumPivots[kj]);
                }
                printf("%lf\n", minDistanceSum[0]);
            }
        } else {
            // 反向入栈保证遍历顺序与递归一致
            for (int i = n-1; i > startI; i--) {
                pivots[currentKi] = i;
                stack[stackPtr++] = currentKi + 1;
                stack[stackPtr++] = i + 1;
            }
            // 处理当前起始索引的情况
            pivots[currentKi] = startI;
            stack[stackPtr++] = currentKi + 1;
            stack[stackPtr++] = startI + 1;
        }
    }
    free(stack);
}

在main函数中替换原递归调用:

// 替换原Combination调用
CombinationIterative(k, n, dim, M, coord, &temp[1], maxDistanceSum, maxDisSumPivots, minDistanceSum, minDisSumPivots);

二、其他性能优化点

  1. 预计算所有点对的欧氏距离
    提前计算并存储所有点对的欧氏距离到二维矩阵,避免每次生成枢轴组合时重复计算:
    // 在main函数中预计算距离矩阵
    double* distMatrix = (double*)malloc(sizeof(double) * n * n);
    for (int i = 0; i <
相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.25 23:33:29