如何在MATLAB中向量化K-NN算法嵌套循环以提升代码效率
Optimizing K-NN Classification with Vectorization in MATLAB
Alright, let's fix that slow nested loop in your K-NN implementation. The double loop over test samples and K-values is killing your efficiency—vectorization in MATLAB will cut runtime drastically here. Here's how to rewrite this properly:
Full Vectorized Implementation
function [Cpreds] = my_knn_classify(Xtrn, Ctrn, Xtst, Ks) % Input: % Xtrn : M-by-D training data matrix % Ctrn : M-by-1 training data label vector % Xtst : N-by-D test data matrix % Ks : L-by-1 vector of k-values for nearest neighbors % Output: % Cpreds : N-by-L matrix of predicted labels % Step 1: Compute squared Euclidean distances (skip sqrt for speed) dist_sq = sum(Xtst.^2, 2) - 2 * Xtst * Xtrn' + sum(Xtrn.^2, 1)'; % Step 2: Get sorted indices of training samples by distance (closest first) [~, sorted_idx] = sort(dist_sq, 2, 'ascend'); % Step 3: Pre-sort training labels based on distance for all test samples trn_labels_sorted = Ctrn(sorted_idx); % N-by-M matrix of ordered labels % Step 4: Compute predictions for each K value L = length(Ks); Cpreds = zeros(size(Xtst, 1), L); for c = 1:L k = Ks(c); % Extract first k nearest neighbor labels for every test sample k_neighbor_labels = trn_labels_sorted(:, 1:k); % Calculate mode (most frequent label) for each test sample [predicted_labels, ~] = mode(k_neighbor_labels, 2); Cpreds(:, c) = predicted_labels; end end
Key Optimizations Explained
- Eliminated the outer test sample loop: We compute distances and sorted neighbor indices for all test samples in one go using matrix operations—this is where MATLAB shines.
- Reduced loop count to only K-values: Instead of
N*Literations, we now only loop over theLK-values, which is usually a tiny number compared to the number of test samplesN. - Skipped unnecessary sqrt calculation: Squared Euclidean distances work just as well for sorting neighbors, and avoiding the square root saves significant computation time.
- Precomputed sorted labels: We only index into the training labels once to create a sorted matrix, so we don't repeat index lookups for each K-value.
Notes
If you already had a precomputed matrix of neighbor indices (the in variable in your original code), you can skip the distance calculation and sorting steps—just replace sorted_idx with your precomputed index matrix.
内容的提问来源于stack exchange,提问作者Sameer Karim
相关产品推荐
相关产品推荐

