[모두의 연구소] Wk_4 SimpleNet - Local Feature 특성 - 채널별 표준 편차 (260809)

WonTerry·2026년 8월 9일

Deep Learning

목록 보기
21/32

논문의 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()
profile
Hello, I'm Terry! 👋 Enjoy every moment of your life! 🌱 My current interests are Signal processing, Machine learning, Python, Database, LLM & RAG, MCP & ADK, Multi-Agents, Physical AI, ROS2...

0개의 댓글