Scipy/Sklearn中基于自定义度量的高维图像距离矩阵计算问询
解决方案
你可以通过两种方式实现对(n_samples, width, height)格式图像的自定义度量距离矩阵计算:
1. 适配Sklearn/Scipy的现有函数(推荐)
Sklearn的pairwise_distances和Scipy的cdist确实仅接受(n_samples, n_features)的2D输入,但你可以先将每张图像flatten为1D数组,再在自定义度量函数中重新reshape回2D格式计算。
示例代码(Sklearn版)
import numpy as np from sklearn.metrics import pairwise_distances ksize = 3 # 封装自定义度量,通过闭包传入图像尺寸 def create_image_metric(img_shape, ksize=3): w, h = img_shape def metric(a_flat, b_flat): # 将flatten后的数组恢复为2D图像 a = a_flat.reshape(w, h) b = b_flat.reshape(w, h) ratio = [] for i in range(w // ksize): for j in range(h // ksize): kernel_a = a[i*ksize:(i+1)*ksize, j*ksize:(j+1)*ksize] kernel_b = b[i*ksize:(i+1)*ksize, j*ksize:(j+1)*ksize] ratio.append(np.mean(kernel_a == kernel_b)) return 1 - max(ratio) return metric # 构造样本集:(n_samples, width, height) a = np.array([ [1,0,0,0,0,1], [0,0,1,0,0,0], [0,0,0,0,0,0], [0,0,0,0,1,0], [0,1,0,0,0,1], [0,0,0,0,0,0] ]) b = np.array([ [0,1,0,0,0,0], [0,0,1,0,0,1], [1,0,0,0,1,0], [0,0,0,0,1,0], [0,0,0,1,0,1], [0,0,0,1,1,1] ]) samples = np.array([a, b]) # 转换为(n_samples, n_features)格式 samples_flat = samples.reshape(samples.shape[0], -1) # 创建适配的度量函数 img_metric = create_image_metric(img_shape=(6,6)) # 计算距离矩阵 distance_matrix = pairwise_distances(samples_flat, metric=img_metric) print(distance_matrix)
2. 手动构建距离矩阵
如果样本量不大,也可以直接遍历所有样本对,调用原自定义度量函数生成距离矩阵,无需flatten操作:
import numpy as np ksize = 3 def custom_metric(a, b): w, h = a.shape ratio = [] for i in range(w // ksize): for j in range(h // ksize): kernel_a = a[i*ksize:(i+1)*ksize, j*ksize:(j+1)*ksize] kernel_b = b[i*ksize:(i+1)*ksize, j*ksize:(j+1)*ksize] ratio.append(np.mean(kernel_a == kernel_b)) return 1 - max(ratio) def compute_distance_matrix(samples): n_samples = samples.shape[0] dist_matrix = np.zeros((n_samples, n_samples)) # 遍历所有样本对(利用对称性减少计算量) for i in range(n_samples): for j in range(i, n_samples): dist = custom_metric(samples[i], samples[j]) dist_matrix[i, j] = dist dist_matrix[j, i] = dist return dist_matrix # 输入为(n_samples, width, height)格式 samples = np.array([a, b]) distance_matrix = compute_distance_matrix(samples) print(distance_matrix)
关键说明
- Sklearn/Scipy的官方距离函数不直接支持高维输入,flatten是最通用的适配方案,不会影响度量计算的结果。
- 手动构建矩阵适合小样本场景,避免了数组reshape的开销,代码更直观。
内容的提问来源于stack exchange,提问作者Tartempion
相关产品推荐
相关产品推荐

