File size: 10,421 Bytes
c8c00f0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 | """
DefectFill Experiment Results Visualization Module
Generates the following visualization charts:
1. Heatmaps - Shows config vs class performance matrix
2. Scatter Plot - KID vs IC-LPIPS quality-diversity tradeoff analysis
3. Grouped Bar Charts - Comparison of each config across different classes
Usage:
python visualize_results.py --csv_path evaluation_results.csv --output_dir ./figures
"""
import os
import argparse
import pandas as pd
import numpy as np
# Set non-interactive backend for headless environments (e.g., servers/clusters)
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import matplotlib.patches as mpatches
from matplotlib.lines import Line2D
import seaborn as sns
# Global Plotting Configuration
plt.rcParams['axes.unicode_minus'] = False
sns.set_style("whitegrid")
sns.set_context("paper", font_scale=1.2)
# Visualization Color Scheme
CONFIG_COLORS = {
'base': '#3498db', # Blue
'tex': '#e74c3c', # Red
'obj': '#2ecc71' # Green
}
# Shapes to distinguish between Object and Texture datasets
CATEGORY_MARKERS = {
'object': 'o', # Circle
'texture': 's' # Square
}
# MVTec AD Dataset Groupings
OBJECT_CLASSES = ['bottle', 'cable', 'hazelnut', 'metal_nut', 'toothbrush']
TEXTURE_CLASSES = ['carpet', 'grid', 'leather', 'tile', 'wood']
def load_and_validate_data(csv_path):
"""Load and validate the evaluation CSV data."""
if not os.path.exists(csv_path):
raise FileNotFoundError(f"CSV file not found: {csv_path}")
df = pd.read_csv(csv_path)
# Required metric columns
required_columns = ['class', 'config', 'category_type', 'KID_mean', 'KID_std',
'IC_LPIPS_mean', 'IC_LPIPS_std']
missing_columns = [col for col in required_columns if col not in df.columns]
if missing_columns:
raise ValueError(f"CSV missing required columns: {missing_columns}")
print(f"Successfully loaded data: {len(df)} records")
print(f"Classes found: {df['class'].unique().tolist()}")
print(f"Configs found: {df['config'].unique().tolist()}")
return df
def create_heatmaps(df, output_dir):
"""
Generate performance matrices.
1. KID Heatmap: Quality evaluation (lower is better).
2. IC-LPIPS Heatmap: Diversity evaluation (higher is better).
"""
print("\nGenerating performance heatmaps...")
fig, axes = plt.subplots(1, 2, figsize=(16, 8))
# === KID Heatmap ===
pivot_kid = df.pivot_table(index='class', columns='config', values='KID_mean', aggfunc='mean')
# Sort by Category (Object classes first, then Texture)
class_order = [c for c in OBJECT_CLASSES if c in pivot_kid.index] + \
[c for c in TEXTURE_CLASSES if c in pivot_kid.index]
pivot_kid = pivot_kid.reindex(class_order)
config_order = ['base', 'tex', 'obj']
pivot_kid = pivot_kid.reindex(columns=[c for c in config_order if c in pivot_kid.columns])
ax1 = axes[0]
sns.heatmap(pivot_kid, annot=True, fmt='.4f', cmap='RdYlGn_r',
linewidths=0.5, ax=ax1,
cbar_kws={'label': 'KID (lower is better)', 'shrink': 0.8},
annot_kws={'size': 11, 'weight': 'bold'})
ax1.set_title('KID Quality Evaluation Matrix', fontsize=14, fontweight='bold', pad=15)
# Visual separators for dataset types
n_object = len([c for c in OBJECT_CLASSES if c in class_order])
if 0 < n_object < len(class_order):
ax1.axhline(y=n_object, color='black', linewidth=2)
# === IC-LPIPS Heatmap ===
pivot_lpips = df.pivot_table(index='class', columns='config', values='IC_LPIPS_mean', aggfunc='mean')
pivot_lpips = pivot_lpips.reindex(class_order)
pivot_lpips = pivot_lpips.reindex(columns=[c for c in config_order if c in pivot_lpips.columns])
ax2 = axes[1]
sns.heatmap(pivot_lpips, annot=True, fmt='.4f', cmap='RdYlGn',
linewidths=0.5, ax=ax2,
cbar_kws={'label': 'IC-LPIPS (higher is better)', 'shrink': 0.8},
annot_kws={'size': 11, 'weight': 'bold'})
ax2.set_title('IC-LPIPS Diversity Evaluation Matrix', fontsize=14, fontweight='bold', pad=15)
if 0 < n_object < len(class_order):
ax2.axhline(y=n_object, color='black', linewidth=2)
plt.tight_layout()
output_path = os.path.join(output_dir, 'heatmaps.png')
plt.savefig(output_path, dpi=300, bbox_inches='tight', facecolor='white')
print(f" Heatmaps saved to: {output_path}")
plt.close()
def create_scatter_plot(df, output_dir):
"""
Generates a Quality vs. Diversity Trade-off scatter plot.
Top-left corner represents the "Ideal Region" (Low KID, High Diversity).
"""
print("\nGenerating Quality vs Diversity scatter plot...")
fig, ax = plt.subplots(figsize=(12, 9))
for _, row in df.iterrows():
color = CONFIG_COLORS.get(row['config'], '#95a5a6')
marker = CATEGORY_MARKERS.get(row['category_type'], 'o')
ax.scatter(row['KID_mean'], row['IC_LPIPS_mean'],
c=color, marker=marker, s=200, alpha=0.8,
edgecolors='black', linewidth=1.5, zorder=3)
# Annotate class names
label_text = row['class'][:4] if len(row['class']) > 4 else row['class']
ax.annotate(label_text, (row['KID_mean'], row['IC_LPIPS_mean']),
fontsize=8, ha='center', va='bottom', xytext=(0, 8),
textcoords='offset points', fontweight='bold')
# Reference Median Lines
ax.axhline(y=df['IC_LPIPS_mean'].median(), color='gray', linestyle='--', alpha=0.5)
ax.axvline(x=df['KID_mean'].median(), color='gray', linestyle='--', alpha=0.5)
# Highlight the Ideal Region
kid_min, lpips_max = df['KID_mean'].min(), df['IC_LPIPS_mean'].max()
ax.annotate('Ideal Region\n(High Quality + High Diversity)',
xy=(kid_min, lpips_max), fontsize=11, color='#27ae60', fontweight='bold',
ha='left', va='top', bbox=dict(boxstyle='round,pad=0.3', facecolor='#d5f4e6', edgecolor='#27ae60'))
ax.set_xlabel('KID (Lower = Higher Fidelity)', fontsize=13, fontweight='bold')
ax.set_ylabel('IC-LPIPS (Higher = More Diverse)', fontsize=13, fontweight='bold')
ax.set_title('Quality vs Diversity Trade-off Analysis', fontsize=15, fontweight='bold', pad=15)
# Dynamic Legend Generation
legend_elements = [Line2D([0], [0], marker='o', color='w', markerfacecolor=c, markersize=12, label=f'Config {k.upper()}', markeredgecolor='black')
for k, c in CONFIG_COLORS.items() if k in df['config'].values]
ax.legend(handles=legend_elements, loc='upper right', framealpha=0.95)
ax.grid(True, alpha=0.3)
plt.tight_layout()
output_path = os.path.join(output_dir, 'scatter_tradeoff.png')
plt.savefig(output_path, dpi=300, bbox_inches='tight', facecolor='white')
print(f" Scatter plot saved to: {output_path}")
plt.close()
def create_grouped_bar_charts(df, output_dir):
"""Generates comparison bar charts across all classes and configs."""
print("\nGenerating comparison bar charts...")
fig, axes = plt.subplots(2, 1, figsize=(16, 12))
all_classes = df['class'].unique().tolist()
class_order = [c for c in OBJECT_CLASSES if c in all_classes] + \
[c for c in TEXTURE_CLASSES if c in all_classes]
config_order = ['base', 'tex', 'obj']
configs = [c for c in config_order if c in df['config'].values]
bar_width = 0.25
x = np.arange(len(class_order))
for idx, (metric, ax_title, ylabel) in enumerate([
('KID', 'KID Quality Evaluation - Config Comparison', 'KID (lower is better)'),
('IC_LPIPS', 'IC-LPIPS Diversity Evaluation - Config Comparison', 'IC-LPIPS (higher is better)')
]):
ax = axes[idx]
for i, config in enumerate(configs):
config_data = df[df['config'] == config].set_index('class')
values = [config_data.loc[c, f'{metric}_mean'] if c in config_data.index else 0 for c in class_order]
errors = [config_data.loc[c, f'{metric}_std'] if c in config_data.index else 0 for c in class_order]
ax.bar(x + i * bar_width, values, bar_width, label=f'Config {config.upper()}',
color=CONFIG_COLORS.get(config, '#95a5a6'), yerr=errors, capsize=3,
alpha=0.85, edgecolor='black', linewidth=0.5)
ax.set_title(ax_title, fontsize=14, fontweight='bold')
ax.set_ylabel(ylabel, fontweight='bold')
ax.set_xticks(x + bar_width * (len(configs) - 1) / 2)
ax.set_xticklabels(class_order, rotation=45, ha='right')
ax.legend()
plt.tight_layout()
output_path = os.path.join(output_dir, 'grouped_bar_charts.png')
plt.savefig(output_path, dpi=300, bbox_inches='tight')
print(f" Bar charts saved to: {output_path}")
plt.close()
def create_summary_table(df, output_dir):
"""Calculates and saves grouped summary statistics."""
print("\nCalculating summary statistics...")
summary = df.groupby(['category_type', 'config']).agg({
'KID_mean': ['mean', 'std', 'min', 'max'],
'IC_LPIPS_mean': ['mean', 'std', 'min', 'max']
}).round(4)
summary.columns = ['_'.join(col).strip() for col in summary.columns.values]
summary_path = os.path.join(output_dir, 'summary_statistics.csv')
summary.to_csv(summary_path)
print("\n" + "="*80 + "\nSummary Statistics:\n" + "="*80)
print(summary.to_string())
return summary
def main():
parser = argparse.ArgumentParser(description="Visualize DefectFill Experiment Results")
parser.add_argument("--csv_path", type=str, required=True, help="Path to evaluation results CSV")
parser.add_argument("--output_dir", type=str, default="./figures", help="Output directory for plots")
args = parser.parse_args()
os.makedirs(args.output_dir, exist_ok=True)
# Process all visualizations
df = load_and_validate_data(args.csv_path)
create_heatmaps(df, args.output_dir)
create_scatter_plot(df, args.output_dir)
create_grouped_bar_charts(df, args.output_dir)
create_summary_table(df, args.output_dir)
print(f"\nAll visualization charts generated in: {args.output_dir}")
if __name__ == "__main__":
main() |