Files
spawner1145-comfyui-lsnet/tests/render_analysis_examples.py
T
2026-10-06 00:30:52 +08:00

146 lines
8.5 KiB
Python

"""Infer each example once, cache features, then exercise every plotting node mode."""
import argparse
import csv
import hashlib
import html
import json
from pathlib import Path
import re
import sys
import numpy as np
from PIL import Image, ImageDraw, ImageFont, ImageOps
import torch
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from feature_analysis import CHART_TYPES
from kaloscope_dinov3.preprocessing import prepare_rgb
from model_loading import load_model_bundle
from test_model_loading import load_nodes
def main():
parser = argparse.ArgumentParser()
parser.add_argument('--input', default='../test')
parser.add_argument('--model-dir', default='../model')
parser.add_argument('--output', default='../outputs')
parser.add_argument('--device', default='cuda')
parser.add_argument('--batch-size', default=2, type=int)
parser.add_argument('--reuse-cache', action='store_true', help='Redraw charts from features.pt without loading the model')
args = parser.parse_args()
torch.set_num_threads(4)
source = Path(args.input).resolve()
output = Path(args.output).resolve()
output.mkdir(parents=True, exist_ok=True)
paths = sorted(path for path in source.iterdir() if path.suffix.lower() in {'.png','.jpg','.jpeg','.webp','.bmp'})
if not paths:
raise ValueError(f'No images in {source}')
if args.batch_size < 1:
raise ValueError('batch-size must be positive')
cache_path=output/'features.pt'
if args.reuse_cache:
cached=torch.load(cache_path,map_location='cpu',weights_only=True)
features,patch_features=cached['features'],cached['patch_tokens']
labels,records=cached['labels'],cached['source_images']
model_info=cached.get('model_info',{'model_type':'dinov3_vitb16','feature_source':'projector',
'input_size':512,'checkpoint':str(Path(args.model_dir).resolve()/'best.pt')})
inference_batches=cached.get('inference_batches',4)
previews=[]
for record in records:
with Image.open(record['path']) as original:
previews.append(ImageOps.pad(prepare_rgb(original),(240,200),color='white'))
thumbnails=torch.stack([torch.from_numpy(np.array(image)).float()/255 for image in previews])
print('Reusing cached features: zero new model forwards',flush=True)
else:
bundle = load_model_bundle(args.model_dir, device=args.device)
model = bundle['model']
model_info = {key: bundle[key] for key in ('model_type','feature_source','input_size','checkpoint')}
if not hasattr(model, 'extract_tensor'):
raise ValueError('This example suite requires a DINOv3 checkpoint for patch analysis')
vectors, patches, previews = [], [], []
labels, records = [], []
inference_batches = 0
with torch.inference_mode():
for start in range(0,len(paths),args.batch_size):
batch = []
for index, path in enumerate(paths[start:start+args.batch_size], start):
with Image.open(path) as original:
image = prepare_rgb(original)
batch.append(bundle['transform'](image))
previews.append(ImageOps.pad(image, (240,200), color='white'))
artist = re.search(r'_drawn_by_(.+?)__[0-9a-f]+$', path.stem)
label = f'{index+1:02d} ' + (artist.group(1) if artist else path.stem[:22])
labels.append(label)
with path.open('rb') as stream:
digest = hashlib.file_digest(stream, 'sha256').hexdigest()
records.append({'index': index, 'label': label, 'path': str(path), 'sha256': digest})
tokens = model.backbone.forward_features(torch.stack(batch).to(args.device))
pooled = model._pool(tokens['x_norm_clstoken'],tokens['x_norm_patchtokens'],model.pooling)
feature = model.projector(pooled) if model.feature_source=='projector' else pooled
vectors.append(feature.float().cpu())
patches.append(tokens['x_norm_patchtokens'].float().cpu())
inference_batches += 1
print(f'Inferred images {start+1}-{min(start+args.batch_size,len(paths))} once',flush=True)
features = torch.cat(vectors)
patch_features = torch.cat(patches)
thumbnails = torch.stack([torch.from_numpy(np.array(image)).float()/255 for image in previews])
torch.save({'features':features, 'patch_tokens':patch_features, 'labels':labels, 'source_images':records,'model_info':model_info,'inference_batches':inference_batches},output/'features.pt')
# Release GPU model; all charts below consume cached CPU tensors only.
del model, bundle
if torch.cuda.is_available():
torch.cuda.empty_cache()
nodes = load_nodes()
node = nodes.KaloscopeFeatureAnalysisNode()
generated = []
for chart in CHART_TYPES:
tensor = patch_features if chart=='patch_energy' else features
result, text, distances = node.analyze(tensor,chart_type=chart,labels=json.dumps(labels),images=thumbnails,
tensor_layout='tokens' if chart=='patch_energy' else 'vectors',n_clusters=3,top_k=3,
metric='cosine',normalize=True,width=1600,height=1100,perplexity=3.0)
report = json.loads(text)
report['source_images']=records
report['feature_source']='patch_tokens' if chart=='patch_energy' else model_info['feature_source']
array = (result[0].clamp(0,1).numpy()*255).round().astype(np.uint8)
Image.fromarray(array).save(output/f'{chart}.png',dpi=(120,120))
(output/f'{chart}.json').write_text(json.dumps(report,ensure_ascii=False,indent=2,allow_nan=False),encoding='utf-8')
if chart=='relationship_graph':
with (output/'distance_matrix.csv').open('w',encoding='utf-8-sig',newline='') as stream:
writer=csv.writer(stream)
writer.writerow(['image',*labels])
writer.writerows([label,*row.tolist()] for label,row in zip(labels,distances))
generated.append({'type':chart,'image':f'{chart}.png','report':f'{chart}.json'})
print(f'Saved {chart}.png',flush=True)
manifest={'input_directory':str(source),'checkpoint':model_info['checkpoint'],
'architecture':model_info['model_type'],'model_info':model_info,'feature_shape':list(features.shape),'patch_shape':list(patch_features.shape),
'source_images':records,'inference_batches':inference_batches,'image_forward_count':len(paths),
'analysis_model_forward_count':0,'this_run_model_forward_count':0 if args.reuse_cache else len(paths),'seed':42,'charts':generated}
(output/'manifest.json').write_text(json.dumps(manifest,ensure_ascii=False,indent=2),encoding='utf-8')
body=''.join(f'<section><h2>{html.escape(item["type"])}</h2><a href="{item["image"]}"><img src="{item["image"]}" alt="{html.escape(item["type"])}"></a><p><a href="{item["report"]}">分析数据 JSON</a></p></section>' for item in generated)
(output/'index.html').write_text('''<!doctype html><html lang="zh-CN"><meta charset="utf-8"><title>Kaloscope 分析图示例</title>
<style>body{font:16px system-ui;background:#eef2f5;color:#243746;max-width:1500px;margin:32px auto;padding:0 24px}section{background:white;margin:24px 0;padding:24px;border-radius:12px}img{max-width:100%;height:auto}a{color:#0072b2}</style>
<h1>Kaloscope 特征分析示例</h1><p>每张输入图片仅推理一次。18 类图表复用缓存特征;二维距离布局是近似结果,具体距离以矩阵和 JSON 为准。</p>
<p><a href="manifest.json">来源与运行记录</a> · <a href="distance_matrix.csv">距离矩阵 CSV</a></p>'''+body+'</html>',encoding='utf-8')
contact = Image.new('RGB',(1600,math_rows(len(generated))*390),'#eef2f5')
drawer=ImageDraw.Draw(contact)
try:
font=ImageFont.truetype('C:/Windows/Fonts/arial.ttf',22)
except OSError:
font=ImageFont.load_default()
for index,item in enumerate(generated):
x=(index%3)*530+10; y=(index//3)*390
with Image.open(output/item['image']) as image:
preview=ImageOps.contain(image,(510,350))
contact.paste(preview,(x,y+32))
drawer.text((x,y+5),item['type'],fill='#243746',font=font)
contact.save(output/'overview.png')
print(f'{len(generated)} chart types written to {output}; analysis performed zero model forwards',flush=True)
def math_rows(count):
return (count+2)//3
if __name__=='__main__':
main()