논문의 Figure 4에 해당하는 내용 : Feature Adapter가 적용되기 전의 결과

from pathlib import Path
import matplotlib.pyplot as plt
import torch
import torch.nn as nn
import torch.nn.functional as F
from PIL import Image
from torchvision import models, transforms
# ------------------------------------------------------------------
# 설정 (Config) - 필요하면 아래 값만 수정하세요
# ------------------------------------------------------------------
DATA_ROOT = "./mvtec_ad/bottle"
SPLIT_DIR = "train/good"
MAX_IMAGES = 30
PATCH_SIZE = 3 # Eq.(1) 이웃 patch 크기
BINS = 100
PRETRAINED = True
OUT_PATH = "figure4_local_feature_std.png"
IMG_EXTENSIONS = (".png", ".jpg", ".jpeg", ".bmp")
# ------------------------------------------------------------------
# Feature Extractor: hierarchy level 2 + 3 결합 (논문 4.3절 기본 설정)
# ------------------------------------------------------------------
class FeatureExtractor(nn.Module):
def __init__(self, pretrained: bool = True):
super().__init__()
weights = models.Wide_ResNet50_2_Weights.IMAGENET1K_V2 if pretrained else None
backbone = models.wide_resnet50_2(weights=weights)
backbone.eval()
self.stem = nn.Sequential(backbone.conv1, backbone.bn1, backbone.relu, backbone.maxpool)
self.layer1 = backbone.layer1
self.layer2 = backbone.layer2 # hierarchy level 2
self.layer3 = backbone.layer3 # hierarchy level 3
for p in self.parameters():
p.requires_grad_(False)
@torch.no_grad()
def forward(self, x: torch.Tensor) -> dict:
x = self.stem(x)
x = self.layer1(x)
feat_l2 = self.layer2(x)
feat_l3 = self.layer3(feat_l2)
return {"level2": feat_l2, "level3": feat_l3}
def local_feature_aggregation(feat_map: torch.Tensor, patch_size: int = 3) -> torch.Tensor:
"""
Eq.(1)-(2): 이웃(patch_size x patch_size) 평균 풀링.
이 단계가 바로 논문이 정의하는 'local feature'입니다.
stride=1, same padding이라 공간 해상도(H, W)는 그대로 유지됩니다.
"""
pad = patch_size // 2
return F.avg_pool2d(feat_map, kernel_size=patch_size, stride=1, padding=pad)
def build_local_feature_map(hierarchy_feats: dict, patch_size: int = 3) -> torch.Tensor:
"""Eq.(3)-(4): 레벨별로 patch 집계 후, 가장 큰 해상도로 resize하여 채널 방향 결합."""
aggregated = {
name: local_feature_aggregation(f, patch_size) for name, f in hierarchy_feats.items()
}
target_h, target_w = aggregated["level2"].shape[-2:] # level2가 더 큰 해상도
resized = [
F.interpolate(f, size=(target_h, target_w), mode="bilinear", align_corners=False)
for f in aggregated.values()
]
return torch.cat(resized, dim=1) # (B, C=1536, H0, W0)
def load_image_tensor(image_path: Path) -> torch.Tensor:
preprocess = transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
img = Image.open(image_path).convert("RGB")
return preprocess(img).unsqueeze(0)
def main():
data_root = Path(DATA_ROOT)
img_dir = data_root / SPLIT_DIR
if not img_dir.exists():
raise SystemExit(f"❌ 이미지 폴더를 찾을 수 없습니다: {img_dir}")
image_paths = sorted(p for p in img_dir.iterdir() if p.suffix.lower() in IMG_EXTENSIONS)
if not image_paths:
raise SystemExit(f"❌ '{img_dir}'에서 이미지를 찾지 못했습니다.")
if MAX_IMAGES:
image_paths = image_paths[:MAX_IMAGES]
device = "cuda" if torch.cuda.is_available() else "cpu"
extractor = FeatureExtractor(pretrained=PRETRAINED).to(device)
print("🚀 데이터 로딩 및 특징 추출 시작 (local feature = patch 집계 O)")
print(f"📂 경로: {img_dir} | 🖼️ 이미지 수: {len(image_paths)}")
all_feats = []
for i, path in enumerate(image_paths):
x = load_image_tensor(path).to(device)
hierarchy_feats = extractor(x)
o = build_local_feature_map(hierarchy_feats, patch_size=PATCH_SIZE) # (1, C, H0, W0)
b, c, h, w = o.shape
flat = o.permute(0, 2, 3, 1).reshape(-1, c) # (H0*W0, C): 위치별 벡터를 행으로
all_feats.append(flat.cpu())
if (i + 1) % 5 == 0 or (i + 1) == len(image_paths):
print(f" [{i + 1}/{len(image_paths)}] 완료...")
all_feats = torch.cat(all_feats, dim=0) # (N_total_locations, C)
# 채널별 표준편차 계산 (Figure 4의 x축 값)
feature_std = all_feats.std(dim=0).numpy()
print("\n📊 결과 요약:")
print(f" - 전체 샘플 수 (위치 수): {all_feats.shape[0]}")
print(f" - 총 채널 수 (C): {len(feature_std)}")
print(f" - 평균 표준편차: {feature_std.mean():.6f}")
print(f" - 최소/최대: {feature_std.min():.6f} / {feature_std.max():.6f}")
# 히스토그램 시각화
plt.figure(figsize=(6, 5))
plt.hist(feature_std, bins=BINS, color="tab:orange", alpha=0.8, label="local")
plt.axvline(feature_std.mean(), color="red", linestyle="dashed", linewidth=1,
label=f"Mean: {feature_std.mean():.4f}")
plt.xlabel("Feature std.")
plt.ylabel("Number")
plt.title(f"Reproduced Figure 4: local feature std.\n({data_root.name})")
plt.legend()
plt.grid(axis="y", alpha=0.3)
plt.tight_layout()
plt.savefig(OUT_PATH, dpi=150)
print(f"\n✅ 히스토그램이 저장되었습니다: {OUT_PATH}")
if __name__ == "__main__":
main()