You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用K-means聚类可视化CNN模型学习到的特征?求复现《Intriguing Properties of Contrastive Losses》论文中特征可视化流程的相关资源

Great question! Visualizing CNN intermediate features with K-means is a powerful way to validate if your model is learning hierarchical, generalizable representations—exactly what the paper you referenced does for SimCLR's ResNet50. Let me walk you through the key steps to reproduce this on VggNet, Xception, or other models, plus share some solid resources to deepen your understanding:

Step-by-Step Guide to Reproduce the Visualization

1. Extract Intermediate Features

First, you need to pull outputs from specific blocks in your target model (e.g., VggNet's block2/block3, Xception's middle flow modules). Here's a quick PyTorch example—TensorFlow uses similar logic with get_layer or custom callbacks:

import torch
import torchvision.models as models

# Load pre-trained model
model = models.vgg16(pretrained=True)
feature_store = {}

# Define a hook to capture layer outputs
def save_feature(layer_name):
    def hook_fn(module, input, output):
        feature_store[layer_name] = output.detach()
    return hook_fn

# Register hook for VggNet's block2 output (adjust index for your target layer)
model.features[10].register_forward_hook(save_feature('block2'))

# Pass sample images through the model
sample_input = torch.randn(1, 3, 224, 224)
model(sample_input)
block2_features = feature_store['block2']

2. Preprocess Features for Clustering

CNN features are 4D tensors ([batch_size, channels, height, width]). You’ll need to flatten the spatial dimensions and normalize features (like L2 normalization, as used in SimCLR) to ensure consistent scaling:

import numpy as np

# Reshape to [total_spatial_points, channels]
flattened_features = block2_features.permute(0, 2, 3, 1).reshape(-1, block2_features.shape[1]).numpy()
# L2 normalization
flattened_features = flattened_features / np.linalg.norm(flattened_features, axis=1, keepdims=True)

3. Run K-means Clustering

Use sklearn.cluster.KMeans to group similar features. If working with large datasets, sample a subset of features to avoid memory bottlenecks:

from sklearn.cluster import KMeans

# Set number of clusters (adjust based on your dataset or paper reference)
kmeans = KMeans(n_clusters=10, random_state=42)
cluster_labels = kmeans.fit_predict(flattened_features)

# Reshape labels back to spatial dimensions for visualization
cluster_map = cluster_labels.reshape(block2_features.shape[0], block2_features.shape[2], block2_features.shape[3])

4. Visualize Cluster Maps

Overlay or display the cluster labels alongside the original image to see which regions the model groups together:

import matplotlib.pyplot as plt

# Plot first sample's original image and cluster map
original_img = sample_input[0].permute(1, 2, 0).numpy()  # Adjust tensor shape for matplotlib

fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 6))
ax1.imshow(original_img)
ax1.set_title('Original Image')
ax2.imshow(cluster_map[0], cmap='viridis')
ax2.set_title('K-means Cluster Map (Block2 Features)')
plt.show()
  • Scikit-learn Official K-means Tutorials: The official docs walk through clustering best practices, including how to tune parameters like n_clusters and handle high-dimensional features—perfect for refining your implementation.
  • PyTorch/TensorFlow Model Interpretation Guides: Both frameworks have dedicated sections on model explainability that cover intermediate feature extraction and visualization. Look for tutorials on forward hooks (PyTorch) or layer output extraction (TensorFlow) paired with clustering.
  • Stanford CS231n Course Materials: The course’s lecture notes and homework assignments include detailed explanations of CNN feature visualization, including clustering-based analysis, with code examples you can adapt to your models.
  • Classic Computer Vision Papers: Beyond the paper you’re studying, check out Visualizing and Understanding Convolutional Networks (the seminal work on CNN feature visualization) and the original SimCLR paper’s appendix—both break down the logic behind using clustering to validate feature quality.
  • Open-Source Code Repositories: Search for keywords like "CNN feature kmeans visualization" to find full implementations that support multiple models (VggNet, Xception, etc.). These repos often include preprocessing and visualization scripts you can reuse directly.

内容的提问来源于stack exchange,提问作者VEDANT JOSHI

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.30 19:32:51