diff --git a/cluster_artist.py b/cluster_artist.py new file mode 100644 index 0000000..9717034 --- /dev/null +++ b/cluster_artist.py @@ -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) diff --git a/comfyui_lsnet_artist_node.py b/comfyui_lsnet_artist_node.py new file mode 100644 index 0000000..64856b6 --- /dev/null +++ b/comfyui_lsnet_artist_node.py @@ -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" +} \ No newline at end of file diff --git a/inference_artist.py b/inference_artist.py new file mode 100644 index 0000000..9f97f3d --- /dev/null +++ b/inference_artist.py @@ -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) diff --git a/model/__init__.py b/model/__init__.py new file mode 100644 index 0000000..bd62110 --- /dev/null +++ b/model/__init__.py @@ -0,0 +1,3 @@ +from .lsnet import * +from .lsnet_artist import * +from .build import * diff --git a/model/build.py b/model/build.py new file mode 100644 index 0000000..136b4f2 --- /dev/null +++ b/model/build.py @@ -0,0 +1 @@ +import model.lsnet \ No newline at end of file diff --git a/model/lsnet.py b/model/lsnet.py new file mode 100644 index 0000000..b214b88 --- /dev/null +++ b/model/lsnet.py @@ -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) \ No newline at end of file diff --git a/model/lsnet_artist.py b/model/lsnet_artist.py new file mode 100644 index 0000000..83f812d --- /dev/null +++ b/model/lsnet_artist.py @@ -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 diff --git a/model/ska.py b/model/ska.py new file mode 100644 index 0000000..ee0ef7a --- /dev/null +++ b/model/ska.py @@ -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 diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..029c60b --- /dev/null +++ b/requirements.txt @@ -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 # 混合精度训练(需要从源码安装) diff --git a/tools/cluster_and_compare.py b/tools/cluster_and_compare.py new file mode 100644 index 0000000..92c084f --- /dev/null +++ b/tools/cluster_and_compare.py @@ -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() diff --git a/tools/compare_vectors.py b/tools/compare_vectors.py new file mode 100644 index 0000000..1515395 --- /dev/null +++ b/tools/compare_vectors.py @@ -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() diff --git a/tools/extract_cluster_features.py b/tools/extract_cluster_features.py new file mode 100644 index 0000000..f0baae6 --- /dev/null +++ b/tools/extract_cluster_features.py @@ -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()