基于自定义成对Wasserstein距离的层次聚类实现咨询
Got it, let's tackle your problem. The main issue right now is that your pairwise_wasserstein function only prints out distances instead of returning a properly structured distance matrix that scipy's hierarchical clustering tools can work with. Let's fix this step by step and get your clustering up and running.
Step 1: Fix the Pairwise Distance Calculation
First, we need to modify pairwise_wasserstein to build a symmetric square distance matrix (where dist_mat[i,j] equals the distance between sample i and j, and dist_mat[i,i] = 0). This is the structure scipy's clustering functions expect.
Also, don't forget to import the missing distance module you're using in your Wasserstein calculation!
from scipy.spatial import distance def pairwise_wasserstein(points): """Compute pairwise Wasserstein distance matrix between samples""" n_samples = points.shape[0] # Initialize an empty symmetric matrix dist_mat = np.zeros((n_samples, n_samples)) for first_index in range(n_samples): for second_index in range(first_index + 1, n_samples): # Calculate distance between the two samples dist = wasserstein_distance_function(points[first_index], points[second_index]) # Fill both symmetric entries since distance is bidirectional dist_mat[first_index, second_index] = dist dist_mat[second_index, first_index] = dist return dist_mat
Step 2: Update the Clustering Function
Scipy's ward linkage can accept either the square distance matrix directly, or a condensed 1D version of it (more efficient for large datasets). We'll adjust find_clusters_formation to use the matrix correctly, then pass it to ward and fcluster.
def find_clusters_formation(data, n_clusters=3): """Method to find clusters using hierarchical clustering with Wasserstein distances""" # Get the symmetric distance matrix from our revised function dist_mat = pairwise_wasserstein(data) # Optional: Convert to condensed matrix (efficient for large datasets) condensed_dist = distance.squareform(dist_mat) # Perform ward linkage (can use condensed_dist or dist_mat directly) Z = ward(condensed_dist) # Alternative: Z = linkage(dist_mat, method='ward') # Assign clusters based on max number of clusters clusters = fcluster(Z, n_clusters, criterion='maxclust') print("Cluster assignments:", clusters) return clusters
Step 3: Test the Full Pipeline
Here's the complete working code with your sample data, all pieces put together:
import numpy as np from scipy.optimize import linear_sum_assignment from scipy.cluster.hierarchy import ward, fcluster from scipy.spatial import distance # Your sample 3D array data data = np.array([[[1, 2], [3, 4], [1, 2], [3, 4], [1, 2], [3, 4], [1, 2], [3, 4], [1, 2], [3, 4]], [[5, 6], [7, 8], [5, 6], [7, 8], [5, 6], [7, 8], [5, 6], [7, 8], [5, 6], [7, 8]], [[1, 15], [3, 2], [1, 2], [5, 4], [1, 2], [3, 4], [1, 2], [3, 4], [1, 2], [3, 4]], [[5, 1], [7, 8], [5, 6], [7, 1], [5, 6], [7, 8], [5, 1], [7, 8], [5, 6], [7, 8]]]) def wasserstein_distance_function(f1, f2): min_cost = np.inf f1 = f1.reshape((10, 2)) f2 = f2.reshape((10, 2)) for l in np.linspace(0.8, 1.2, 3): for k in np.linspace(0.8, 1.2, 3): cost = distance.cdist(l * f1, k * f2, 'sqeuclidean') row_ind, col_ind = linear_sum_assignment(cost) curr_cost = cost[row_ind, col_ind].sum() if curr_cost < min_cost: min_cost = curr_cost return min_cost # Include the revised pairwise_wasserstein and find_clusters_formation functions here # Run the clustering clusters = find_clusters_formation(data, n_clusters=3)
Key Notes
- Symmetric Matrix Requirement: Hierarchical clustering relies on distance matrices being symmetric (distance from i to j is the same as j to i), which our revised
pairwise_wassersteinensures. - Condensed vs Square Matrix: The
squareformfunction converts the square matrix to a 1D array of upper-triangular values (excluding the diagonal), which is a more efficient format for large datasets. - Preserved Custom Wasserstein Logic: Your custom scaling and optimal assignment logic stays intact—we just wrapped it into a structure that works with scipy's clustering tools.
内容的提问来源于stack exchange,提问作者m1gnoc

