This commit is contained in:
spawner1145
2025-10-18 16:14:26 +08:00
parent 76769b9eb2
commit e7e63e844d
12 changed files with 2659 additions and 0 deletions
+251
View File
@@ -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)
+448
View File
@@ -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"
}
+475
View File
@@ -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)
+3
View File
@@ -0,0 +1,3 @@
from .lsnet import *
from .lsnet_artist import *
from .build import *
+1
View File
@@ -0,0 +1 @@
import model.lsnet
+405
View File
@@ -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)
+271
View File
@@ -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
View File
@@ -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
+28
View File
@@ -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 # 混合精度训练(需要从源码安装)
+304
View File
@@ -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()
+139
View File
@@ -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()
+166
View File
@@ -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()