【问题标题】:Improve performance of nested loops MATLAB提高嵌套循环的性能 MATLAB
【发布时间】:2013-10-21 22:26:38
【问题描述】:

我正在尝试分析大量数据,这些数据使我的程序运行缓慢。 我正在将数据集从 .txt 文件读取到元胞数组。 我正在使用一个单元格数组来对我的数据进行分类,它是两个属性的形式,我需要字符类。

我想使用最接近的平均分类器找到重新替换错误。 我有一个主要的外循环,它遍历我的数据集的每一行(数万行)。依次删除每一行,每次迭代一个。每次迭代都会重新计算两个属性的平均值,并删除线。主要的挂点似乎是我要计算数据集中每一行的下一部分:

  • 该行数据(2 个属性值)与 我每个班级的平均值。
  • 然后我想记录其属性平均值最接近的类,这将是其分配的类。
  • 最后我想检查这个分配的类是否正确 类。

目前这个循环是这样的。

errorCount = 0;
for l = 1:20000
    closest = 100;
    class = 0;
    attribute1 = d{2}(l);
    attribute2 = d{3}(l);
    for m = 1:numel(classes)
        dist = sqrt((attribute1-meansattr1(m))*(attribute1-meansattr1(m)) + (attribute2-meansattr2(m))*(attribute2-meansattr2(m)));
        if dist < closest
            closest = dist;
            class = m;
        end
    end

    if strcmp(d{1}(l),classes(class))
        %correct
    else
        errorCount = errorCount + 1;
    end
end

d 是我的单元格数组,其中d{2} 是包含我的属性 1 值的列。对于该列的第一行,我使用 d{1}(1) 访问这些值。

classes 是我数据集中的唯一类,因此对于我的每个类,我都会计算到它的欧几里得距离。

meansattr1meansattr2 是包含我的每个属性的平均值的数组。当删除一行时,这些在外部循环的每次迭代中都会更新。

希望这可以帮助您理解我拥有的代码。非常感谢在优化和加速这些计算方面的任何帮助。

【问题讨论】:

  • 最简单的速度改进是删除sqrt 调用。找到最近距离的平方与最近距离完全相同。

标签: performance matlab vectorization nested-loops


【解决方案1】:

最简单的速度改进是删除sqrt 调用。求最近距离的平方与最近距离完全一样。

接下来,您可以对内部循环进行矢量化。自从我做任何 MatLab 以来已经很久了,所以我可能会弄错下面的代码,但我的想法是把这两个属性变成一个长度为 numel(classes) 的向量。然后,您可以直接计算差异并将它们平方。

类似这样的:

d1 = attribute1 - meansattr1;
d2 = attribute2 - meansattr2;
[closest, class] = min( d1 .* d1 + d2 .* d2 );

顺便说一句,将class 用作变量(如果可以的话)并不是一个好主意。这是一个保留字。

【讨论】:

  • 'closest=strt(closest)' 在使用距离的情况下丢失,变量包含平方距离。
  • 当然,但原始代码也没有显示它的使用位置。 所有循环之后很容易采取sqrt
【解决方案2】:

您实质上是在优化 k-means 算法的迭代部分,因此您可以参考 my previous solution 了解对其进行矢量化的方法。但是,这里是针对您的问题和数据格式的方法。

取一个如下所示的随机数据集,

numClasses = 5;
numPoints = 20e3;
numDims = 2;

classes = strsplit(num2str(1:numClasses));

% generate random data (expected error rate of (numClasses-1)/numClasses)
d{1} = classes(randi(numClasses,numPoints,1));
d{2} = rand(numPoints,1);
d{3} = rand(numPoints,1);

% random initial class centers
meansattr1 = rand(5,1);
meansattr2 = rand(5,1);

您的代码,压缩并存储每个点最近的类 ID 以及到该类的距离变为:

closestDistance = zeros(numPoints,1);  nearestCluster = zeros(numPoints,1);
errorCount = 0;
for l = 1:numPoints
    closest = 100; iclass = 0;
    attribute1 = d{2}(l); attribute2 = d{3}(l);

    for m = 1:numel(classes)
        dist = sqrt((attribute1-meansattr1(m))*(attribute1-meansattr1(m)) + ...
            (attribute2-meansattr2(m))*(attribute2-meansattr2(m)));
        if dist < closest
            closest = dist; closestDistance(l) = closest;
            iclass = m; nearestCluster(l) = iclass;
        end
    end

    if ~strcmp(d{1}(l),classes(iclass))
        errorCount = errorCount + 1;
    end
end

上面的矢量化版本是:

data = [d{2}(:) d{3}(:)];
meansattr = [meansattr1(:) meansattr2(:)];

kdiffs = bsxfun(@minus,data,permute(meansattr,[3 2 1]));

allDistances = sqrt(sum(kdiffs.^2,2)); % no need to do sqrt
allDistances = squeeze(allDistances); % Nx1xK => NxK

[closestDistance,nearestCluster] = min(allDistances,[],2); % Nx1

correctClassIds = str2num(char(d{1}(:)));
errorCount = nnz(nearestCluster ~= correctClassIds);

errorCountclosestDistancenearestCluster 中的结果与之前的解决方案等价。如代码注释所示,您可以删除 sqrt 并在 errorCountnearestCluster 中获得相同的结果。

假设你想做下一步更新meansattr1meansattr2

% Calculate the NEW cluster centers (mean the data)
meansattr_new = zeros(numClasses,numDims);
clustersizes = zeros(numClasses,1);
for ii=1:numClasses,
    indk = nearestCluster==ii;
    clustersizes(ii) = nnz(indk);
    meansattr_new(ii,:) = mean(data(indk,:))';
end

meansattr1_next = meansattr_new(:,1);
meansattr2_next = meansattr_new(:,2);

把这一切都放在while errorCount&gt;THRESHfor jj = 1:MAXITER 中,你应该得到你想要的。

【讨论】:

  • 谢谢,这样做肯定会提高性能。我将不得不查看我的程序的其他方面,看看我可以在哪里进行类似的增强。
【解决方案3】:

我从水稻的解决方案开始,变量名的简单替换:

[closest, cl] = min( (d{2}(m) - meansattr1).^2 +(d{3}(m) - meansattr2).^
2);

因此我们有一个单行for循环,常见的策略:将其制作一个函数并将其放入arrayfun:

f=@(x)min( (d{2}(x) - meansattr1).^2 +(d{3}(x) - meansattr2).^2)
[sqclosest,cl]=arrayfun(f,1:numel(d{2}));

%If necessary real distances could be calculated:
%closest=sqrt(sqclosest)

errorCount=sum(arrayfun(@(x,c)(1-strcmp(x,classes(c))),d{1},cl))

注意:请勿将“类”或任何其他保留字用于其他目的。

【讨论】:

    猜你喜欢
    • 2013-01-29
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2012-11-25
    • 2021-12-01
    • 2019-04-01
    • 1970-01-01
    相关资源
    最近更新 更多