Matlab中的速度有效分类

cas*_*sen 7 performance matlab classification machine-learning data-mining

我有一个大小为RGB的图像uint8(576,720,3),我想将每个像素分类为一组颜色.我已经使用rgb2labRGB 转换为LAB空间,然后删除了L层,因此它现在double(576,720,2)由AB组成.

现在,我想将其归类为我在另一幅图像上训练的一些颜色,并将它们各自的AB表示计算为:

Cluster 1: -17.7903  -13.1170
Cluster 2: -30.1957   40.3520
Cluster 3:  -4.4608   47.2543
Cluster 4:  46.3738   36.5225
Cluster 5:  43.3134  -17.6443
Cluster 6:  -0.9003    1.4042
Cluster 7:   7.3884   11.5584
Run Code Online (Sandbox Code Playgroud)

现在,为了将每个像素分类/标记到簇1-7,我目前执行以下操作(伪代码):

clusters;
for each x
  for each y
    ab = im(x,y,2:3);
    dist = norm(ab - clusters); // norm of dist between ab and each cluster
    [~, idx] = min(dist);
  end
end
Run Code Online (Sandbox Code Playgroud)

然而,由于图像分辨率和我手动遍历每个x和y,这非常慢(52秒).

是否有一些我可以使用的内置函数执行相同的工作?必须有.

总结一下:我需要一种分类方法,将像素图像分类为已定义的一组聚类.

Div*_*kar 11

方法#1

对于一个N x 2大小的点/像素数组,你可以避免Luispermute其他解决方案中的建议,这可能会减慢一些东西,有一种"permute-unrolled"版本的它,并且让我们的bsxfun工作朝向2D数组而不是3D数组,这必须表现更好.

因此,假设要按照N x 2大小的数组排序集群,您可以尝试使用其他bsxfun方法 -

%// Get a's and b's
im_a = im(:,:,2);
im_b = im(:,:,3);

%// Get the minimum indices that correspond to the cluster IDs
[~,idx]  = min(bsxfun(@minus,im_a(:),clusters(:,1).').^2 + ...
    bsxfun(@minus,im_b(:),clusters(:,2).').^2,[],2);
idx = reshape(idx,size(im,1),[]);
Run Code Online (Sandbox Code Playgroud)

方法#2

您可以尝试另一种利用 fast matrix multiplication in MATLAB并基于此智能解决方案的方法 -

d = 2; %// dimension of the problem size

im23 = reshape(im(:,:,2:3),[],2);

numA = size(im23,1);
numB = size(clusters,1);

A_ext = zeros(numA,3*d);
B_ext = zeros(numB,3*d);
for id = 1:d
    A_ext(:,3*id-2:3*id) = [ones(numA,1), -2*im23(:,id), im23(:,id).^2 ];
    B_ext(:,3*id-2:3*id) = [clusters(:,id).^2 ,  clusters(:,id), ones(numB,1)];
end
[~, idx] = min(A_ext * B_ext',[],2); %//'
idx = reshape(idx, size(im,1),[]); %// Desired IDs
Run Code Online (Sandbox Code Playgroud)

基于矩阵乘法的距离矩阵计算会发生什么?

让我们考虑两个矩阵A以及B我们想要计算距离矩阵的人.对于后面接下来的一个更简单的解释起见,让我们考虑A作为3 x 2B作为4 x 2大小的数组,从而表明我们正在与XY点工作.如果我们有A作为N x 3B作为M x 3大小的数组,那么这些将是X-Y-Z点.

现在,如果我们必须手动计算距离矩阵平方的第一个元素,它看起来像这样 -

first_element = ( A(1,1) – B(1,1) )^2 + ( A(1,2) – B(1,2) )^2         
Run Code Online (Sandbox Code Playgroud)

这将是 -

first_element = A(1,1)^2 + B(1,1)^2 -2*A(1,1)* B(1,1)   +  ...
                A(1,2)^2 + B(1,2)^2 -2*A(1,2)* B(1,2)    … Equation  (1)
Run Code Online (Sandbox Code Playgroud)

现在,根据我们提出的矩阵乘法,如果检查前面代码中循环的输出A_extB_ext结束后,它们将如下所示 -

在此输入图像描述

在此输入图像描述

因此,如果你执行矩阵乘法A_ext和转置B_ext,产品的第一个元素将是第一行A_ext和之间元素乘法的总和B_ext,即这些的总和 -

在此输入图像描述

结果与Equation (1)之前获得的结果相同.这将继续针对AB该列相同的列中的所有元素的所有元素A.因此,我们最终会得到完整的平方距离矩阵.这就是全部!!

矢量化变体

基于矩阵乘法的距离矩阵计算的矢量化变化是可能的,尽管它们没有看到任何大的性能改进.接下来列出两个这样的变化.

变化#1

[nA,dim] = size(A);
nB = size(B,1);

A_ext = ones(nA,dim*3);
A_ext(:,2:3:end) = -2*A;
A_ext(:,3:3:end) = A.^2;

B_ext = ones(nB,dim*3);
B_ext(:,1:3:end) = B.^2;
B_ext(:,2:3:end) = B;

distmat = A_ext * B_ext.';
Run Code Online (Sandbox Code Playgroud)

变化#2

[nA,dim] = size(A);
nB = size(B,1);

A_ext = [ones(nA*dim,1) -2*A(:) A(:).^2];
B_ext = [B(:).^2 B(:) ones(nB*dim,1)];

A_ext = reshape(permute(reshape(A_ext,nA,dim,[]),[1 3 2]),nA,[]);
B_ext = reshape(permute(reshape(B_ext,nB,dim,[]),[1 3 2]),nB,[]);

distmat = A_ext * B_ext.';
Run Code Online (Sandbox Code Playgroud)

因此,这些也可以被视为实验版本.