Files
2026-10-06 00:30:52 +08:00

320 lines
13 KiB
Python

"""
画师风格模型推理脚本
支持两种模式:
1. 聚类模式:提取特征向量用于聚类
2. 分类模式:直接输出分类结果
"""
import argparse
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 model_loading import (
FEATURE_OUTPUTS, load_model_bundle, load_checkpoint_state, normalize_state_dict_keys, load_class_mapping,
)
from backend_lsnet.analysis import extract_tensor_batch, cache_bytes
def get_args_parser():
parser = argparse.ArgumentParser('Artist Style Inference', add_help=False)
# 模型参数
parser.add_argument('--model', default=None, type=str,
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('--input-size', default=None, type=int,
help='Input image size')
# 推理模式
parser.add_argument('--mode', default='auto', type=str,
choices=['auto', 'classify', 'cluster', 'both'],
help='Inference mode: classify (with head), cluster (features only), or both')
parser.add_argument('--output-type', choices=FEATURE_OUTPUTS, default='default')
parser.add_argument('--layers', default='-1', help='Intermediate layer indices, comma separated')
parser.add_argument('--no-intermediate-norm', action='store_true')
# 输入输出
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('--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 resolved_mode(model, mode):
if mode not in ('auto', 'classify', 'cluster', 'both'):
raise ValueError(f'Unknown inference mode: {mode}')
if mode == 'auto':
return 'classify' if model.has_classifier else 'cluster'
if mode in ('classify', 'both') and not model.has_classifier:
raise ValueError('Model has no classification head; use --mode auto or --mode cluster')
return mode
def load_model(args, state_dict=None):
bundle = load_model_bundle(checkpoint=args.checkpoint, device=args.device, model_name=args.model,
class_csv=args.class_csv, input_size=args.input_size)
args.mode = resolved_mode(bundle['model'], args.mode)
return bundle['model']
def preprocess_image(image_path, transform):
"""预处理单张图像"""
with Image.open(image_path) as image:
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):
"""对图像进行分类"""
if not model.has_classifier:
raise ValueError('This checkpoint has no classification head')
with torch.no_grad():
image_tensor = image_tensor.to(device)
# 使用分类头
logits = model(image_tensor, return_features=False)
# 计算概率
probs = F.softmax(logits.float(), 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, output_type='default', layers='-1', intermediate_norm=True):
"""提取特征向量用于聚类"""
with torch.no_grad():
image_tensor = image_tensor.to(device)
# 不使用分类头,直接返回特征
features = extract_tensor_batch(model, image_tensor, output_type, layers, intermediate_norm)
return features.float().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, args.output_type, args.layers, not args.no_intermediate_norm)
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 = sorted(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 = []
valid_paths = []
for path in batch_paths:
try:
tensor = preprocess_image(path, transform)
batch_tensors.append(tensor)
valid_paths.append(path)
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.float(), 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 = extract_tensor_batch(model, batch_tensor, args.output_type, args.layers, not args.no_intermediate_norm)
# 保存结果
for j, path in enumerate(valid_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[str(path.relative_to(input_dir))] = result
print(f"Processed {min(i + args.batch_size, len(image_paths))}/{len(image_paths)} images")
return all_results
def main(args):
if args.batch_size < 1 or args.top_k < 1:
raise ValueError('batch-size and top-k must be positive')
bundle = load_model_bundle(checkpoint=args.checkpoint, model_name=args.model, device=args.device,
class_csv=args.class_csv, input_size=args.input_size)
model, transform, class_mapping = bundle['model'], bundle['transform'], bundle['class_mapping']
args.mode = resolved_mode(model, args.mode)
print(f"Loaded {bundle['model_type']}: classifier={bundle['has_classifier']}, "
f"features={bundle['feature_dim']} ({bundle['feature_source']}), input={bundle['input_size']}")
output_dir = Path(args.output)
output_dir.mkdir(parents=True, exist_ok=True)
# 判断输入类型
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}")
if results and 'features' in results:
(output_dir / 'features.npz').write_bytes(cache_bytes(torch.tensor([results['features']]),
[input_path.name], args.output_type, args.layers))
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)
(output_dir / 'features.npz').write_bytes(cache_bytes(torch.from_numpy(features_array),
image_names, args.output_type, args.layers))
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)