111
This commit is contained in:
@@ -0,0 +1,251 @@
|
||||
"""
|
||||
使用提取的特征进行画师风格聚类
|
||||
支持多种聚类算法和可视化
|
||||
"""
|
||||
import argparse
|
||||
import numpy as np
|
||||
import json
|
||||
from pathlib import Path
|
||||
import matplotlib.pyplot as plt
|
||||
from sklearn.cluster import KMeans, DBSCAN, AgglomerativeClustering
|
||||
from sklearn.manifold import TSNE
|
||||
from sklearn.decomposition import PCA
|
||||
import seaborn as sns
|
||||
|
||||
|
||||
def get_args_parser():
|
||||
parser = argparse.ArgumentParser('Artist Style Clustering', add_help=False)
|
||||
|
||||
parser.add_argument('--features', required=True, type=str,
|
||||
help='Path to features.npy file')
|
||||
parser.add_argument('--image-names', required=True, type=str,
|
||||
help='Path to image_names.txt file')
|
||||
parser.add_argument('--output', default='./output/clustering', type=str,
|
||||
help='Output directory')
|
||||
|
||||
# 聚类参数
|
||||
parser.add_argument('--method', default='kmeans', type=str,
|
||||
choices=['kmeans', 'dbscan', 'hierarchical'],
|
||||
help='Clustering method')
|
||||
parser.add_argument('--n-clusters', default=10, type=int,
|
||||
help='Number of clusters (for kmeans and hierarchical)')
|
||||
parser.add_argument('--eps', default=0.5, type=float,
|
||||
help='DBSCAN eps parameter')
|
||||
parser.add_argument('--min-samples', default=5, type=int,
|
||||
help='DBSCAN min_samples parameter')
|
||||
|
||||
# 可视化参数
|
||||
parser.add_argument('--visualize', action='store_true', default=True,
|
||||
help='Create visualization')
|
||||
parser.add_argument('--viz-method', default='tsne', type=str,
|
||||
choices=['tsne', 'pca'],
|
||||
help='Dimensionality reduction method for visualization')
|
||||
parser.add_argument('--perplexity', default=30, type=int,
|
||||
help='t-SNE perplexity parameter')
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
def load_features(features_path, image_names_path):
|
||||
"""加载特征和图像名称"""
|
||||
features = np.load(features_path)
|
||||
with open(image_names_path, 'r') as f:
|
||||
image_names = [line.strip() for line in f]
|
||||
|
||||
print(f"Loaded features: {features.shape}")
|
||||
print(f"Number of images: {len(image_names)}")
|
||||
|
||||
return features, image_names
|
||||
|
||||
|
||||
def perform_clustering(features, method='kmeans', n_clusters=10, eps=0.5, min_samples=5):
|
||||
"""执行聚类"""
|
||||
print(f"\nPerforming clustering with method: {method}")
|
||||
|
||||
if method == 'kmeans':
|
||||
clusterer = KMeans(n_clusters=n_clusters, random_state=42, n_init=10)
|
||||
labels = clusterer.fit_predict(features)
|
||||
print(f"K-Means clustering completed with {n_clusters} clusters")
|
||||
|
||||
elif method == 'dbscan':
|
||||
clusterer = DBSCAN(eps=eps, min_samples=min_samples)
|
||||
labels = clusterer.fit_predict(features)
|
||||
n_clusters = len(set(labels)) - (1 if -1 in labels else 0)
|
||||
n_noise = list(labels).count(-1)
|
||||
print(f"DBSCAN clustering completed")
|
||||
print(f"Number of clusters: {n_clusters}")
|
||||
print(f"Number of noise points: {n_noise}")
|
||||
|
||||
elif method == 'hierarchical':
|
||||
clusterer = AgglomerativeClustering(n_clusters=n_clusters)
|
||||
labels = clusterer.fit_predict(features)
|
||||
print(f"Hierarchical clustering completed with {n_clusters} clusters")
|
||||
|
||||
return labels
|
||||
|
||||
|
||||
def reduce_dimensions(features, method='tsne', perplexity=30):
|
||||
"""降维用于可视化"""
|
||||
print(f"\nReducing dimensions with {method}")
|
||||
|
||||
if method == 'tsne':
|
||||
reducer = TSNE(n_components=2, perplexity=perplexity, random_state=42)
|
||||
features_2d = reducer.fit_transform(features)
|
||||
print("t-SNE reduction completed")
|
||||
|
||||
elif method == 'pca':
|
||||
reducer = PCA(n_components=2, random_state=42)
|
||||
features_2d = reducer.fit_transform(features)
|
||||
explained_var = reducer.explained_variance_ratio_
|
||||
print(f"PCA reduction completed")
|
||||
print(f"Explained variance: {explained_var[0]:.3f}, {explained_var[1]:.3f}")
|
||||
|
||||
return features_2d
|
||||
|
||||
|
||||
def visualize_clusters(features_2d, labels, output_dir, method_name):
|
||||
"""可视化聚类结果"""
|
||||
print("\nCreating visualization...")
|
||||
|
||||
plt.figure(figsize=(12, 10))
|
||||
|
||||
# 获取唯一的标签(聚类)
|
||||
unique_labels = set(labels)
|
||||
n_clusters = len(unique_labels) - (1 if -1 in unique_labels else 0)
|
||||
|
||||
# 使用不同颜色
|
||||
colors = plt.cm.Spectral(np.linspace(0, 1, len(unique_labels)))
|
||||
|
||||
for label, color in zip(unique_labels, colors):
|
||||
if label == -1:
|
||||
# 噪声点用黑色
|
||||
color = [0, 0, 0, 1]
|
||||
marker = 'x'
|
||||
label_name = 'Noise'
|
||||
else:
|
||||
marker = 'o'
|
||||
label_name = f'Cluster {label}'
|
||||
|
||||
mask = labels == label
|
||||
plt.scatter(features_2d[mask, 0], features_2d[mask, 1],
|
||||
c=[color], label=label_name, marker=marker, s=50, alpha=0.6)
|
||||
|
||||
plt.title(f'Artist Style Clustering ({method_name})\nTotal Clusters: {n_clusters}',
|
||||
fontsize=14, fontweight='bold')
|
||||
plt.xlabel('Dimension 1', fontsize=12)
|
||||
plt.ylabel('Dimension 2', fontsize=12)
|
||||
plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left', fontsize=10)
|
||||
plt.tight_layout()
|
||||
|
||||
# 保存图像
|
||||
output_path = output_dir / f'clustering_{method_name}.png'
|
||||
plt.savefig(output_path, dpi=300, bbox_inches='tight')
|
||||
print(f"Visualization saved to: {output_path}")
|
||||
|
||||
plt.close()
|
||||
|
||||
|
||||
def save_clustering_results(labels, image_names, output_dir):
|
||||
"""保存聚类结果"""
|
||||
# 按聚类分组
|
||||
clusters = {}
|
||||
for label, name in zip(labels, image_names):
|
||||
label = int(label)
|
||||
if label not in clusters:
|
||||
clusters[label] = []
|
||||
clusters[label].append(name)
|
||||
|
||||
# 保存JSON格式
|
||||
json_path = output_dir / 'clustering_results.json'
|
||||
with open(json_path, 'w', encoding='utf-8') as f:
|
||||
json.dump(clusters, f, indent=2, ensure_ascii=False)
|
||||
print(f"Clustering results saved to: {json_path}")
|
||||
|
||||
# 保存文本格式(易读)
|
||||
txt_path = output_dir / 'clustering_results.txt'
|
||||
with open(txt_path, 'w', encoding='utf-8') as f:
|
||||
for label in sorted(clusters.keys()):
|
||||
if label == -1:
|
||||
f.write(f"Noise ({len(clusters[label])} images):\n")
|
||||
else:
|
||||
f.write(f"Cluster {label} ({len(clusters[label])} images):\n")
|
||||
for name in clusters[label]:
|
||||
f.write(f" - {name}\n")
|
||||
f.write("\n")
|
||||
print(f"Clustering results (text) saved to: {txt_path}")
|
||||
|
||||
# 统计信息
|
||||
stats = {
|
||||
'total_images': len(image_names),
|
||||
'n_clusters': len([k for k in clusters.keys() if k != -1]),
|
||||
'cluster_sizes': {int(k): len(v) for k, v in clusters.items()}
|
||||
}
|
||||
|
||||
stats_path = output_dir / 'clustering_stats.json'
|
||||
with open(stats_path, 'w', encoding='utf-8') as f:
|
||||
json.dump(stats, f, indent=2)
|
||||
print(f"Clustering statistics saved to: {stats_path}")
|
||||
|
||||
return clusters
|
||||
|
||||
|
||||
def print_cluster_statistics(clusters):
|
||||
"""打印聚类统计信息"""
|
||||
print("\n" + "="*50)
|
||||
print("Clustering Statistics")
|
||||
print("="*50)
|
||||
|
||||
total_images = sum(len(v) for v in clusters.values())
|
||||
n_clusters = len([k for k in clusters.keys() if k != -1])
|
||||
|
||||
print(f"Total images: {total_images}")
|
||||
print(f"Number of clusters: {n_clusters}")
|
||||
|
||||
if -1 in clusters:
|
||||
print(f"Noise points: {len(clusters[-1])}")
|
||||
|
||||
print("\nCluster sizes:")
|
||||
for label in sorted(clusters.keys()):
|
||||
if label == -1:
|
||||
print(f" Noise: {len(clusters[label])} images")
|
||||
else:
|
||||
print(f" Cluster {label}: {len(clusters[label])} images")
|
||||
|
||||
print("="*50)
|
||||
|
||||
|
||||
def main(args):
|
||||
# 创建输出目录
|
||||
output_dir = Path(args.output)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 加载特征
|
||||
features, image_names = load_features(args.features, args.image_names)
|
||||
|
||||
# 执行聚类
|
||||
labels = perform_clustering(
|
||||
features,
|
||||
method=args.method,
|
||||
n_clusters=args.n_clusters,
|
||||
eps=args.eps,
|
||||
min_samples=args.min_samples
|
||||
)
|
||||
|
||||
# 保存聚类结果
|
||||
clusters = save_clustering_results(labels, image_names, output_dir)
|
||||
|
||||
# 打印统计信息
|
||||
print_cluster_statistics(clusters)
|
||||
|
||||
# 可视化
|
||||
if args.visualize:
|
||||
features_2d = reduce_dimensions(features, method=args.viz_method, perplexity=args.perplexity)
|
||||
visualize_clusters(features_2d, labels, output_dir, args.method)
|
||||
|
||||
print(f"\n✓ Clustering completed! Results saved to: {output_dir}")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser('Artist Style Clustering', parents=[get_args_parser()])
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,448 @@
|
||||
import os
|
||||
import sys
|
||||
import json
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
from typing import Dict, Optional
|
||||
|
||||
sys.path.append(os.path.dirname(__file__))
|
||||
|
||||
import folder_paths
|
||||
|
||||
from timm.data import resolve_data_config
|
||||
from timm.data.transforms_factory import create_transform
|
||||
from timm.models import create_model
|
||||
|
||||
from model import lsnet_artist # noqa: F401
|
||||
|
||||
from inference_artist import (
|
||||
load_checkpoint_state,
|
||||
normalize_state_dict_keys,
|
||||
resolve_num_classes,
|
||||
resolve_feature_dim,
|
||||
load_class_mapping
|
||||
)
|
||||
|
||||
from sklearn.cluster import KMeans, DBSCAN, AgglomerativeClustering
|
||||
from sklearn.manifold import TSNE
|
||||
from sklearn.decomposition import PCA
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
|
||||
class LSNetModelLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
base_dir = os.path.join(folder_paths.models_dir, 'lsnet')
|
||||
subfolders = []
|
||||
if os.path.exists(base_dir):
|
||||
subfolders = [f for f in os.listdir(base_dir) if os.path.isdir(os.path.join(base_dir, f))]
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"model_folder": (subfolders, {"default": subfolders[0] if subfolders else ""}),
|
||||
"device": ("STRING", {"default": "cuda"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LSNET_MODEL",)
|
||||
FUNCTION = "load"
|
||||
CATEGORY = "LSNet"
|
||||
|
||||
def load(self, model_folder, device):
|
||||
base_dir = os.path.join(folder_paths.models_dir, 'lsnet')
|
||||
model_dir = os.path.join(base_dir, model_folder)
|
||||
checkpoint_path = os.path.join(model_dir, "checkpoint.pth")
|
||||
csv_path = os.path.join(model_dir, "class_mapping.csv")
|
||||
|
||||
if not os.path.exists(checkpoint_path):
|
||||
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
|
||||
if not os.path.exists(csv_path):
|
||||
raise FileNotFoundError(f"Class mapping CSV not found: {csv_path}")
|
||||
class_mapping = load_class_mapping(csv_path)
|
||||
state_dict = load_checkpoint_state(checkpoint_path)
|
||||
state_dict = normalize_state_dict_keys(state_dict)
|
||||
num_classes = resolve_num_classes(None, class_mapping, state_dict)
|
||||
feature_dim = resolve_feature_dim(None, state_dict)
|
||||
model = create_model(
|
||||
'lsnet_xl_artist',
|
||||
pretrained=False,
|
||||
num_classes=num_classes,
|
||||
feature_dim=feature_dim,
|
||||
)
|
||||
model.load_state_dict(state_dict, strict=False)
|
||||
model.to(device)
|
||||
model.eval()
|
||||
config = resolve_data_config({}, model=model)
|
||||
transform = create_transform(**config)
|
||||
model_bundle = {
|
||||
'model': model,
|
||||
'transform': transform,
|
||||
'class_mapping': class_mapping,
|
||||
'device': device
|
||||
}
|
||||
|
||||
return (model_bundle,)
|
||||
|
||||
class LSNetArtistInferenceNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"model": ("LSNET_MODEL",),
|
||||
"top_k": ("INT", {"default": 5, "min": 1, "max": 100}),
|
||||
"threshold": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING", "STRING")
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LSNet"
|
||||
|
||||
def process(self, image, model, top_k, threshold):
|
||||
model_bundle = model
|
||||
model = model_bundle['model']
|
||||
transform = model_bundle['transform']
|
||||
class_mapping = model_bundle['class_mapping']
|
||||
device = model_bundle['device']
|
||||
|
||||
if image.ndim == 4:
|
||||
image = image[0]
|
||||
image = (image * 255).clamp(0, 255).byte().cpu().numpy()
|
||||
pil_image = Image.fromarray(image)
|
||||
|
||||
# Preprocess image
|
||||
image_tensor = transform(pil_image).unsqueeze(0) # Add batch dimension
|
||||
|
||||
# Classify
|
||||
with torch.no_grad():
|
||||
image_tensor = image_tensor.to(device)
|
||||
logits = model(image_tensor, return_features=False)
|
||||
probs = F.softmax(logits, dim=-1)
|
||||
top_probs, top_indices = torch.topk(probs, k=min(top_k, probs.size(-1)), dim=-1)
|
||||
|
||||
results = []
|
||||
for prob, idx in zip(top_probs[0].cpu().numpy(), top_indices[0].cpu().numpy()):
|
||||
if prob >= threshold:
|
||||
class_id = int(idx)
|
||||
class_name = class_mapping.get(class_id, f"Class {class_id}")
|
||||
results.append({
|
||||
'class_id': class_id,
|
||||
'class_name': class_name,
|
||||
'probability': float(prob)
|
||||
})
|
||||
|
||||
# Limit to top_k if more results after filtering
|
||||
if len(results) > top_k:
|
||||
results = results[:top_k]
|
||||
|
||||
# Prepare outputs
|
||||
tags = [res['class_name'] for res in results]
|
||||
tag_string = ",".join(tags)
|
||||
json_output = json.dumps(results, ensure_ascii=False)
|
||||
|
||||
return (tag_string, json_output)
|
||||
|
||||
class LSNetArtistSimilarityNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"processed_image": ("IMAGE",),
|
||||
"reference_images": ("IMAGE",),
|
||||
"model": ("LSNET_MODEL",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LSNet"
|
||||
|
||||
def process(self, processed_image, reference_images, model):
|
||||
model_bundle = model
|
||||
model = model_bundle['model']
|
||||
transform = model_bundle['transform']
|
||||
device = model_bundle['device']
|
||||
|
||||
def image_to_tensor(img):
|
||||
if img.ndim == 4:
|
||||
img = img[0]
|
||||
img = (img * 255).clamp(0, 255).byte().cpu().numpy()
|
||||
pil_img = Image.fromarray(img)
|
||||
return transform(pil_img).unsqueeze(0)
|
||||
|
||||
processed_tensor = image_to_tensor(processed_image)
|
||||
with torch.no_grad():
|
||||
processed_tensor = processed_tensor.to(device)
|
||||
processed_features = model(processed_tensor, return_features=True).cpu().numpy()[0]
|
||||
|
||||
references = []
|
||||
similarities = []
|
||||
num_refs = reference_images.shape[0] if reference_images.ndim == 4 else 1
|
||||
for i in range(num_refs):
|
||||
ref_img = reference_images[i] if reference_images.ndim == 4 else reference_images
|
||||
ref_tensor = image_to_tensor(ref_img)
|
||||
with torch.no_grad():
|
||||
ref_tensor = ref_tensor.to(device)
|
||||
ref_features = model(ref_tensor, return_features=True).cpu().numpy()[0]
|
||||
references.append(ref_features.tolist())
|
||||
sim = np.dot(processed_features, ref_features) / (np.linalg.norm(processed_features) * np.linalg.norm(ref_features))
|
||||
similarities.append(float(sim))
|
||||
|
||||
result = {
|
||||
"processed_features": processed_features.tolist(),
|
||||
"reference_features": references,
|
||||
"similarities": similarities
|
||||
}
|
||||
json_output = json.dumps(result, ensure_ascii=False)
|
||||
|
||||
return (json_output,)
|
||||
|
||||
class LSNetCommonFeaturesNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"reference_images": ("IMAGE",),
|
||||
"model": ("LSNET_MODEL",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("TENSOR",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LSNet"
|
||||
|
||||
def process(self, reference_images, model):
|
||||
model_bundle = model
|
||||
model = model_bundle['model']
|
||||
transform = model_bundle['transform']
|
||||
device = model_bundle['device']
|
||||
|
||||
def image_to_tensor(img):
|
||||
if img.ndim == 4:
|
||||
img = img[0]
|
||||
img = (img * 255).clamp(0, 255).byte().cpu().numpy()
|
||||
pil_img = Image.fromarray(img)
|
||||
return transform(pil_img).unsqueeze(0)
|
||||
|
||||
references = []
|
||||
num_refs = reference_images.shape[0] if reference_images.ndim == 4 else 1
|
||||
for i in range(num_refs):
|
||||
ref_img = reference_images[i] if reference_images.ndim == 4 else reference_images
|
||||
ref_tensor = image_to_tensor(ref_img)
|
||||
with torch.no_grad():
|
||||
ref_tensor = ref_tensor.to(device)
|
||||
ref_features = model(ref_tensor, return_features=True).cpu().numpy()[0]
|
||||
references.append(ref_features)
|
||||
|
||||
if references:
|
||||
common_features = np.mean(np.array(references), axis=0)
|
||||
else:
|
||||
common_features = np.zeros(384)
|
||||
return (torch.tensor(common_features),)
|
||||
|
||||
class LSNetClusteringNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"method": (["kmeans", "dbscan", "hierarchical"], {"default": "kmeans"}),
|
||||
"n_clusters": ("INT", {"default": 10, "min": 2, "max": 100}),
|
||||
"eps": ("FLOAT", {"default": 0.5, "min": 0.1, "max": 10.0}),
|
||||
"min_samples": ("INT", {"default": 5, "min": 1, "max": 50}),
|
||||
"visualize": ("BOOLEAN", {"default": True}),
|
||||
"viz_method": (["tsne", "pca"], {"default": "tsne"}),
|
||||
"perplexity": ("INT", {"default": 30, "min": 5, "max": 100}),
|
||||
},
|
||||
"optional": {
|
||||
"group_1": ("TENSOR",),
|
||||
"group_2": ("TENSOR",),
|
||||
"group_3": ("TENSOR",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING", "IMAGE")
|
||||
FUNCTION = "cluster"
|
||||
CATEGORY = "LSNet"
|
||||
|
||||
def cluster(self, method, n_clusters, eps, min_samples, visualize, viz_method, perplexity, group_1=None, group_2=None, group_3=None):
|
||||
groups = []
|
||||
group_sizes = []
|
||||
for g in [group_1, group_2, group_3]:
|
||||
if g is not None:
|
||||
groups.append(g.cpu().numpy())
|
||||
group_sizes.append(g.shape[0])
|
||||
|
||||
if not groups:
|
||||
return (json.dumps({"error": "No groups provided"}), torch.zeros(1, 64, 64, 3))
|
||||
|
||||
features_np = np.vstack(groups)
|
||||
if method == "kmeans":
|
||||
clusterer = KMeans(n_clusters=n_clusters, random_state=42)
|
||||
labels = clusterer.fit_predict(features_np)
|
||||
centers = clusterer.cluster_centers_
|
||||
elif method == "dbscan":
|
||||
clusterer = DBSCAN(eps=eps, min_samples=min_samples)
|
||||
labels = clusterer.fit_predict(features_np)
|
||||
centers = None
|
||||
elif method == "hierarchical":
|
||||
clusterer = AgglomerativeClustering(n_clusters=n_clusters)
|
||||
labels = clusterer.fit_predict(features_np)
|
||||
centers = None
|
||||
|
||||
result = {
|
||||
"method": method,
|
||||
"n_samples": len(features_np),
|
||||
"group_sizes": group_sizes,
|
||||
"labels": labels.tolist(),
|
||||
}
|
||||
if centers is not None:
|
||||
result["centers"] = centers.tolist()
|
||||
|
||||
json_output = json.dumps(result, ensure_ascii=False)
|
||||
|
||||
if visualize and len(features_np) > 1:
|
||||
if viz_method == "tsne":
|
||||
reducer = TSNE(n_components=2, perplexity=min(perplexity, len(features_np)-1), random_state=42)
|
||||
else:
|
||||
reducer = PCA(n_components=2, random_state=42)
|
||||
|
||||
reduced_features = reducer.fit_transform(features_np)
|
||||
|
||||
plt.figure(figsize=(10, 8))
|
||||
unique_labels = np.unique(labels)
|
||||
colors = plt.cm.rainbow(np.linspace(0, 1, len(unique_labels)))
|
||||
|
||||
for label, color in zip(unique_labels, colors):
|
||||
mask = labels == label
|
||||
plt.scatter(reduced_features[mask, 0], reduced_features[mask, 1],
|
||||
color=color, label=f'Cluster {label}', alpha=0.7)
|
||||
|
||||
plt.title(f'{method.upper()} Clustering ({viz_method.upper()})')
|
||||
plt.legend()
|
||||
plt.tight_layout()
|
||||
|
||||
fig = plt.gcf()
|
||||
fig.canvas.draw()
|
||||
img_array = np.frombuffer(fig.canvas.tostring_rgb(), dtype=np.uint8)
|
||||
img_array = img_array.reshape(fig.canvas.get_width_height()[::-1] + (3,))
|
||||
pil_image = Image.fromarray(img_array)
|
||||
plt.close()
|
||||
|
||||
viz_tensor = torch.from_numpy(np.array(pil_image)).float() / 255.0
|
||||
if viz_tensor.ndim == 3:
|
||||
viz_tensor = viz_tensor.unsqueeze(0)
|
||||
else:
|
||||
viz_tensor = torch.zeros(1, 64, 64, 3)
|
||||
|
||||
return (json_output, viz_tensor)
|
||||
|
||||
class LSNetFeatureComparisonNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"model": ("LSNET_MODEL",),
|
||||
},
|
||||
"optional": {
|
||||
"group_1": ("TENSOR",),
|
||||
"group_2": ("TENSOR",),
|
||||
"group_3": ("TENSOR",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "compare"
|
||||
CATEGORY = "LSNet"
|
||||
|
||||
def compare(self, image, model, group_1=None, group_2=None, group_3=None):
|
||||
model_bundle = model
|
||||
model = model_bundle['model']
|
||||
transform = model_bundle['transform']
|
||||
device = model_bundle['device']
|
||||
|
||||
if image.ndim == 4:
|
||||
image = image[0]
|
||||
image_np = (image * 255).clamp(0, 255).byte().cpu().numpy()
|
||||
pil_image = Image.fromarray(image_np)
|
||||
image_tensor = transform(pil_image).unsqueeze(0)
|
||||
|
||||
with torch.no_grad():
|
||||
image_tensor = image_tensor.to(device)
|
||||
query_features = model(image_tensor, return_features=True).cpu().numpy()[0]
|
||||
|
||||
groups = []
|
||||
for g in [group_1, group_2, group_3]:
|
||||
if g is not None:
|
||||
groups.append(g.cpu().numpy())
|
||||
|
||||
if not groups:
|
||||
return (json.dumps({"error": "No groups provided"}),)
|
||||
|
||||
similarities = []
|
||||
for group_feat in groups:
|
||||
sim = np.dot(query_features, group_feat) / (np.linalg.norm(query_features) * np.linalg.norm(group_feat))
|
||||
similarities.append(float(sim))
|
||||
|
||||
best_index = np.argmax(similarities)
|
||||
best_similarity = similarities[best_index]
|
||||
result = {
|
||||
"best_group_index": int(best_index),
|
||||
"best_similarity": best_similarity,
|
||||
"all_similarities": similarities
|
||||
}
|
||||
json_output = json.dumps(result, ensure_ascii=False)
|
||||
|
||||
return (json_output,)
|
||||
|
||||
class LSNetArtistImageConnector:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image_1": ("IMAGE",),
|
||||
"image_2": ("IMAGE",),
|
||||
"image_3": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "connect"
|
||||
CATEGORY = "LSNet"
|
||||
|
||||
def connect(self, image_1, image_2, image_3):
|
||||
def normalize_image(img):
|
||||
if img.ndim == 4:
|
||||
img = img[0]
|
||||
return img.unsqueeze(0)
|
||||
|
||||
img1 = normalize_image(image_1)
|
||||
img2 = normalize_image(image_2)
|
||||
img3 = normalize_image(image_3)
|
||||
|
||||
stacked = torch.cat([img1, img2, img3], dim=0)
|
||||
return (stacked,)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"LSNetModelLoader": LSNetModelLoader,
|
||||
"LSNetArtistInference": LSNetArtistInferenceNode,
|
||||
"LSNetArtistSimilarity": LSNetArtistSimilarityNode,
|
||||
"LSNetCommonFeatures": LSNetCommonFeaturesNode,
|
||||
"LSNetClustering": LSNetClusteringNode,
|
||||
"LSNetFeatureComparison": LSNetFeatureComparisonNode,
|
||||
"LSNetArtistImageConnector": LSNetArtistImageConnector
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"LSNetModelLoader": "LSNet Model Loader",
|
||||
"LSNetArtistInference": "LSNet Artist Inference",
|
||||
"LSNetArtistSimilarity": "LSNet Artist Similarity",
|
||||
"LSNetCommonFeatures": "LSNet Common Features",
|
||||
"LSNetClustering": "LSNet Clustering",
|
||||
"LSNetFeatureComparison": "LSNet Feature Comparison",
|
||||
"LSNetArtistImageConnector": "LSNet Image Connector"
|
||||
}
|
||||
@@ -0,0 +1,475 @@
|
||||
"""
|
||||
画师风格模型推理脚本
|
||||
支持两种模式:
|
||||
1. 聚类模式:提取特征向量用于聚类
|
||||
2. 分类模式:直接输出分类结果
|
||||
"""
|
||||
import argparse
|
||||
import csv
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Dict, Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from PIL import Image
|
||||
from timm.data import resolve_data_config
|
||||
from timm.data.transforms_factory import create_transform
|
||||
from timm.models import create_model
|
||||
|
||||
# Ensure custom LSNet artist models are registered with timm
|
||||
from model import lsnet_artist # noqa: F401
|
||||
|
||||
|
||||
def get_args_parser():
|
||||
parser = argparse.ArgumentParser('Artist Style Inference', add_help=False)
|
||||
|
||||
# 模型参数
|
||||
parser.add_argument('--model', default='lsnet_t_artist', type=str,
|
||||
choices=['lsnet_t_artist', 'lsnet_s_artist', 'lsnet_b_artist', 'lsnet_l_artist', 'lsnet_xl_artist'],
|
||||
help='Model architecture')
|
||||
parser.add_argument('--checkpoint', required=True, type=str,
|
||||
help='Path to model checkpoint')
|
||||
parser.add_argument('--num-classes', default=None, type=int,
|
||||
help='Number of classes. If omitted, will try to infer from checkpoint or CSV mapping.')
|
||||
parser.add_argument('--feature-dim', default=None, type=int,
|
||||
help='Feature dimension')
|
||||
|
||||
# 推理模式
|
||||
parser.add_argument('--mode', default='classify', type=str,
|
||||
choices=['classify', 'cluster', 'both'],
|
||||
help='Inference mode: classify (with head), cluster (features only), or both')
|
||||
|
||||
# 输入输出
|
||||
parser.add_argument('--input', required=True, type=str,
|
||||
help='Input image path or directory')
|
||||
parser.add_argument('--output', default='./output/inference', type=str,
|
||||
help='Output directory')
|
||||
parser.add_argument('--class-csv', default=None, type=str,
|
||||
help='Path to class mapping CSV exported during training')
|
||||
|
||||
# 其他参数
|
||||
parser.add_argument('--device', default='cuda', type=str,
|
||||
help='Device to use')
|
||||
parser.add_argument('--batch-size', default=32, type=int,
|
||||
help='Batch size for batch inference')
|
||||
parser.add_argument('--allow-head-reinit', action='store_true', default=False,
|
||||
help='Allow re-initializing classification head when checkpoint classes mismatch')
|
||||
parser.add_argument('--top-k', default=5, type=int,
|
||||
help='Number of top predictions to show (default: 5)')
|
||||
parser.add_argument('--threshold', default=0.0, type=float,
|
||||
help='Probability threshold to filter predictions (default: 0.0)')
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
def load_checkpoint_state(checkpoint_path: str):
|
||||
"""加载 checkpoint 并返回模型权重"""
|
||||
checkpoint = torch.load(checkpoint_path, map_location='cpu', weights_only=False)
|
||||
if isinstance(checkpoint, dict):
|
||||
if 'model' in checkpoint:
|
||||
return checkpoint['model']
|
||||
if 'model_ema' in checkpoint:
|
||||
return checkpoint['model_ema']
|
||||
return checkpoint
|
||||
|
||||
|
||||
def normalize_state_dict_keys(state_dict):
|
||||
"""移除分布式训练前缀等冗余标记"""
|
||||
normalized = {}
|
||||
for key, value in state_dict.items():
|
||||
if key.startswith('module.'):
|
||||
new_key = key[len('module.'):]
|
||||
else:
|
||||
new_key = key
|
||||
normalized[new_key] = value
|
||||
return normalized
|
||||
|
||||
|
||||
def resolve_num_classes(num_classes_arg: Optional[int],
|
||||
class_mapping: Optional[Dict[int, str]],
|
||||
state_dict) -> int:
|
||||
"""根据参数、CSV 或 checkpoint 推断类别数"""
|
||||
# 优先使用CSV中的类别数
|
||||
if class_mapping:
|
||||
csv_classes = len(class_mapping)
|
||||
if num_classes_arg is not None and num_classes_arg != csv_classes:
|
||||
print(f"[Warning] 提供的 num_classes={num_classes_arg} 与 CSV 中的类别数 {csv_classes} 不一致,已使用 CSV 的值。")
|
||||
return csv_classes
|
||||
|
||||
# 如果没有CSV,使用参数
|
||||
if num_classes_arg is not None:
|
||||
return num_classes_arg
|
||||
|
||||
# 最后尝试从权重中解析分类头大小
|
||||
for key, value in state_dict.items():
|
||||
if key.endswith('head.weight') or key.endswith('head.l.weight'):
|
||||
return value.shape[0]
|
||||
|
||||
raise ValueError('无法推断 num_classes,请提供 CSV 映射文件或显式指定 num_classes 参数。')
|
||||
|
||||
|
||||
def resolve_feature_dim(feature_dim_arg: Optional[int], state_dict) -> int:
|
||||
"""根据参数或 checkpoint 推断特征维度"""
|
||||
if feature_dim_arg is not None:
|
||||
return feature_dim_arg
|
||||
|
||||
# 尝试从权重中解析特征维度
|
||||
# 查找head.bn.weight的维度,这通常是特征维度
|
||||
for key, value in state_dict.items():
|
||||
if key.endswith('head.bn.weight'):
|
||||
return value.shape[0]
|
||||
|
||||
# 如果找不到,尝试查找其他可能的特征维度指示器
|
||||
for key, value in state_dict.items():
|
||||
if 'head' in key and 'weight' in key and len(value.shape) >= 2:
|
||||
# 对于线性层,输入维度通常是特征维度
|
||||
return value.shape[1] if len(value.shape) > 1 else value.shape[0]
|
||||
|
||||
# 默认值
|
||||
print("[Warning] 无法从checkpoint推断特征维度,使用默认值384")
|
||||
return 384
|
||||
|
||||
|
||||
def load_model(args, state_dict):
|
||||
"""加载模型"""
|
||||
print(f"Loading model: {args.model}")
|
||||
state_dict = normalize_state_dict_keys(state_dict)
|
||||
|
||||
model = create_model(
|
||||
args.model,
|
||||
pretrained=False,
|
||||
num_classes=args.num_classes,
|
||||
feature_dim=args.feature_dim,
|
||||
)
|
||||
|
||||
model_state = model.state_dict()
|
||||
adapted_state = {}
|
||||
mismatched = {}
|
||||
|
||||
for key, value in state_dict.items():
|
||||
if key in model_state:
|
||||
if model_state[key].shape != value.shape:
|
||||
mismatched[key] = (model_state[key].shape, value.shape)
|
||||
continue
|
||||
adapted_state[key] = value
|
||||
|
||||
classifier_keys = [key for key in mismatched if 'head' in key or 'classifier' in key]
|
||||
other_mismatched = {key: shapes for key, shapes in mismatched.items() if key not in classifier_keys}
|
||||
|
||||
if other_mismatched:
|
||||
details = '\n'.join([
|
||||
f" - {key}: checkpoint {found} -> model {expected}"
|
||||
for key, (expected, found) in other_mismatched.items()
|
||||
])
|
||||
raise RuntimeError(
|
||||
"以下权重尺寸与当前模型不兼容,且无法自动处理,请检查 checkpoint 或模型配置:\n" + details
|
||||
)
|
||||
|
||||
require_strict_head = (args.mode in ['classify', 'both']) and not args.allow_head_reinit
|
||||
|
||||
if classifier_keys and require_strict_head:
|
||||
details = '\n'.join([
|
||||
f" - {key}: checkpoint {mismatched[key][1]} -> model {mismatched[key][0]}"
|
||||
for key in classifier_keys
|
||||
])
|
||||
raise RuntimeError(
|
||||
"分类模式下检测到 checkpoint 分类头与当前 num_classes 不一致,已终止加载以避免随机初始化结果。\n"
|
||||
"请使用与训练数据一致的 checkpoint,或在确认需要重新初始化分类头时添加 --allow-head-reinit,"
|
||||
"或者切换到 --mode cluster 仅提取特征。\n" + details
|
||||
)
|
||||
|
||||
if classifier_keys:
|
||||
print("[Warning] 分类头权重尺寸与当前 num_classes 不一致,将重新初始化以下权重:")
|
||||
for key in classifier_keys:
|
||||
expected, found = mismatched[key]
|
||||
print(f" - {key}: checkpoint {found} -> model {expected}")
|
||||
# 冲突键已在上文过滤掉,无需额外处理
|
||||
|
||||
load_result = model.load_state_dict(adapted_state, strict=False)
|
||||
|
||||
if load_result.missing_keys:
|
||||
print(f"[Info] Missing keys during load: {load_result.missing_keys}")
|
||||
if load_result.unexpected_keys:
|
||||
print(f"[Info] Unexpected keys ignored: {load_result.unexpected_keys}")
|
||||
|
||||
if classifier_keys and args.mode in ['classify', 'both']:
|
||||
if args.allow_head_reinit:
|
||||
print("[Warning] 分类模式在 --allow-head-reinit 下运行,分类头为随机初始化;为获得可靠结果请提供匹配的数据集 checkpoint。")
|
||||
else:
|
||||
print("[Info] 分类模式未开启或无需分类头,已忽略冲突的分类权重。")
|
||||
|
||||
model.to(args.device)
|
||||
model.eval()
|
||||
|
||||
print(f"Model loaded from {args.checkpoint}")
|
||||
return model
|
||||
|
||||
|
||||
def load_class_mapping(class_csv_path: Optional[str]) -> Optional[Dict[int, str]]:
|
||||
"""加载 CSV 类别映射,返回 class_id -> name 的字典"""
|
||||
if not class_csv_path:
|
||||
return None
|
||||
|
||||
csv_path = Path(class_csv_path)
|
||||
if not csv_path.exists():
|
||||
raise FileNotFoundError(f"Class mapping CSV not found: {csv_path}")
|
||||
|
||||
with csv_path.open('r', encoding='utf-8-sig') as f:
|
||||
reader = csv.DictReader(f)
|
||||
if not reader.fieldnames or 'class_id' not in reader.fieldnames or 'class_name' not in reader.fieldnames:
|
||||
raise ValueError('CSV 必须包含 class_id 和 class_name 两列。')
|
||||
|
||||
mapping: Dict[int, str] = {}
|
||||
for row in reader:
|
||||
class_id = int(row['class_id'])
|
||||
class_name = row['class_name']
|
||||
mapping[class_id] = class_name
|
||||
|
||||
if not mapping:
|
||||
raise ValueError(f"CSV {csv_path} 中未找到任何类别映射。")
|
||||
|
||||
return mapping
|
||||
|
||||
|
||||
def preprocess_image(image_path, transform):
|
||||
"""预处理单张图像"""
|
||||
image = Image.open(image_path).convert('RGB')
|
||||
tensor = transform(image)
|
||||
return tensor.unsqueeze(0)
|
||||
|
||||
|
||||
def classify_image(model, image_tensor, device, class_mapping: Optional[Dict[int, str]] = None, top_k=5, threshold=0.0):
|
||||
"""对图像进行分类"""
|
||||
with torch.no_grad():
|
||||
image_tensor = image_tensor.to(device)
|
||||
# 使用分类头
|
||||
logits = model(image_tensor, return_features=False)
|
||||
|
||||
# 计算概率
|
||||
probs = F.softmax(logits, dim=-1)
|
||||
|
||||
# Top-K结果
|
||||
top_probs, top_indices = torch.topk(probs, k=min(top_k, probs.size(-1)), dim=-1)
|
||||
|
||||
results = []
|
||||
for prob, idx in zip(top_probs[0].cpu().numpy(), top_indices[0].cpu().numpy()):
|
||||
if prob >= threshold:
|
||||
class_name = class_mapping.get(int(idx), f"Class {idx}") if class_mapping else f"Class {idx}"
|
||||
results.append({
|
||||
'class_id': int(idx),
|
||||
'class_name': class_name,
|
||||
'probability': float(prob)
|
||||
})
|
||||
|
||||
# 如果过滤后结果少于 top_k,保持原样;否则取前 top_k
|
||||
if len(results) > top_k:
|
||||
results = results[:top_k]
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def extract_features(model, image_tensor, device):
|
||||
"""提取特征向量用于聚类"""
|
||||
with torch.no_grad():
|
||||
image_tensor = image_tensor.to(device)
|
||||
# 不使用分类头,直接返回特征
|
||||
features = model(image_tensor, return_features=True)
|
||||
return features.cpu().numpy()
|
||||
|
||||
|
||||
def process_single_image(args, model, transform, class_mapping: Optional[Dict[int, str]] = None):
|
||||
"""处理单张图像"""
|
||||
image_path = Path(args.input)
|
||||
if not image_path.exists():
|
||||
print(f"Error: Image not found: {image_path}")
|
||||
return
|
||||
|
||||
print(f"\nProcessing: {image_path.name}")
|
||||
|
||||
# 预处理
|
||||
image_tensor = preprocess_image(image_path, transform)
|
||||
|
||||
results = {}
|
||||
|
||||
# 分类模式
|
||||
if args.mode in ['classify', 'both']:
|
||||
print("\n[Classification Results]")
|
||||
classification = classify_image(model, image_tensor, args.device, class_mapping, args.top_k, args.threshold)
|
||||
results['classification'] = classification
|
||||
|
||||
for i, result in enumerate(classification, 1):
|
||||
print(f"{i}. {result['class_name']}: {result['probability']:.4f}")
|
||||
|
||||
# 聚类模式(提取特征)
|
||||
if args.mode in ['cluster', 'both']:
|
||||
print("\n[Feature Extraction]")
|
||||
features = extract_features(model, image_tensor, args.device)
|
||||
results['features'] = features[0].tolist()
|
||||
print(f"Feature vector shape: {features.shape}")
|
||||
print(f"Feature vector (first 10 dims): {features[0][:10]}")
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def process_directory(args, model, transform, class_mapping: Optional[Dict[int, str]] = None):
|
||||
"""批量处理目录中的图像"""
|
||||
input_dir = Path(args.input)
|
||||
if not input_dir.is_dir():
|
||||
print(f"Error: Directory not found: {input_dir}")
|
||||
return
|
||||
|
||||
# 支持的图像格式
|
||||
image_extensions = {'.jpg', '.jpeg', '.png', '.bmp', '.tiff', '.webp'}
|
||||
image_paths = [p for p in input_dir.glob('**/*') if p.suffix.lower() in image_extensions]
|
||||
|
||||
if not image_paths:
|
||||
print(f"No images found in {input_dir}")
|
||||
return
|
||||
|
||||
print(f"Found {len(image_paths)} images")
|
||||
|
||||
all_results = {}
|
||||
|
||||
# 批量处理
|
||||
for i in range(0, len(image_paths), args.batch_size):
|
||||
batch_paths = image_paths[i:i + args.batch_size]
|
||||
|
||||
# 预处理批次
|
||||
batch_tensors = []
|
||||
for path in batch_paths:
|
||||
try:
|
||||
tensor = preprocess_image(path, transform)
|
||||
batch_tensors.append(tensor)
|
||||
except Exception as e:
|
||||
print(f"Error processing {path.name}: {e}")
|
||||
continue
|
||||
|
||||
if not batch_tensors:
|
||||
continue
|
||||
|
||||
batch_tensor = torch.cat(batch_tensors, dim=0)
|
||||
|
||||
# 推理
|
||||
with torch.no_grad():
|
||||
batch_tensor = batch_tensor.to(args.device)
|
||||
|
||||
if args.mode in ['classify', 'both']:
|
||||
logits = model(batch_tensor, return_features=False)
|
||||
probs = F.softmax(logits, dim=-1)
|
||||
top_probs, top_indices = torch.topk(probs, k=min(args.top_k, probs.size(-1)), dim=-1)
|
||||
|
||||
if args.mode in ['cluster', 'both']:
|
||||
features = model(batch_tensor, return_features=True)
|
||||
|
||||
# 保存结果
|
||||
for j, path in enumerate(batch_paths):
|
||||
if j >= len(batch_tensors):
|
||||
continue
|
||||
|
||||
result = {'image': path.name}
|
||||
|
||||
if args.mode in ['classify', 'both']:
|
||||
# 获取该图像的 top-k 结果
|
||||
img_top_probs = top_probs[j].cpu().numpy()
|
||||
img_top_indices = top_indices[j].cpu().numpy()
|
||||
|
||||
classifications = []
|
||||
for prob, idx in zip(img_top_probs, img_top_indices):
|
||||
if prob >= args.threshold:
|
||||
class_id = int(idx)
|
||||
class_name = class_mapping.get(class_id, f"Class {class_id}") if class_mapping else f"Class {class_id}"
|
||||
classifications.append({
|
||||
'class_id': class_id,
|
||||
'class_name': class_name,
|
||||
'probability': float(prob)
|
||||
})
|
||||
|
||||
# 如果需要,取前 top_k
|
||||
if len(classifications) > args.top_k:
|
||||
classifications = classifications[:args.top_k]
|
||||
|
||||
result['classification'] = classifications
|
||||
|
||||
if args.mode in ['cluster', 'both']:
|
||||
result['features'] = features[j].cpu().numpy().tolist()
|
||||
|
||||
all_results[path.name] = result
|
||||
|
||||
print(f"Processed {min(i + args.batch_size, len(image_paths))}/{len(image_paths)} images")
|
||||
|
||||
return all_results
|
||||
|
||||
|
||||
def main(args):
|
||||
# 创建输出目录
|
||||
output_dir = Path(args.output)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if args.mode in ['classify', 'both'] and not args.class_csv:
|
||||
raise ValueError('分类或混合模式下必须提供 --class-csv,且需使用训练阶段导出的映射文件。')
|
||||
|
||||
# 加载类别映射
|
||||
class_mapping = load_class_mapping(args.class_csv)
|
||||
|
||||
# 加载 checkpoint 并解析类别数
|
||||
state_dict = load_checkpoint_state(args.checkpoint)
|
||||
state_dict = normalize_state_dict_keys(state_dict)
|
||||
args.num_classes = resolve_num_classes(args.num_classes, class_mapping, state_dict)
|
||||
args.feature_dim = resolve_feature_dim(args.feature_dim, state_dict)
|
||||
|
||||
# 加载模型
|
||||
model = load_model(args, state_dict)
|
||||
|
||||
# 创建数据转换
|
||||
config = resolve_data_config({}, model=model)
|
||||
transform = create_transform(**config)
|
||||
|
||||
# 判断输入类型
|
||||
input_path = Path(args.input)
|
||||
|
||||
if input_path.is_file():
|
||||
# 单张图像
|
||||
results = process_single_image(args, model, transform, class_mapping)
|
||||
|
||||
# 保存结果
|
||||
output_file = output_dir / f"{input_path.stem}_result.json"
|
||||
with open(output_file, 'w', encoding='utf-8') as f:
|
||||
json.dump(results, f, indent=2, ensure_ascii=False)
|
||||
print(f"\nResults saved to: {output_file}")
|
||||
|
||||
elif input_path.is_dir():
|
||||
# 目录批量处理
|
||||
results = process_directory(args, model, transform, class_mapping)
|
||||
|
||||
# 保存结果
|
||||
output_file = output_dir / "batch_results.json"
|
||||
with open(output_file, 'w', encoding='utf-8') as f:
|
||||
json.dump(results, f, indent=2, ensure_ascii=False)
|
||||
print(f"\nResults saved to: {output_file}")
|
||||
|
||||
# 如果是聚类模式,额外保存特征矩阵
|
||||
if args.mode in ['cluster', 'both']:
|
||||
features_list = []
|
||||
image_names = []
|
||||
for name, result in results.items():
|
||||
if 'features' in result:
|
||||
features_list.append(result['features'])
|
||||
image_names.append(name)
|
||||
|
||||
if features_list:
|
||||
features_array = np.array(features_list)
|
||||
np.save(output_dir / "features.npy", features_array)
|
||||
with open(output_dir / "image_names.txt", 'w') as f:
|
||||
f.write('\n'.join(image_names))
|
||||
print(f"Feature matrix saved: {output_dir / 'features.npy'}")
|
||||
print(f"Feature matrix shape: {features_array.shape}")
|
||||
else:
|
||||
print(f"Error: Invalid input path: {input_path}")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser('Artist Style Inference', parents=[get_args_parser()])
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,3 @@
|
||||
from .lsnet import *
|
||||
from .lsnet_artist import *
|
||||
from .build import *
|
||||
@@ -0,0 +1 @@
|
||||
import model.lsnet
|
||||
+405
@@ -0,0 +1,405 @@
|
||||
import torch
|
||||
import itertools
|
||||
|
||||
from timm.models.vision_transformer import trunc_normal_
|
||||
from timm.layers import SqueezeExcite
|
||||
from timm.models import register_model
|
||||
from .ska import SKA
|
||||
|
||||
from timm.models import build_model_with_cfg
|
||||
from timm.data import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD
|
||||
|
||||
class Conv2d_BN(torch.nn.Sequential):
|
||||
def __init__(self, a, b, ks=1, stride=1, pad=0, dilation=1,
|
||||
groups=1, bn_weight_init=1):
|
||||
super().__init__()
|
||||
self.add_module('c', torch.nn.Conv2d(
|
||||
a, b, ks, stride, pad, dilation, groups, bias=False))
|
||||
self.add_module('bn', torch.nn.BatchNorm2d(b))
|
||||
torch.nn.init.constant_(self.bn.weight, bn_weight_init)
|
||||
torch.nn.init.constant_(self.bn.bias, 0)
|
||||
|
||||
@torch.no_grad()
|
||||
def fuse(self):
|
||||
c, bn = self._modules.values()
|
||||
w = bn.weight / (bn.running_var + bn.eps)**0.5
|
||||
w = c.weight * w[:, None, None, None]
|
||||
b = bn.bias - bn.running_mean * bn.weight / \
|
||||
(bn.running_var + bn.eps)**0.5
|
||||
m = torch.nn.Conv2d(w.size(1) * self.c.groups, w.size(
|
||||
0), w.shape[2:], stride=self.c.stride, padding=self.c.padding, dilation=self.c.dilation, groups=self.c.groups,
|
||||
device=c.weight.device)
|
||||
m.weight.data.copy_(w)
|
||||
m.bias.data.copy_(b)
|
||||
return m
|
||||
|
||||
|
||||
class BN_Linear(torch.nn.Sequential):
|
||||
def __init__(self, a, b, bias=True, std=0.02):
|
||||
super().__init__()
|
||||
self.add_module('bn', torch.nn.BatchNorm1d(a))
|
||||
self.add_module('l', torch.nn.Linear(a, b, bias=bias))
|
||||
trunc_normal_(self.l.weight, std=std)
|
||||
if bias:
|
||||
torch.nn.init.constant_(self.l.bias, 0)
|
||||
|
||||
@torch.no_grad()
|
||||
def fuse(self):
|
||||
bn, l = self._modules.values()
|
||||
w = bn.weight / (bn.running_var + bn.eps)**0.5
|
||||
b = bn.bias - self.bn.running_mean * \
|
||||
self.bn.weight / (bn.running_var + bn.eps)**0.5
|
||||
w = l.weight * w[None, :]
|
||||
if l.bias is None:
|
||||
b = b @ self.l.weight.T
|
||||
else:
|
||||
b = (l.weight @ b[:, None]).view(-1) + self.l.bias
|
||||
m = torch.nn.Linear(w.size(1), w.size(0), device=l.weight.device)
|
||||
m.weight.data.copy_(w)
|
||||
m.bias.data.copy_(b)
|
||||
return m
|
||||
|
||||
class Residual(torch.nn.Module):
|
||||
def __init__(self, m, drop=0.):
|
||||
super().__init__()
|
||||
self.m = m
|
||||
self.drop = drop
|
||||
|
||||
def forward(self, x):
|
||||
if self.training and self.drop > 0:
|
||||
return x + self.m(x) * torch.rand(x.size(0), 1, 1, 1,
|
||||
device=x.device).ge_(self.drop).div(1 - self.drop).detach()
|
||||
else:
|
||||
return x + self.m(x)
|
||||
|
||||
class FFN(torch.nn.Module):
|
||||
def __init__(self, ed, h):
|
||||
super().__init__()
|
||||
self.pw1 = Conv2d_BN(ed, h)
|
||||
self.act = torch.nn.ReLU()
|
||||
self.pw2 = Conv2d_BN(h, ed, bn_weight_init=0)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.pw2(self.act(self.pw1(x)))
|
||||
return x
|
||||
|
||||
class Attention(torch.nn.Module):
|
||||
def __init__(self, dim, key_dim, num_heads=8,
|
||||
attn_ratio=4,
|
||||
resolution=14):
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
self.scale = key_dim ** -0.5
|
||||
self.key_dim = key_dim
|
||||
self.nh_kd = nh_kd = key_dim * num_heads
|
||||
self.d = int(attn_ratio * key_dim)
|
||||
self.dh = int(attn_ratio * key_dim) * num_heads
|
||||
self.attn_ratio = attn_ratio
|
||||
h = self.dh + nh_kd * 2
|
||||
self.qkv = Conv2d_BN(dim, h, ks=1)
|
||||
self.proj = torch.nn.Sequential(torch.nn.ReLU(), Conv2d_BN(
|
||||
self.dh, dim, bn_weight_init=0))
|
||||
self.dw = Conv2d_BN(nh_kd, nh_kd, 3, 1, 1, groups=nh_kd)
|
||||
points = list(itertools.product(range(resolution), range(resolution)))
|
||||
N = len(points)
|
||||
attention_offsets = {}
|
||||
idxs = []
|
||||
for p1 in points:
|
||||
for p2 in points:
|
||||
offset = (abs(p1[0] - p2[0]), abs(p1[1] - p2[1]))
|
||||
if offset not in attention_offsets:
|
||||
attention_offsets[offset] = len(attention_offsets)
|
||||
idxs.append(attention_offsets[offset])
|
||||
self.attention_biases = torch.nn.Parameter(
|
||||
torch.zeros(num_heads, len(attention_offsets)))
|
||||
self.register_buffer('attention_bias_idxs',
|
||||
torch.LongTensor(idxs).view(N, N))
|
||||
|
||||
@torch.no_grad()
|
||||
def train(self, mode=True):
|
||||
super().train(mode)
|
||||
if mode and hasattr(self, 'ab'):
|
||||
del self.ab
|
||||
else:
|
||||
self.ab = self.attention_biases[:, self.attention_bias_idxs]
|
||||
|
||||
def forward(self, x):
|
||||
B, _, H, W = x.shape
|
||||
N = H * W
|
||||
qkv = self.qkv(x)
|
||||
q, k, v = qkv.view(B, -1, H, W).split([self.nh_kd, self.nh_kd, self.dh], dim=1)
|
||||
q = self.dw(q)
|
||||
q, k, v = q.view(B, self.num_heads, -1, N), k.view(B, self.num_heads, -1, N), v.view(B, self.num_heads, -1, N)
|
||||
attn = (
|
||||
(q.transpose(-2, -1) @ k) * self.scale
|
||||
+
|
||||
(self.attention_biases[:, self.attention_bias_idxs]
|
||||
if self.training else self.ab)
|
||||
)
|
||||
attn = attn.softmax(dim=-1)
|
||||
x = (v @ attn.transpose(-2, -1)).reshape(B, -1, H, W)
|
||||
x = self.proj(x)
|
||||
return x
|
||||
|
||||
class RepVGGDW(torch.nn.Module):
|
||||
def __init__(self, ed) -> None:
|
||||
super().__init__()
|
||||
self.conv = Conv2d_BN(ed, ed, 3, 1, 1, groups=ed)
|
||||
self.conv1 = Conv2d_BN(ed, ed, 1, 1, 0, groups=ed)
|
||||
self.dim = ed
|
||||
|
||||
def forward(self, x):
|
||||
return self.conv(x) + self.conv1(x) + x
|
||||
|
||||
@torch.no_grad()
|
||||
def fuse(self):
|
||||
conv = self.conv.fuse()
|
||||
conv1 = self.conv1.fuse()
|
||||
|
||||
conv_w = conv.weight
|
||||
conv_b = conv.bias
|
||||
conv1_w = conv1.weight
|
||||
conv1_b = conv1.bias
|
||||
|
||||
conv1_w = torch.nn.functional.pad(conv1_w, [1,1,1,1])
|
||||
|
||||
identity = torch.nn.functional.pad(torch.ones(conv1_w.shape[0], conv1_w.shape[1], 1, 1, device=conv1_w.device), [1,1,1,1])
|
||||
|
||||
final_conv_w = conv_w + conv1_w + identity
|
||||
final_conv_b = conv_b + conv1_b
|
||||
|
||||
conv.weight.data.copy_(final_conv_w)
|
||||
conv.bias.data.copy_(final_conv_b)
|
||||
return conv
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
class LKP(nn.Module):
|
||||
def __init__(self, dim, lks, sks, groups):
|
||||
super().__init__()
|
||||
self.cv1 = Conv2d_BN(dim, dim // 2)
|
||||
self.act = nn.ReLU()
|
||||
self.cv2 = Conv2d_BN(dim // 2, dim // 2, ks=lks, pad=(lks - 1) // 2, groups=dim // 2)
|
||||
self.cv3 = Conv2d_BN(dim // 2, dim // 2)
|
||||
self.cv4 = nn.Conv2d(dim // 2, sks ** 2 * dim // groups, kernel_size=1)
|
||||
self.norm = nn.GroupNorm(num_groups=dim // groups, num_channels=sks ** 2 * dim // groups)
|
||||
|
||||
self.sks = sks
|
||||
self.groups = groups
|
||||
self.dim = dim
|
||||
|
||||
def forward(self, x):
|
||||
x = self.act(self.cv3(self.cv2(self.act(self.cv1(x)))))
|
||||
w = self.norm(self.cv4(x))
|
||||
b, _, h, width = w.size()
|
||||
w = w.view(b, self.dim // self.groups, self.sks ** 2, h, width)
|
||||
return w
|
||||
|
||||
class LSConv(nn.Module):
|
||||
def __init__(self, dim):
|
||||
super(LSConv, self).__init__()
|
||||
self.lkp = LKP(dim, lks=7, sks=3, groups=8)
|
||||
self.ska = SKA()
|
||||
self.bn = nn.BatchNorm2d(dim)
|
||||
|
||||
def forward(self, x):
|
||||
return self.bn(self.ska(x, self.lkp(x))) + x
|
||||
|
||||
class Block(torch.nn.Module):
|
||||
def __init__(self,
|
||||
ed, kd, nh=8,
|
||||
ar=4,
|
||||
resolution=14,
|
||||
stage=-1, depth=-1):
|
||||
super().__init__()
|
||||
|
||||
if depth % 2 == 0:
|
||||
self.mixer = RepVGGDW(ed)
|
||||
self.se = SqueezeExcite(ed, 0.25)
|
||||
else:
|
||||
self.se = torch.nn.Identity()
|
||||
if stage == 3:
|
||||
self.mixer = Residual(Attention(ed, kd, nh, ar, resolution=resolution))
|
||||
else:
|
||||
self.mixer = LSConv(ed)
|
||||
|
||||
self.ffn = Residual(FFN(ed, int(ed * 2)))
|
||||
|
||||
def forward(self, x):
|
||||
return self.ffn(self.se(self.mixer(x)))
|
||||
|
||||
class LSNet(torch.nn.Module):
|
||||
def __init__(self, img_size=224,
|
||||
patch_size=16,
|
||||
in_chans=3,
|
||||
num_classes=1000,
|
||||
embed_dim=[64, 128, 192, 256],
|
||||
key_dim=[16, 16, 16, 16],
|
||||
depth=[1, 2, 3, 4],
|
||||
num_heads=[4, 4, 4, 4],
|
||||
distillation=False,
|
||||
**kwargs):
|
||||
super().__init__()
|
||||
|
||||
default_cfg = kwargs.pop('default_cfg', None)
|
||||
pretrained_cfg = kwargs.pop('pretrained_cfg', None)
|
||||
pretrained_cfg_overlay = kwargs.pop('pretrained_cfg_overlay', None)
|
||||
|
||||
if default_cfg is not None:
|
||||
self.default_cfg = default_cfg
|
||||
if pretrained_cfg is not None:
|
||||
self.pretrained_cfg = pretrained_cfg
|
||||
if pretrained_cfg_overlay is not None:
|
||||
self.pretrained_cfg_overlay = pretrained_cfg_overlay
|
||||
|
||||
if kwargs:
|
||||
self.extra_init_kwargs = kwargs
|
||||
|
||||
resolution = img_size
|
||||
self.patch_embed = torch.nn.Sequential(Conv2d_BN(in_chans, embed_dim[0] // 4, 3, 2, 1), torch.nn.ReLU(),
|
||||
Conv2d_BN(embed_dim[0] // 4, embed_dim[0] // 2, 3, 2, 1), torch.nn.ReLU(),
|
||||
Conv2d_BN(embed_dim[0] // 2, embed_dim[0], 3, 2, 1)
|
||||
)
|
||||
|
||||
resolution = img_size // patch_size
|
||||
attn_ratio = [embed_dim[i] / (key_dim[i] * num_heads[i]) for i in range(len(embed_dim))]
|
||||
self.blocks1 = nn.Sequential()
|
||||
self.blocks2 = nn.Sequential()
|
||||
self.blocks3 = nn.Sequential()
|
||||
self.blocks4 = nn.Sequential()
|
||||
blocks = [self.blocks1, self.blocks2, self.blocks3, self.blocks4]
|
||||
|
||||
for i, (ed, kd, dpth, nh, ar) in enumerate(
|
||||
zip(embed_dim, key_dim, depth, num_heads, attn_ratio)):
|
||||
for d in range(dpth):
|
||||
blocks[i].append(Block(ed, kd, nh, ar, resolution, stage=i, depth=d))
|
||||
|
||||
if i != len(depth) - 1:
|
||||
blk = blocks[i+1]
|
||||
resolution_ = (resolution - 1) // 2 + 1
|
||||
blk.append(Conv2d_BN(embed_dim[i], embed_dim[i], ks=3, stride=2, pad=1, groups=embed_dim[i]))
|
||||
blk.append(Conv2d_BN(embed_dim[i], embed_dim[i+1], ks=1, stride=1, pad=0))
|
||||
resolution = resolution_
|
||||
|
||||
self.head = BN_Linear(embed_dim[-1], num_classes) if num_classes > 0 else torch.nn.Identity()
|
||||
self.distillation = distillation
|
||||
if distillation:
|
||||
self.head_dist = BN_Linear(embed_dim[-1], num_classes) if num_classes > 0 else torch.nn.Identity()
|
||||
|
||||
self.num_classes = num_classes
|
||||
self.num_features = embed_dim[-1]
|
||||
|
||||
@torch.jit.ignore # type: ignore
|
||||
def no_weight_decay(self):
|
||||
return {x for x in self.state_dict().keys() if 'attention_biases' in x}
|
||||
|
||||
def forward(self, x):
|
||||
x = self.patch_embed(x)
|
||||
x = self.blocks1(x)
|
||||
x = self.blocks2(x)
|
||||
x = self.blocks3(x)
|
||||
x = self.blocks4(x)
|
||||
x = torch.nn.functional.adaptive_avg_pool2d(x, 1).flatten(1)
|
||||
if self.distillation:
|
||||
x = self.head(x), self.head_dist(x)
|
||||
if not self.training:
|
||||
x = (x[0] + x[1]) / 2
|
||||
else:
|
||||
x = self.head(x)
|
||||
return x
|
||||
|
||||
def _cfg(url='', **kwargs):
|
||||
return {
|
||||
'url': url,
|
||||
'num_classes': 1000, 'input_size': (3, 224, 224), 'pool_size': (4, 4),
|
||||
'crop_pct': .9, 'interpolation': 'bicubic',
|
||||
'mean': IMAGENET_DEFAULT_MEAN, 'std': IMAGENET_DEFAULT_STD,
|
||||
'first_conv': 'patch_embed.0.c', 'classifier': ('head.linear', 'head_dist.linear'),
|
||||
**kwargs
|
||||
}
|
||||
|
||||
def _with_hf_hub(kwargs):
|
||||
"""兼容不同 timm 版本的 hf hub 配置字段"""
|
||||
if 'hf_hub' in kwargs and 'hf_hub_id' not in kwargs:
|
||||
kwargs['hf_hub_id'] = kwargs.pop('hf_hub')
|
||||
return kwargs
|
||||
|
||||
|
||||
default_cfgs = dict(
|
||||
lsnet_t=_cfg(**_with_hf_hub({'hf_hub': 'jameslahm/lsnet_t'})),
|
||||
lsnet_t_distill=_cfg(**_with_hf_hub({'hf_hub': 'jameslahm/lsnet_t_distill'})),
|
||||
lsnet_s=_cfg(**_with_hf_hub({'hf_hub': 'jameslahm/lsnet_s'})),
|
||||
lsnet_s_distill=_cfg(**_with_hf_hub({'hf_hub': 'jameslahm/lsnet_s_distill'})),
|
||||
lsnet_b=_cfg(**_with_hf_hub({'hf_hub': 'jameslahm/lsnet_b'})),
|
||||
lsnet_b_distill=_cfg(**_with_hf_hub({'hf_hub': 'jameslahm/lsnet_b_distill'})),
|
||||
)
|
||||
|
||||
def _create_lsnet(variant, pretrained=False, **kwargs):
|
||||
cfg = default_cfgs.get(variant, None)
|
||||
if cfg is not None:
|
||||
kwargs.setdefault('default_cfg', cfg)
|
||||
kwargs.setdefault('pretrained_cfg', cfg)
|
||||
model = build_model_with_cfg(
|
||||
LSNet,
|
||||
variant,
|
||||
pretrained,
|
||||
**kwargs,
|
||||
)
|
||||
return model
|
||||
|
||||
@register_model
|
||||
def lsnet_t(num_classes=1000, distillation=False, pretrained=False, **kwargs):
|
||||
model = _create_lsnet("lsnet_t" + ("_distill" if distillation else ""),
|
||||
pretrained=pretrained,
|
||||
num_classes=num_classes,
|
||||
distillation=distillation,
|
||||
img_size=224,
|
||||
patch_size=8,
|
||||
embed_dim=[64, 128, 256, 384],
|
||||
depth=[0, 2, 8, 10],
|
||||
num_heads=[3, 3, 3, 4],
|
||||
)
|
||||
return model
|
||||
|
||||
@register_model
|
||||
def lsnet_s(num_classes=1000, distillation=False, pretrained=False, **kwargs):
|
||||
model = _create_lsnet("lsnet_s" + ("_distill" if distillation else ""),
|
||||
pretrained=pretrained,
|
||||
num_classes=num_classes,
|
||||
distillation=distillation,
|
||||
img_size=224,
|
||||
patch_size=8,
|
||||
embed_dim=[96, 192, 320, 448],
|
||||
depth=[1, 2, 8, 10],
|
||||
num_heads=[3, 3, 3, 4],
|
||||
)
|
||||
return model
|
||||
|
||||
@register_model
|
||||
def lsnet_b(num_classes=1000, distillation=False, pretrained=False, **kwargs):
|
||||
model = _create_lsnet("lsnet_b" + ("_distill" if distillation else ""),
|
||||
pretrained=pretrained,
|
||||
num_classes=num_classes,
|
||||
distillation=distillation,
|
||||
img_size=224,
|
||||
patch_size=8,
|
||||
embed_dim=[128, 256, 384, 512],
|
||||
depth=[4, 6, 8, 10],
|
||||
num_heads=[3, 3, 3, 4],
|
||||
)
|
||||
return model
|
||||
|
||||
@register_model
|
||||
def lsnet_t_distill(**kwargs):
|
||||
kwargs["distillation"] = True
|
||||
return lsnet_t(**kwargs)
|
||||
|
||||
@register_model
|
||||
def lsnet_s_distill(**kwargs):
|
||||
kwargs["distillation"] = True
|
||||
return lsnet_s(**kwargs)
|
||||
|
||||
@register_model
|
||||
def lsnet_b_distill(**kwargs):
|
||||
kwargs["distillation"] = True
|
||||
return lsnet_b(**kwargs)
|
||||
@@ -0,0 +1,271 @@
|
||||
"""
|
||||
LSNet for Artist Style Classification and Clustering
|
||||
支持画师风格的分类和聚类任务
|
||||
"""
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from .lsnet import LSNet, Conv2d_BN, BN_Linear
|
||||
from timm.models import register_model
|
||||
from timm.models import build_model_with_cfg
|
||||
|
||||
|
||||
class LSNetArtist(LSNet):
|
||||
"""
|
||||
LSNet模型用于画师风格分类和聚类
|
||||
|
||||
特点:
|
||||
- 训练时使用分类头进行监督学习
|
||||
- 推理时可选择是否使用分类头
|
||||
- 去掉分类头输出特征向量用于聚类
|
||||
- 保留分类头可以做风格分类
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
img_size=224,
|
||||
patch_size=8,
|
||||
in_chans=3,
|
||||
num_classes=1000,
|
||||
embed_dim=[64, 128, 256, 384],
|
||||
key_dim=[16, 16, 16, 16],
|
||||
depth=[0, 2, 8, 10],
|
||||
num_heads=[3, 3, 3, 4],
|
||||
distillation=False,
|
||||
feature_dim=None, # 特征向量维度,默认为embed_dim[-1]
|
||||
use_projection=True, # 是否使用projection层
|
||||
**kwargs):
|
||||
default_cfg = kwargs.pop('default_cfg', None)
|
||||
pretrained_cfg = kwargs.pop('pretrained_cfg', None)
|
||||
pretrained_cfg_overlay = kwargs.pop('pretrained_cfg_overlay', None)
|
||||
|
||||
super().__init__(
|
||||
img_size=img_size,
|
||||
patch_size=patch_size,
|
||||
in_chans=in_chans,
|
||||
num_classes=num_classes,
|
||||
embed_dim=embed_dim,
|
||||
key_dim=key_dim,
|
||||
depth=depth,
|
||||
num_heads=num_heads,
|
||||
distillation=distillation,
|
||||
default_cfg=default_cfg,
|
||||
pretrained_cfg=pretrained_cfg,
|
||||
pretrained_cfg_overlay=pretrained_cfg_overlay,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
self.feature_dim = feature_dim if feature_dim is not None else embed_dim[-1]
|
||||
self.use_projection = use_projection
|
||||
|
||||
# 如果使用projection层,添加一个映射层来生成固定维度的特征
|
||||
if self.use_projection and self.feature_dim != embed_dim[-1]:
|
||||
self.projection = nn.Sequential(
|
||||
BN_Linear(embed_dim[-1], self.feature_dim),
|
||||
nn.ReLU(),
|
||||
)
|
||||
else:
|
||||
self.projection = nn.Identity()
|
||||
|
||||
# 重新定义分类头(基于特征维度)
|
||||
if num_classes > 0:
|
||||
self.head = BN_Linear(self.feature_dim, num_classes)
|
||||
if distillation:
|
||||
self.head_dist = BN_Linear(self.feature_dim, num_classes)
|
||||
|
||||
def forward_features(self, x):
|
||||
"""
|
||||
提取特征,不经过分类头
|
||||
用于聚类或特征提取
|
||||
"""
|
||||
x = self.patch_embed(x)
|
||||
x = self.blocks1(x)
|
||||
x = self.blocks2(x)
|
||||
x = self.blocks3(x)
|
||||
x = self.blocks4(x)
|
||||
x = torch.nn.functional.adaptive_avg_pool2d(x, 1).flatten(1)
|
||||
x = self.projection(x)
|
||||
return x
|
||||
|
||||
def forward(self, x, return_features=False):
|
||||
"""
|
||||
前向传播
|
||||
|
||||
Args:
|
||||
x: 输入图像
|
||||
return_features: 是否只返回特征向量(用于聚类)
|
||||
False时返回分类logits(用于分类)
|
||||
|
||||
Returns:
|
||||
如果return_features=True: 返回特征向量 (batch_size, feature_dim)
|
||||
如果return_features=False: 返回分类logits (batch_size, num_classes)
|
||||
"""
|
||||
features = self.forward_features(x)
|
||||
|
||||
if return_features:
|
||||
# 返回特征向量用于聚类
|
||||
return features
|
||||
|
||||
# 返回分类结果
|
||||
if self.distillation:
|
||||
x = self.head(features), self.head_dist(features)
|
||||
if not self.training:
|
||||
x = (x[0] + x[1]) / 2
|
||||
else:
|
||||
x = self.head(features)
|
||||
|
||||
return x
|
||||
|
||||
def get_features(self, x):
|
||||
"""
|
||||
便捷方法:提取特征向量
|
||||
"""
|
||||
return self.forward(x, return_features=True)
|
||||
|
||||
def classify(self, x):
|
||||
"""
|
||||
便捷方法:进行分类
|
||||
"""
|
||||
return self.forward(x, return_features=False)
|
||||
|
||||
|
||||
def _cfg_artist(url='', **kwargs):
|
||||
return {
|
||||
'url': url,
|
||||
'num_classes': 1000,
|
||||
'input_size': (3, 224, 224),
|
||||
'pool_size': (4, 4),
|
||||
'crop_pct': .9,
|
||||
'interpolation': 'bicubic',
|
||||
'mean': (0.485, 0.456, 0.406),
|
||||
'std': (0.229, 0.224, 0.225),
|
||||
'first_conv': 'patch_embed.0.c',
|
||||
'classifier': ('head.linear', 'head_dist.linear'),
|
||||
**kwargs
|
||||
}
|
||||
|
||||
|
||||
default_cfgs_artist = dict(
|
||||
lsnet_t_artist = _cfg_artist(),
|
||||
lsnet_s_artist = _cfg_artist(),
|
||||
lsnet_b_artist = _cfg_artist(),
|
||||
lsnet_l_artist = _cfg_artist(), # Large model for massive training
|
||||
lsnet_xl_artist = _cfg_artist(), # Extra Large model for 100k+ classes
|
||||
)
|
||||
|
||||
|
||||
def _create_lsnet_artist(variant, pretrained=False, **kwargs):
|
||||
cfg = default_cfgs_artist.get(variant, None)
|
||||
if cfg is not None:
|
||||
kwargs.setdefault('default_cfg', cfg)
|
||||
kwargs.setdefault('pretrained_cfg', cfg)
|
||||
model = build_model_with_cfg(
|
||||
LSNetArtist,
|
||||
variant,
|
||||
pretrained,
|
||||
**kwargs,
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def lsnet_t_artist(num_classes=1000, distillation=False, pretrained=False,
|
||||
feature_dim=None, use_projection=True, **kwargs):
|
||||
"""LSNet-T for Artist Style Classification"""
|
||||
model = _create_lsnet_artist(
|
||||
"lsnet_t_artist",
|
||||
pretrained=pretrained,
|
||||
num_classes=num_classes,
|
||||
distillation=distillation,
|
||||
img_size=224,
|
||||
patch_size=8,
|
||||
embed_dim=[64, 128, 256, 384],
|
||||
depth=[0, 2, 8, 10],
|
||||
num_heads=[3, 3, 3, 4],
|
||||
feature_dim=feature_dim,
|
||||
use_projection=use_projection,
|
||||
**kwargs
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def lsnet_s_artist(num_classes=1000, distillation=False, pretrained=False,
|
||||
feature_dim=None, use_projection=True, **kwargs):
|
||||
"""LSNet-S for Artist Style Classification"""
|
||||
model = _create_lsnet_artist(
|
||||
"lsnet_s_artist",
|
||||
pretrained=pretrained,
|
||||
num_classes=num_classes,
|
||||
distillation=distillation,
|
||||
img_size=224,
|
||||
patch_size=8,
|
||||
embed_dim=[96, 192, 320, 448],
|
||||
depth=[1, 2, 8, 10],
|
||||
num_heads=[3, 3, 3, 4],
|
||||
feature_dim=feature_dim,
|
||||
use_projection=use_projection,
|
||||
**kwargs
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def lsnet_b_artist(num_classes=1000, distillation=False, pretrained=False,
|
||||
feature_dim=None, use_projection=True, **kwargs):
|
||||
"""LSNet-B for Artist Style Classification"""
|
||||
model = _create_lsnet_artist(
|
||||
"lsnet_b_artist",
|
||||
pretrained=pretrained,
|
||||
num_classes=num_classes,
|
||||
distillation=distillation,
|
||||
img_size=224,
|
||||
patch_size=8,
|
||||
embed_dim=[128, 256, 384, 512],
|
||||
depth=[4, 6, 8, 10],
|
||||
num_heads=[3, 3, 3, 4],
|
||||
feature_dim=feature_dim,
|
||||
use_projection=use_projection,
|
||||
**kwargs
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def lsnet_l_artist(num_classes=1000, distillation=False, pretrained=False,
|
||||
feature_dim=None, use_projection=True, **kwargs):
|
||||
"""LSNet-L for Artist Style Classification (Large model for massive training)"""
|
||||
model = _create_lsnet_artist(
|
||||
"lsnet_l_artist",
|
||||
pretrained=pretrained,
|
||||
num_classes=num_classes,
|
||||
distillation=distillation,
|
||||
img_size=224,
|
||||
patch_size=8,
|
||||
embed_dim=[160, 320, 480, 640], # 更大的embed_dim
|
||||
depth=[6, 8, 12, 14], # 更深的网络
|
||||
num_heads=[4, 4, 4, 4], # 更多的注意力头
|
||||
feature_dim=feature_dim,
|
||||
use_projection=use_projection,
|
||||
**kwargs
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def lsnet_xl_artist(num_classes=1000, distillation=False, pretrained=False,
|
||||
feature_dim=None, use_projection=True, **kwargs):
|
||||
"""LSNet-XL for Artist Style Classification (Extra Large model for massive datasets with 100k+ classes)"""
|
||||
model = _create_lsnet_artist(
|
||||
"lsnet_xl_artist",
|
||||
pretrained=pretrained,
|
||||
num_classes=num_classes,
|
||||
distillation=distillation,
|
||||
img_size=224,
|
||||
patch_size=8,
|
||||
embed_dim=[192, 384, 576, 768], # 超大embed_dim,支持10万+类别
|
||||
depth=[8, 12, 16, 20], # 超深网络,学习复杂特征
|
||||
num_heads=[6, 6, 6, 6], # 更多注意力头
|
||||
feature_dim=feature_dim,
|
||||
use_projection=use_projection,
|
||||
**kwargs
|
||||
)
|
||||
return model
|
||||
+168
@@ -0,0 +1,168 @@
|
||||
import torch
|
||||
from torch.autograd import Function
|
||||
import triton
|
||||
import triton.language as tl
|
||||
from torch.amp import custom_fwd, custom_bwd
|
||||
import math
|
||||
|
||||
def _grid(numel: int, bs: int) -> tuple:
|
||||
return (triton.cdiv(numel, bs),)
|
||||
|
||||
@triton.jit
|
||||
def _idx(i, n: int, c: int, h: int, w: int):
|
||||
ni = i // (c * h * w)
|
||||
ci = (i // (h * w)) % c
|
||||
hi = (i // w) % h
|
||||
wi = i % w
|
||||
m = i < (n * c * h * w)
|
||||
return ni, ci, hi, wi, m
|
||||
|
||||
@triton.jit
|
||||
def ska_fwd(
|
||||
x_ptr, w_ptr, o_ptr,
|
||||
n, ic, h, w, ks, pad, wc,
|
||||
BS: tl.constexpr,
|
||||
CT: tl.constexpr, AT: tl.constexpr
|
||||
):
|
||||
pid = tl.program_id(0)
|
||||
start = pid * BS
|
||||
offs = start + tl.arange(0, BS)
|
||||
|
||||
ni, ci, hi, wi, m = _idx(offs, n, ic, h, w)
|
||||
val = tl.zeros((BS,), dtype=AT)
|
||||
|
||||
for kh in range(ks):
|
||||
hin = hi - pad + kh
|
||||
hb = (hin >= 0) & (hin < h)
|
||||
for kw in range(ks):
|
||||
win = wi - pad + kw
|
||||
b = hb & (win >= 0) & (win < w)
|
||||
|
||||
x_off = ((ni * ic + ci) * h + hin) * w + win
|
||||
w_off = ((ni * wc + ci % wc) * ks * ks + (kh * ks + kw)) * h * w + hi * w + wi
|
||||
|
||||
x_val = tl.load(x_ptr + x_off, mask=m & b, other=0.0).to(CT)
|
||||
w_val = tl.load(w_ptr + w_off, mask=m, other=0.0).to(CT)
|
||||
val += tl.where(b & m, x_val * w_val, 0.0).to(AT)
|
||||
|
||||
tl.store(o_ptr + offs, val.to(CT), mask=m)
|
||||
|
||||
@triton.jit
|
||||
def ska_bwd_x(
|
||||
go_ptr, w_ptr, gi_ptr,
|
||||
n, ic, h, w, ks, pad, wc,
|
||||
BS: tl.constexpr,
|
||||
CT: tl.constexpr, AT: tl.constexpr
|
||||
):
|
||||
pid = tl.program_id(0)
|
||||
start = pid * BS
|
||||
offs = start + tl.arange(0, BS)
|
||||
|
||||
ni, ci, hi, wi, m = _idx(offs, n, ic, h, w)
|
||||
val = tl.zeros((BS,), dtype=AT)
|
||||
|
||||
for kh in range(ks):
|
||||
ho = hi + pad - kh
|
||||
hb = (ho >= 0) & (ho < h)
|
||||
for kw in range(ks):
|
||||
wo = wi + pad - kw
|
||||
b = hb & (wo >= 0) & (wo < w)
|
||||
|
||||
go_off = ((ni * ic + ci) * h + ho) * w + wo
|
||||
w_off = ((ni * wc + ci % wc) * ks * ks + (kh * ks + kw)) * h * w + ho * w + wo
|
||||
|
||||
go_val = tl.load(go_ptr + go_off, mask=m & b, other=0.0).to(CT)
|
||||
w_val = tl.load(w_ptr + w_off, mask=m, other=0.0).to(CT)
|
||||
val += tl.where(b & m, go_val * w_val, 0.0).to(AT)
|
||||
|
||||
tl.store(gi_ptr + offs, val.to(CT), mask=m)
|
||||
|
||||
@triton.jit
|
||||
def ska_bwd_w(
|
||||
go_ptr, x_ptr, gw_ptr,
|
||||
n, wc, h, w, ic, ks, pad,
|
||||
BS: tl.constexpr,
|
||||
CT: tl.constexpr, AT: tl.constexpr
|
||||
):
|
||||
pid = tl.program_id(0)
|
||||
start = pid * BS
|
||||
offs = start + tl.arange(0, BS)
|
||||
|
||||
ni, ci, hi, wi, m = _idx(offs, n, wc, h, w)
|
||||
|
||||
for kh in range(ks):
|
||||
hin = hi - pad + kh
|
||||
hb = (hin >= 0) & (hin < h)
|
||||
for kw in range(ks):
|
||||
win = wi - pad + kw
|
||||
b = hb & (win >= 0) & (win < w)
|
||||
w_off = ((ni * wc + ci) * ks * ks + (kh * ks + kw)) * h * w + hi * w + wi
|
||||
|
||||
val = tl.zeros((BS,), dtype=AT)
|
||||
steps = (ic - ci + wc - 1) // wc
|
||||
for s in range(tl.max(steps, axis=0)):
|
||||
cc = ci + s * wc
|
||||
cm = (cc < ic) & m & b
|
||||
|
||||
x_off = ((ni * ic + cc) * h + hin) * w + win
|
||||
go_off = ((ni * ic + cc) * h + hi) * w + wi
|
||||
|
||||
x_val = tl.load(x_ptr + x_off, mask=cm, other=0.0).to(CT)
|
||||
go_val = tl.load(go_ptr + go_off, mask=cm, other=0.0).to(CT)
|
||||
val += tl.where(cm, x_val * go_val, 0.0).to(AT)
|
||||
|
||||
tl.store(gw_ptr + w_off, val.to(CT), mask=m)
|
||||
|
||||
class SkaFn(Function):
|
||||
@staticmethod
|
||||
@custom_fwd(device_type='cuda')
|
||||
def forward(ctx, x: torch.Tensor, w: torch.Tensor) -> torch.Tensor:
|
||||
ks = int(math.sqrt(w.shape[2]))
|
||||
pad = (ks - 1) // 2
|
||||
ctx.ks, ctx.pad = ks, pad
|
||||
n, ic, h, width = x.shape
|
||||
wc = w.shape[1]
|
||||
o = torch.empty(n, ic, h, width, device=x.device, dtype=x.dtype)
|
||||
numel = o.numel()
|
||||
|
||||
x = x.contiguous()
|
||||
w = w.contiguous()
|
||||
|
||||
grid = lambda meta: _grid(numel, meta["BS"])
|
||||
|
||||
ct = tl.float16 if x.dtype == torch.float16 else (tl.float32 if x.dtype == torch.float32 else tl.float64)
|
||||
at = tl.float32 if x.dtype == torch.float16 else ct
|
||||
|
||||
ska_fwd[grid](x, w, o, n, ic, h, width, ks, pad, wc, BS=1024, CT=ct, AT=at)
|
||||
|
||||
ctx.save_for_backward(x, w)
|
||||
ctx.ct, ctx.at = ct, at
|
||||
return o
|
||||
|
||||
@staticmethod
|
||||
@custom_bwd(device_type='cuda')
|
||||
def backward(ctx, go: torch.Tensor) -> tuple:
|
||||
ks, pad = ctx.ks, ctx.pad
|
||||
x, w = ctx.saved_tensors
|
||||
n, ic, h, width = x.shape
|
||||
wc = w.shape[1]
|
||||
|
||||
go = go.contiguous()
|
||||
gx = gw = None
|
||||
ct, at = ctx.ct, ctx.at
|
||||
|
||||
if ctx.needs_input_grad[0]:
|
||||
gx = torch.empty_like(x)
|
||||
numel = gx.numel()
|
||||
ska_bwd_x[lambda meta: _grid(numel, meta["BS"])](go, w, gx, n, ic, h, width, ks, pad, wc, BS=1024, CT=ct, AT=at)
|
||||
|
||||
if ctx.needs_input_grad[1]:
|
||||
gw = torch.empty_like(w)
|
||||
numel = gw.numel() // w.shape[2]
|
||||
ska_bwd_w[lambda meta: _grid(numel, meta["BS"])](go, x, gw, n, wc, h, width, ic, ks, pad, BS=1024, CT=ct, AT=at)
|
||||
|
||||
return gx, gw, None, None
|
||||
|
||||
class SKA(torch.nn.Module):
|
||||
def forward(self, x: torch.Tensor, w: torch.Tensor) -> torch.Tensor:
|
||||
return SkaFn.apply(x, w) # type: ignore
|
||||
@@ -0,0 +1,28 @@
|
||||
einops==0.4.1
|
||||
fvcore
|
||||
easydict
|
||||
matplotlib
|
||||
yacs
|
||||
scikit-image==0.19.3
|
||||
wandb
|
||||
torch==2.4.1
|
||||
scikit-learn>=1.0
|
||||
|
||||
torch>=1.10.0
|
||||
torchvision>=0.11.0
|
||||
|
||||
timm
|
||||
|
||||
numpy>=1.19.0
|
||||
Pillow>=8.0.0
|
||||
|
||||
scikit-learn>=1.0.0
|
||||
|
||||
matplotlib>=3.3.0
|
||||
seaborn>=0.11.0
|
||||
|
||||
tqdm>=4.60.0
|
||||
|
||||
# 可选:加速训练
|
||||
# tensorboard>=2.8.0
|
||||
# apex # 混合精度训练(需要从源码安装)
|
||||
@@ -0,0 +1,304 @@
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Iterable, List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from sklearn.cluster import KMeans
|
||||
from timm.data import resolve_data_config
|
||||
from timm.data.transforms_factory import create_transform
|
||||
|
||||
from inference_artist import (
|
||||
load_checkpoint_state,
|
||||
load_model,
|
||||
preprocess_image,
|
||||
resolve_num_classes,
|
||||
)
|
||||
|
||||
IMAGE_EXTENSIONS = {'.jpg', '.jpeg', '.png', '.bmp', '.tiff', '.webp'}
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(
|
||||
'Artist feature clustering and similarity utilities',
|
||||
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
||||
)
|
||||
|
||||
# 模型与数据相关参数
|
||||
parser.add_argument('--images-dir', type=str, default=None,
|
||||
help='待聚类的图像文件夹路径(将对所有支持的图像进行特征提取并聚类)')
|
||||
parser.add_argument('--model', type=str, default='lsnet_t_artist',
|
||||
choices=['lsnet_t_artist', 'lsnet_s_artist', 'lsnet_b_artist'],
|
||||
help='用于特征提取的模型名称')
|
||||
parser.add_argument('--checkpoint', type=str, default=None,
|
||||
help='模型 checkpoint 路径(提取特征或分类需提供)')
|
||||
parser.add_argument('--feature-dim', type=int, default=None,
|
||||
help='特征维度(若模型需要可指定)')
|
||||
parser.add_argument('--device', type=str, default='cuda',
|
||||
help='推理设备:cuda 或 cpu')
|
||||
parser.add_argument('--batch-size', type=int, default=64,
|
||||
help='批量推理时的 batch size')
|
||||
parser.add_argument('--num-clusters', type=int, default=5,
|
||||
help='KMeans 聚类簇数量')
|
||||
parser.add_argument('--cluster-output', type=str, default='./output/cluster',
|
||||
help='聚类结果输出目录')
|
||||
parser.add_argument('--seed', type=int, default=42,
|
||||
help='随机种子,用于聚类可复现')
|
||||
|
||||
# 向量相似度比较相关参数
|
||||
parser.add_argument('--query-vector', type=str, default=None,
|
||||
help='待比对的目标向量文件(.npy 或包含向量的 .json/.txt)')
|
||||
parser.add_argument('--reference-vectors', type=str, nargs='*', default=None,
|
||||
help='参考向量文件列表(.npy 或 .json/.txt);也可传入目录,程序会读取目录下所有 .npy 文件')
|
||||
parser.add_argument('--top-k', type=int, default=5,
|
||||
help='返回相似度前 K 个结果(<= 参考向量数量)')
|
||||
parser.add_argument('--similarity-output', type=str, default=None,
|
||||
help='相似度结果保存路径(JSON)')
|
||||
parser.add_argument('--normalize', action='store_true',
|
||||
help='在计算相似度前对向量进行 L2 归一化')
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
if not args.images_dir and not args.query_vector:
|
||||
parser.error('至少需要指定 --images-dir 或 --query-vector 中的一个功能。')
|
||||
|
||||
if args.images_dir and not args.checkpoint:
|
||||
parser.error('--images-dir 模式需要提供 --checkpoint 以加载模型特征提取。')
|
||||
|
||||
return args
|
||||
|
||||
|
||||
def _collect_image_paths(images_dir: Path) -> List[Path]:
|
||||
paths: List[Path] = []
|
||||
for entry in images_dir.iterdir():
|
||||
if entry.is_file() and entry.suffix.lower() in IMAGE_EXTENSIONS:
|
||||
paths.append(entry)
|
||||
return sorted(paths)
|
||||
|
||||
|
||||
def _load_transform(model) -> Tuple[callable, dict]:
|
||||
config = resolve_data_config({}, model=model)
|
||||
transform = create_transform(**config)
|
||||
return transform, config
|
||||
|
||||
|
||||
def _process_batch(model, batch_tensors: List[torch.Tensor], device: torch.device) -> np.ndarray:
|
||||
batch_tensor = torch.cat(batch_tensors, dim=0).to(device)
|
||||
features = model(batch_tensor, return_features=True)
|
||||
return features.cpu().numpy()
|
||||
|
||||
|
||||
def _extract_directory_features(args) -> Tuple[np.ndarray, List[str]]:
|
||||
images_dir = Path(args.images_dir)
|
||||
if not images_dir.is_dir():
|
||||
raise FileNotFoundError(f'找不到图像文件夹: {images_dir}')
|
||||
|
||||
image_paths = _collect_image_paths(images_dir)
|
||||
if not image_paths:
|
||||
raise ValueError(f'在 {images_dir} 未找到支持的图像文件,支持扩展名: {sorted(IMAGE_EXTENSIONS)}')
|
||||
|
||||
print(f'共找到 {len(image_paths)} 张图像,开始提取特征...')
|
||||
|
||||
state_dict = load_checkpoint_state(args.checkpoint)
|
||||
num_classes = resolve_num_classes(None, None, state_dict)
|
||||
device = torch.device(args.device)
|
||||
feature_args = argparse.Namespace(
|
||||
model=args.model,
|
||||
num_classes=num_classes,
|
||||
feature_dim=args.feature_dim,
|
||||
device=args.device,
|
||||
)
|
||||
model = load_model(feature_args, state_dict)
|
||||
|
||||
transform, _ = _load_transform(model)
|
||||
|
||||
features: List[np.ndarray] = []
|
||||
names: List[str] = []
|
||||
|
||||
batch_tensors: List[torch.Tensor] = []
|
||||
batch_names: List[str] = []
|
||||
|
||||
for path in image_paths:
|
||||
try:
|
||||
tensor = preprocess_image(path, transform)
|
||||
except Exception as exc:
|
||||
print(f'[Warning] 无法处理 {path.name}: {exc}')
|
||||
continue
|
||||
batch_tensors.append(tensor)
|
||||
batch_names.append(path.name)
|
||||
|
||||
if len(batch_tensors) == args.batch_size:
|
||||
batch_features = _process_batch(model, batch_tensors, device)
|
||||
features.append(batch_features)
|
||||
names.extend(batch_names)
|
||||
batch_tensors.clear()
|
||||
batch_names.clear()
|
||||
|
||||
if batch_tensors:
|
||||
batch_features = _process_batch(model, batch_tensors, device)
|
||||
features.append(batch_features)
|
||||
names.extend(batch_names)
|
||||
|
||||
if not features:
|
||||
raise RuntimeError('特征提取失败,没有有效的图像。')
|
||||
|
||||
feature_matrix = np.concatenate(features, axis=0)
|
||||
print(f'特征提取完成,矩阵形状: {feature_matrix.shape}')
|
||||
return feature_matrix, names
|
||||
|
||||
|
||||
def _ensure_output_dir(path: Path) -> None:
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
|
||||
def _run_clustering(args, features: np.ndarray, names: List[str]) -> dict:
|
||||
print(f'开始执行 KMeans 聚类,簇数 = {args.num_clusters} ...')
|
||||
kmeans = KMeans(n_clusters=args.num_clusters, random_state=args.seed, n_init='auto')
|
||||
labels = kmeans.fit_predict(features)
|
||||
|
||||
clusters: dict = {}
|
||||
for name, label in zip(names, labels):
|
||||
clusters.setdefault(int(label), []).append(name)
|
||||
|
||||
centroids = kmeans.cluster_centers_.tolist()
|
||||
inertia = float(kmeans.inertia_)
|
||||
|
||||
clustering_result = {
|
||||
'num_clusters': args.num_clusters,
|
||||
'inertia': inertia,
|
||||
'clusters': clusters,
|
||||
'cluster_sizes': {str(idx): len(items) for idx, items in clusters.items()},
|
||||
'centroids': centroids,
|
||||
}
|
||||
return clustering_result
|
||||
|
||||
|
||||
def load_vector_file(path: Path) -> np.ndarray:
|
||||
suffix = path.suffix.lower()
|
||||
if suffix == '.npy':
|
||||
vector = np.load(path)
|
||||
elif suffix in {'.json', '.txt'}:
|
||||
with path.open('r', encoding='utf-8') as f:
|
||||
data = json.load(f)
|
||||
vector = np.asarray(data, dtype=np.float32)
|
||||
else:
|
||||
raise ValueError(f'不支持的向量文件格式: {path}')
|
||||
|
||||
vector = np.asarray(vector, dtype=np.float32)
|
||||
if vector.ndim > 1:
|
||||
vector = vector.reshape(-1)
|
||||
return vector
|
||||
|
||||
|
||||
def _expand_reference_paths(ref_inputs: Optional[Iterable[str]]) -> List[Path]:
|
||||
if not ref_inputs:
|
||||
return []
|
||||
|
||||
paths: List[Path] = []
|
||||
for item in ref_inputs:
|
||||
p = Path(item)
|
||||
if p.is_dir():
|
||||
paths.extend(sorted(child for child in p.iterdir() if child.suffix.lower() == '.npy'))
|
||||
elif p.exists():
|
||||
paths.append(p)
|
||||
else:
|
||||
raise FileNotFoundError(f'参考向量不存在: {item}')
|
||||
return paths
|
||||
|
||||
|
||||
def _compute_similarity(query_vector: np.ndarray,
|
||||
reference_vectors: List[Tuple[Path, np.ndarray]],
|
||||
normalize: bool) -> List[Tuple[str, float]]:
|
||||
if normalize:
|
||||
query_vector = _normalize_vector(query_vector)
|
||||
ref_vectors = [(path, _normalize_vector(vec)) for path, vec in reference_vectors]
|
||||
else:
|
||||
ref_vectors = reference_vectors
|
||||
|
||||
similarities: List[Tuple[str, float]] = []
|
||||
q_norm = np.linalg.norm(query_vector)
|
||||
if q_norm == 0:
|
||||
raise ValueError('查询向量范数为 0,无法计算相似度。')
|
||||
|
||||
for path, vec in ref_vectors:
|
||||
denom = np.linalg.norm(vec) * q_norm
|
||||
if denom == 0:
|
||||
sim = 0.0
|
||||
else:
|
||||
sim = float(np.dot(query_vector, vec) / denom)
|
||||
similarities.append((str(path), sim))
|
||||
|
||||
similarities.sort(key=lambda x: x[1], reverse=True)
|
||||
return similarities
|
||||
|
||||
|
||||
def _normalize_vector(vector: np.ndarray) -> np.ndarray:
|
||||
norm = np.linalg.norm(vector)
|
||||
if norm == 0:
|
||||
return vector
|
||||
return vector / norm
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
clustering_result = None
|
||||
similarity_result = None
|
||||
|
||||
if args.images_dir:
|
||||
features, names = _extract_directory_features(args)
|
||||
output_dir = Path(args.cluster_output)
|
||||
_ensure_output_dir(output_dir)
|
||||
|
||||
clustering_result = _run_clustering(args, features, names)
|
||||
|
||||
features_path = output_dir / 'features.npy'
|
||||
np.save(features_path, features)
|
||||
with (output_dir / 'cluster_assignments.json').open('w', encoding='utf-8') as f:
|
||||
json.dump(clustering_result, f, indent=2, ensure_ascii=False)
|
||||
|
||||
print(f'聚类结果已保存到 {output_dir}')
|
||||
|
||||
if args.query_vector:
|
||||
query_path = Path(args.query_vector)
|
||||
if not query_path.exists():
|
||||
raise FileNotFoundError(f'查询向量文件不存在: {query_path}')
|
||||
|
||||
query_vector = load_vector_file(query_path)
|
||||
reference_paths = _expand_reference_paths(args.reference_vectors)
|
||||
if not reference_paths:
|
||||
raise ValueError('需要至少提供一个参考向量文件或目录。')
|
||||
|
||||
reference_vectors = [(path, load_vector_file(path)) for path in reference_paths]
|
||||
similarity_pairs = _compute_similarity(query_vector, reference_vectors, args.normalize)
|
||||
|
||||
top_k = min(args.top_k, len(similarity_pairs))
|
||||
similarity_result = similarity_pairs[:top_k]
|
||||
|
||||
print('相似度 Top-{} 结果:'.format(top_k))
|
||||
for idx, (path, score) in enumerate(similarity_result, 1):
|
||||
print(f' {idx}. {path}: {score:.6f}')
|
||||
|
||||
if args.similarity_output:
|
||||
output_path = Path(args.similarity_output)
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with output_path.open('w', encoding='utf-8') as f:
|
||||
json.dump(
|
||||
{
|
||||
'query': str(query_path),
|
||||
'top_k': top_k,
|
||||
'results': [{'path': p, 'cosine_similarity': s} for p, s in similarity_result],
|
||||
},
|
||||
f,
|
||||
indent=2,
|
||||
ensure_ascii=False,
|
||||
)
|
||||
print(f'相似度结果已保存到 {output_path}')
|
||||
|
||||
if clustering_result is None and similarity_result is None:
|
||||
raise RuntimeError('未执行任何操作,请检查输入参数。')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -0,0 +1,139 @@
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Iterable, List, Tuple
|
||||
|
||||
import numpy as np
|
||||
|
||||
SUPPORTED_SUFFIXES = {'.npy', '.json', '.txt'}
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
'Compare a query vector with multiple reference vectors',
|
||||
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
||||
)
|
||||
parser.add_argument('--query-vector', required=True, type=str,
|
||||
help='待比较的查询向量文件 (.npy / .json / .txt)')
|
||||
parser.add_argument('--reference-vectors', required=True, nargs='+', type=str,
|
||||
help='参考向量文件或目录列表(目录中会读取所有 .npy 文件)')
|
||||
parser.add_argument('--top-k', default=5, type=int,
|
||||
help='返回相似度前 K 名结果')
|
||||
parser.add_argument('--normalize', action='store_true',
|
||||
help='计算相似度前对向量执行 L2 归一化')
|
||||
parser.add_argument('--output', default=None, type=str,
|
||||
help='可选的输出 JSON 文件路径')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def load_vector_file(path: Path) -> np.ndarray:
|
||||
suffix = path.suffix.lower()
|
||||
if suffix == '.npy':
|
||||
vector = np.load(path)
|
||||
elif suffix in {'.json', '.txt'}:
|
||||
with path.open('r', encoding='utf-8') as f:
|
||||
data = json.load(f)
|
||||
vector = np.asarray(data, dtype=np.float32)
|
||||
else:
|
||||
raise ValueError(f'不支持的向量文件格式: {path}')
|
||||
|
||||
vector = np.asarray(vector, dtype=np.float32)
|
||||
if vector.ndim > 1:
|
||||
vector = vector.reshape(-1)
|
||||
if vector.size == 0:
|
||||
raise ValueError(f'向量文件为空: {path}')
|
||||
return vector
|
||||
|
||||
|
||||
def expand_reference_paths(ref_inputs: Iterable[str]) -> List[Path]:
|
||||
paths: List[Path] = []
|
||||
for item in ref_inputs:
|
||||
p = Path(item)
|
||||
if not p.exists():
|
||||
raise FileNotFoundError(f'参考向量不存在: {item}')
|
||||
if p.is_dir():
|
||||
paths.extend(sorted(child for child in p.iterdir() if child.suffix.lower() in SUPPORTED_SUFFIXES))
|
||||
else:
|
||||
if p.suffix.lower() not in SUPPORTED_SUFFIXES:
|
||||
raise ValueError(f'不支持的参考向量格式: {p}')
|
||||
paths.append(p)
|
||||
if not paths:
|
||||
raise ValueError('未找到任何参考向量文件。')
|
||||
return paths
|
||||
|
||||
|
||||
def normalize_vector(vector: np.ndarray) -> np.ndarray:
|
||||
norm = np.linalg.norm(vector)
|
||||
if norm == 0:
|
||||
raise ValueError('向量范数为 0,无法进行归一化。')
|
||||
return vector / norm
|
||||
|
||||
|
||||
def compute_similarity(query_vector: np.ndarray,
|
||||
reference_vectors: List[Tuple[Path, np.ndarray]],
|
||||
normalize: bool) -> List[Tuple[str, float]]:
|
||||
if normalize:
|
||||
query_vector = normalize_vector(query_vector)
|
||||
reference_vectors = [(path, normalize_vector(vec)) for path, vec in reference_vectors]
|
||||
|
||||
q_norm = np.linalg.norm(query_vector)
|
||||
if q_norm == 0:
|
||||
raise ValueError('查询向量范数为 0,无法计算相似度。')
|
||||
|
||||
similarities: List[Tuple[str, float]] = []
|
||||
for path, vec in reference_vectors:
|
||||
denom = np.linalg.norm(vec) * q_norm
|
||||
if denom == 0:
|
||||
sim = 0.0
|
||||
else:
|
||||
sim = float(np.dot(query_vector, vec) / denom)
|
||||
similarities.append((str(path), sim))
|
||||
|
||||
similarities.sort(key=lambda x: x[1], reverse=True)
|
||||
return similarities
|
||||
|
||||
|
||||
def save_results(path: Path, query: Path, top_k: int, results: List[Tuple[str, float]]) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with path.open('w', encoding='utf-8') as f:
|
||||
json.dump(
|
||||
{
|
||||
'query': str(query),
|
||||
'top_k': top_k,
|
||||
'results': [
|
||||
{'path': result_path, 'cosine_similarity': score}
|
||||
for result_path, score in results
|
||||
],
|
||||
},
|
||||
f,
|
||||
indent=2,
|
||||
ensure_ascii=False,
|
||||
)
|
||||
print(f'相似度结果已保存到 {path}')
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
query_path = Path(args.query_vector)
|
||||
if not query_path.exists():
|
||||
raise FileNotFoundError(f'查询向量不存在: {query_path}')
|
||||
query_vector = load_vector_file(query_path)
|
||||
|
||||
reference_paths = expand_reference_paths(args.reference_vectors)
|
||||
reference_vectors = [(path, load_vector_file(path)) for path in reference_paths]
|
||||
|
||||
similarity_pairs = compute_similarity(query_vector, reference_vectors, args.normalize)
|
||||
top_k = min(args.top_k, len(similarity_pairs))
|
||||
top_results = similarity_pairs[:top_k]
|
||||
|
||||
print(f'相似度 Top-{top_k} 结果:')
|
||||
for idx, (path, score) in enumerate(top_results, 1):
|
||||
print(f' {idx}. {path}: {score:.6f}')
|
||||
|
||||
if args.output:
|
||||
save_results(Path(args.output), query_path, top_k, top_results)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -0,0 +1,166 @@
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import List, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from sklearn.cluster import KMeans
|
||||
from timm.data import resolve_data_config
|
||||
from timm.data.transforms_factory import create_transform
|
||||
|
||||
from inference_artist import (
|
||||
load_checkpoint_state,
|
||||
load_model,
|
||||
preprocess_image,
|
||||
resolve_num_classes,
|
||||
)
|
||||
|
||||
IMAGE_EXTENSIONS = {'.jpg', '.jpeg', '.png', '.bmp', '.tiff', '.webp'}
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
'Extract LSNet artist features and cluster a folder',
|
||||
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
||||
)
|
||||
parser.add_argument('--images-dir', required=True, type=str,
|
||||
help='包含图像的文件夹路径,将对其中所有支持格式的图像提取特征并聚类')
|
||||
parser.add_argument('--model', default='lsnet_t_artist', type=str,
|
||||
choices=['lsnet_t_artist', 'lsnet_s_artist', 'lsnet_b_artist'],
|
||||
help='用于特征提取的模型名称')
|
||||
parser.add_argument('--checkpoint', required=True, type=str,
|
||||
help='模型 checkpoint 路径')
|
||||
parser.add_argument('--feature-dim', default=None, type=int,
|
||||
help='特征维度(若模型需要可显式指定)')
|
||||
parser.add_argument('--device', default='cuda', type=str,
|
||||
help='推理设备')
|
||||
parser.add_argument('--batch-size', default=64, type=int,
|
||||
help='批量推理时的 batch size')
|
||||
parser.add_argument('--num-clusters', default=5, type=int,
|
||||
help='KMeans 聚类簇数量')
|
||||
parser.add_argument('--seed', default=42, type=int,
|
||||
help='随机种子,确保聚类可复现')
|
||||
parser.add_argument('--output-dir', default='./output/cluster', type=str,
|
||||
help='输出目录,将保存特征矩阵与聚类结果 JSON 文件')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def _collect_image_paths(images_dir: Path) -> List[Path]:
|
||||
return sorted(
|
||||
[p for p in images_dir.iterdir() if p.is_file() and p.suffix.lower() in IMAGE_EXTENSIONS]
|
||||
)
|
||||
|
||||
|
||||
def _load_transform(model) -> Tuple[torch.nn.Module, dict]:
|
||||
config = resolve_data_config({}, model=model)
|
||||
transform = create_transform(**config)
|
||||
return transform, config
|
||||
|
||||
|
||||
def _process_batch(model, tensors: List[torch.Tensor], device: torch.device) -> np.ndarray:
|
||||
batch_tensor = torch.cat(tensors, dim=0).to(device)
|
||||
features = model(batch_tensor, return_features=True)
|
||||
return features.cpu().numpy()
|
||||
|
||||
|
||||
def extract_features(args: argparse.Namespace) -> Tuple[np.ndarray, List[str]]:
|
||||
images_dir = Path(args.images_dir)
|
||||
if not images_dir.is_dir():
|
||||
raise FileNotFoundError(f'找不到图像文件夹: {images_dir}')
|
||||
|
||||
image_paths = _collect_image_paths(images_dir)
|
||||
if not image_paths:
|
||||
raise ValueError(
|
||||
f'在 {images_dir} 未找到支持的图像文件,支持扩展名: {sorted(IMAGE_EXTENSIONS)}'
|
||||
)
|
||||
|
||||
state_dict = load_checkpoint_state(args.checkpoint)
|
||||
num_classes = resolve_num_classes(None, None, state_dict)
|
||||
feature_args = argparse.Namespace(
|
||||
model=args.model,
|
||||
num_classes=num_classes,
|
||||
feature_dim=args.feature_dim,
|
||||
device=args.device,
|
||||
)
|
||||
model = load_model(feature_args, state_dict)
|
||||
device = torch.device(args.device)
|
||||
|
||||
transform, _ = _load_transform(model)
|
||||
|
||||
features: List[np.ndarray] = []
|
||||
names: List[str] = []
|
||||
batch_tensors: List[torch.Tensor] = []
|
||||
batch_names: List[str] = []
|
||||
|
||||
for path in image_paths:
|
||||
try:
|
||||
tensor = preprocess_image(path, transform)
|
||||
except Exception as exc: # pylint: disable=broad-except
|
||||
print(f'[Warning] 无法处理 {path.name}: {exc}')
|
||||
continue
|
||||
|
||||
batch_tensors.append(tensor)
|
||||
batch_names.append(path.name)
|
||||
|
||||
if len(batch_tensors) == args.batch_size:
|
||||
batch_features = _process_batch(model, batch_tensors, device)
|
||||
features.append(batch_features)
|
||||
names.extend(batch_names)
|
||||
batch_tensors.clear()
|
||||
batch_names.clear()
|
||||
|
||||
if batch_tensors:
|
||||
batch_features = _process_batch(model, batch_tensors, device)
|
||||
features.append(batch_features)
|
||||
names.extend(batch_names)
|
||||
|
||||
if not features:
|
||||
raise RuntimeError('特征提取失败,没有成功处理的图像。')
|
||||
|
||||
feature_matrix = np.concatenate(features, axis=0)
|
||||
print(f'特征提取完成,共 {feature_matrix.shape[0]} 张图像,特征维度 {feature_matrix.shape[1]}')
|
||||
return feature_matrix, names
|
||||
|
||||
|
||||
def run_clustering(args: argparse.Namespace, features: np.ndarray, names: List[str]) -> dict:
|
||||
print(f'开始执行 KMeans 聚类,簇数量 = {args.num_clusters}')
|
||||
kmeans = KMeans(n_clusters=args.num_clusters, random_state=args.seed, n_init='auto')
|
||||
labels = kmeans.fit_predict(features)
|
||||
|
||||
clusters = {}
|
||||
for name, label in zip(names, labels):
|
||||
clusters.setdefault(int(label), []).append(name)
|
||||
|
||||
result = {
|
||||
'num_clusters': args.num_clusters,
|
||||
'inertia': float(kmeans.inertia_),
|
||||
'cluster_sizes': {str(k): len(v) for k, v in clusters.items()},
|
||||
'clusters': clusters,
|
||||
'centroids': kmeans.cluster_centers_.tolist(),
|
||||
}
|
||||
return result
|
||||
|
||||
|
||||
def save_outputs(args: argparse.Namespace, features: np.ndarray, clustering: dict) -> None:
|
||||
output_dir = Path(args.output_dir)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
features_path = output_dir / 'features.npy'
|
||||
np.save(features_path, features)
|
||||
|
||||
with (output_dir / 'cluster_assignments.json').open('w', encoding='utf-8') as f:
|
||||
json.dump(clustering, f, indent=2, ensure_ascii=False)
|
||||
|
||||
print(f'特征矩阵与聚类结果已保存到 {output_dir}')
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
features, names = extract_features(args)
|
||||
clustering = run_clustering(args, features, names)
|
||||
save_outputs(args, features, clustering)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
Reference in New Issue
Block a user