【问题标题】:K-means for color quantization - Code not vectorized用于颜色量化的 K 均值 - 代码未矢量化
【发布时间】:2016-10-09 12:01:08
【问题描述】:

我正在做这个由 Andrew NG 编写的关于使用 k-means 减少图像中颜色数量的练习。它工作正常,但由于代码中的所有 for 循环,我担心它有点慢,所以我想对它们进行矢量化。但是有些循环我似乎无法有效地矢量化。请帮帮我,非常感谢!

如果可能的话,请对我的编码风格提供一些反馈:)

这里是link of the exercise,这里是dataset。 正确的结果在练习的链接中给出。

这是我的代码:

function [] = KMeans()

    Image = double(imread('bird_small.tiff'));
    [rows,cols, RGB] = size(Image);
    Points = reshape(Image,rows * cols, RGB);
    K = 16;
    Centroids = zeros(K,RGB);    
    s = RandStream('mt19937ar','Seed',0);
    % Initialization :
    % Pick out K random colours and make sure they are all different
    % from each other! This prevents the situation where two of the means
    % are assigned to the exact same colour, therefore we don't have to 
    % worry about division by zero in the E-step 
    % However, if K = 16 for example, and there are only 15 colours in the
    % image, then this while loop will never exit!!! This needs to be
    % addressed in the future :( 
    % TODO : Vectorize this part!
    done = false;
    while done == false
        RowIndex = randperm(s,rows);
        ColIndex = randperm(s,cols);
        RowIndex = RowIndex(1:K);
        ColIndex = ColIndex(1:K);
        for i = 1 : K
            for j = 1 : RGB
                Centroids(i,j) = Image(RowIndex(i),ColIndex(i),j);
            end
        end
        Centroids = sort(Centroids,2);
        Centroids = unique(Centroids,'rows'); 
        if size(Centroids,1) == K
            done = true;
        end
    end;
%     imshow(imread('bird_small.tiff'))
%    
%     for i = 1 : K
%         hold on;
%         plot(RowIndex(i),ColIndex(i),'r+','MarkerSize',50)
%     end



    eps = 0.01; % Epsilon
    IterNum = 0;
    while 1
        % E-step: Estimate membership given parameters 
        % Membership: The centroid that each colour is assigned to
        % Parameters: Location of centroids
        Dist = pdist2(Points,Centroids,'euclidean');

        [~, WhichCentroid] = min(Dist,[],2);

        % M-step: Estimate parameters given membership
        % Membership: The centroid that each colour is assigned to
        % Parameters: Location of centroids
        % TODO: Vectorize this part!
        OldCentroids = Centroids;
        for i = 1 : K
            PointsInCentroid = Points((find(WhichCentroid == i))',:);
            NumOfPoints = size(PointsInCentroid,1);
            % Note that NumOfPoints is never equal to 0, as a result of
            % the initialization. Or .... ???????
            if NumOfPoints ~= 0 
                Centroids(i,:) = sum(PointsInCentroid , 1) / NumOfPoints ;
            end
        end    

        % Check for convergence: Here we use the L2 distance
        IterNum = IterNum + 1;
        Margins = sqrt(sum((Centroids - OldCentroids).^2, 2));
        if sum(Margins > eps) == 0
            break;
        end

    end
    IterNum;
    Centroids ;


    % Load the larger image
    [LargerImage,ColorMap] = imread('bird_large.tiff');
    LargerImage = double(LargerImage);
    [largeRows,largeCols,NewRGB] = size(LargerImage);  % RGB is always 3     
    % TODO: Vectorize this part!    
    largeRows
    largeCols
    NewRGB
    % Replace each of the pixel with the nearest centroid    
    NewPoints = reshape(LargerImage,largeRows * largeCols, NewRGB);
    Dist = pdist2(NewPoints,Centroids,'euclidean');
    [~,WhichCentroid] = min(Dist,[],2);
    NewPoints = Centroids(WhichCentroid,:);
    LargerImage = reshape(NewPoints,largeRows,largeCols,NewRGB);

%     for i = 1 : largeRows 
%         for j = 1 : largeCols
%             Dist = pdist2(Centroids,reshape(LargerImage(i,j,:),1,RGB),'euclidean');
%             [~,WhichCentroid] = min(Dist);    
%             LargerImage(i,j,:) = Centroids(WhichCentroid,:);            
%         end
%     end

    % Display new image
    imshow(uint8(round(LargerImage)),ColorMap)

更新:替换

for i = 1 : K
            for j = 1 : RGB
                Centroids(i,j) = Image(RowIndex(i),ColIndex(i),j);
            end
        end

与

for i = 1 : K
            Centroids(i,:) = Image(RowIndex(i),ColIndex(i),:);
        end

我认为这可以通过使用线性索引进一步向量化,但现在我应该只关注 while 循环,因为它需要大部分时间。 同样当我尝试@Dev-iL 的建议并替换时

for i = 1 : K
        PointsInCentroid = Points((find(WhichCentroid == i))',:);
        NumOfPoints = size(PointsInCentroid,1);
        % Note that NumOfPoints is never equal to 0, as a result of
        % the initialization. Or .... ???????
        if NumOfPoints ~= 0 
            Centroids(i,:) = sum(PointsInCentroid , 1) / NumOfPoints ;
        end
    end    

与

E = sparse(1:size(WhichCentroid), WhichCentroid' , 1, Num, K, Num);
Centroids = (E * spdiags(1./sum(E,1)',0,K,K))' * Points ;

结果总是更糟:当 K = 16 时,第一个需要 2,414s ,第二个需要 2,455s ; K = 32,第一个需要 4,529s,第二个需要 5,022s。似乎矢量化没有帮助,但也许我的代码有问题:(。

【问题讨论】:

  • 下次考虑将此类问题上传至Code Review。此外,您是否尝试将您的代码与已知可工作的 MATLAB 实现(例如 this)进行比较?
  • @Dev-iL 我没有。我只是在链接中的练习上测试过这个,结果和作者的一样,虽然他的实现比我的时间长
  • 我应该更好地解释一下自己:我的意思是您向矢量化寻求帮助,而我链接的代码包含 k-means 的矢量化版本,就像您想要的一样。您可以将实现的相关部分与链接代码的相应部分进行比较,以了解如何对它们进行矢量化。如果您对链接示例中算法的某些部分是如何向量化有疑问的,您可以询问它。如果链接的代码仍未针对您的需求进行足够优化,您也可以询问that。作为第一步,请确保您完全理解链接代码。
  • @Dev-iL 好的,我有一个问题。在该行中: m = X*(E*spdiags(1./sum(E,1)',0,k,k)); ,Matlab 页面说:“A = spdiags(B,d,m,n) 通过获取 B 的列并将它们沿着 d 指定的对角线放置来创建一个 m×n 稀疏矩阵。”但是这里 1/sum(E,1)' 只是一个向量!这怎么可能?我错过了什么吗? :(
  • 我假设人们看到这个问题已经得到解答,并且不再理会它。可能是因为它不是正确的站点。也许他们没有什么要补充的。

标签: performance matlab image-processing vectorization k-means


【解决方案1】:

替换

for i = 1 : K
            for j = 1 : RGB
                Centroids(i,j) = Image(RowIndex(i),ColIndex(i),j);
            end
        end

与

for i = 1 : K
            Centroids(i,:) = Image(RowIndex(i),ColIndex(i),:);
        end

我认为这可以通过使用线性索引进一步矢量化,但现在我应该只关注 while 循环,因为它需要大部分时间。 同样当我尝试@Dev-iL 的建议并替换时

for i = 1 : K
        PointsInCentroid = Points((find(WhichCentroid == i))',:);
        NumOfPoints = size(PointsInCentroid,1);
        % Note that NumOfPoints is never equal to 0, as a result of
        % the initialization. Or .... ???????
        if NumOfPoints ~= 0 
            Centroids(i,:) = sum(PointsInCentroid , 1) / NumOfPoints ;
        end
    end    

与

E = sparse(1:size(WhichCentroid), WhichCentroid' , 1, Num, K, Num);
Centroids = (E * spdiags(1./sum(E,1)',0,K,K))' * Points ;

结果总是更糟:当 K = 16 时,第一个需要 2,414s ,第二个需要 2,455s ; K = 32,第一个用了 4,529s,第二个用了 5,022s。在这种情况下,矢量化似乎没有帮助。

但是,当我更换时

 Dist = pdist2(Points,Centroids,'euclidean');
 [~, WhichCentroid] = min(Dist,[],2);

(在 while 循环中)与

    Dist = bsxfun(@minus,dot(Centroids',Centroids',1)' / 2 , Centroids * Points'  );
    [~, WhichCentroid] = min(Dist,[],1);
    WhichCentroid = WhichCentroid';

代码运行得更快,尤其是当 K 很大时 (K=32)

谢谢大家!

【讨论】:

    猜你喜欢
    • 2019-12-30
    • 1970-01-01
    • 2015-06-03
    • 2016-02-11
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多