[논문 연구] Self-Supervised Learning for ASD - MIMII Dataset : 판단 근거 시각화 비교 (260904)

WonTerry·2026년 9월 4일

Deep Learning

목록 보기
33/41

normal #03번 데이터 - 정상 판단

abnormal #03번 데이터 - 비정상 판단

정규화전 Grad-CAM의 Peak값을 비교하면 비슷한 수준이다. 오히려 정상의 경우 값이 더 크다.
면적도 오히려 정상이 더 크게 나타난다.

어떤 근거로 정상과 비정상이라 판단하는가?

좀 더 돌려보면,

normal #04번 데이터 - 정상 판단

abnormal #04번 데이터 - 비정상 판단

이번 경우는 확실히 정규화 전 값의 크기가 상당히 큰 것을 알 수 있다.

판단 근거 시각화 코드

"""
./data/MIMII/abnormal 폴더에서 임의의 wav 파일 하나를 골라
1) 정상/비정상 여부를 판정하고
2) 스펙트로그램 브랜치(SpectrogramBranch)의 마지막 conv 블록에 대해
   Grad-CAM 방식으로 "어느 시간-주파수 영역이 이상 판정에 기여했는지"를
   컬러맵으로 오버랩하여 보여주는 스크립트.

anomaly score(코사인 거리)를 스칼라 타겟으로 역전파하여 얻은 그래디언트로
CAM(Class Activation Map과 동일한 원리, 여기서는 "이상치 활성화 맵")을 계산한다.

실행 예:
    python visualize_test_result.py
    python visualize_test_result.py --file ./data/MIMII/abnormal/xxx.wav
    python visualize_test_result.py --seed 42 --output result.png
"""
import argparse
import json
import os
import random

import numpy as np
import torch
import torch.nn.functional as F
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import matplotlib.font_manager as fm

# 한글 라벨이 깨지지 않도록 시스템에 설치된 CJK 폰트를 사용 (없으면 기본 폰트로 대체)
for _font_name in ["Noto Sans CJK KR", "Noto Sans CJK JP", "NanumGothic", "Malgun Gothic"]:
    if any(_font_name.lower() in f.name.lower() for f in fm.fontManager.ttflist):
        plt.rcParams["font.family"] = _font_name
        break
plt.rcParams["axes.unicode_minus"] = False

from mimii_dataset import load_wav_fixed_length, _list_wavs
from asd_model import TwoBranchEmbeddingNet
from test_asd_pytorch import compute_train_threshold, length_normalize


def parse_args():
    p = argparse.ArgumentParser()
    p.add_argument("--config", type=str, default="asd_config.json")
    p.add_argument("--model_path", type=str, default="asd_model_pytorch.pt")
    p.add_argument("--centers_path", type=str, default="kmeans_centers.npy")
    p.add_argument("--data_root", type=str, default=None,
                    help="지정하지 않으면 config.json에 저장된 경로를 사용")
    p.add_argument("--file", type=str, default=None,
                    help="특정 파일을 지정하면 abnormal 폴더에서 무작위로 고르지 않고 이 파일을 사용")
    p.add_argument("--threshold_percentile", type=float, default=90.0)
    p.add_argument("--seed", type=int, default=None, help="무작위 파일 선택 시드")
    p.add_argument("--output", type=str, default="anomaly_gradcam.png")
    p.add_argument("--device", type=str,
                    default="cuda" if torch.cuda.is_available() else "cpu")
    return p.parse_args()


def pick_random_abnormal_file(data_root: str, seed=None) -> str:
    files = _list_wavs(os.path.join(data_root, "abnormal"))
    if len(files) == 0:
        raise FileNotFoundError(f"{data_root}/abnormal 에서 wav 파일을 찾지 못했습니다.")
    rng = random.Random(seed)
    return rng.choice(files)


def compute_gradcam(model: TwoBranchEmbeddingNet, wav: torch.Tensor,
                     centers: np.ndarray, device: torch.device):
    """
    wav: (1, L) 텐서.
    반환: (anomaly_score, cam_norm(H,W) numpy, cam_raw(H,W) numpy,
           spec_db(F,T) numpy, used_center_idx)
    cam_norm: 0~1로 정규화된 CAM (파일 내 상대적 위치 비교용)
    cam_raw : 정규화 이전, 그래디언트-가중 활성화의 원래 스케일 (파일 간 반응 강도 비교용)
    """
    model.eval()
    wav = wav.to(device)

    activations = {}

    def fwd_hook(module, inp, out):
        out.retain_grad()
        activations["feat"] = out

    handle = model.spectrogram_branch.block3.register_forward_hook(fwd_hook)

    # 표시용 스펙트로그램(선형 magnitude, 크롭됨)은 grad 없이 별도로 계산
    with torch.no_grad():
        spec_display = model.raw_spectrogram(wav)  # (1,1,F,T)

    emb_fft, emb_mel = model.embed_clean(wav)  # grad 필요하므로 no_grad 밖에서 실행
    emb = torch.cat([emb_fft, emb_mel], dim=-1)
    emb_n = F.normalize(emb, dim=-1)

    centers_t = torch.from_numpy(length_normalize(centers)).float().to(device)
    cos_sim = emb_n @ centers_t.t()          # (1, K)
    cos_dist = 1 - cos_sim                    # (1, K)
    min_idx = int(cos_dist.argmin(dim=-1).item())
    score = cos_dist[0, min_idx]              # 실제 anomaly score로 사용되는 항

    model.zero_grad(set_to_none=True)
    score.backward()

    grad = activations["feat"].grad            # (1, C, h, w)
    act = activations["feat"].detach()          # (1, C, h, w)
    handle.remove()

    weights = grad.mean(dim=(2, 3), keepdim=True)          # (1, C, 1, 1) - GAP
    cam = F.relu((weights * act).sum(dim=1, keepdim=True))  # (1, 1, h, w)

    target_hw = spec_display.shape[-2:]
    cam = F.interpolate(cam, size=target_hw, mode="bilinear", align_corners=False)
    cam_raw = cam.squeeze().detach().cpu().numpy()          # 정규화 전 (원 스케일)
    cam_norm = cam_raw / (cam_raw.max() + 1e-8)              # 0~1로 정규화

    spec_db = 20 * torch.log10(spec_display.squeeze() + 1e-8).detach().cpu().numpy()

    return float(score.item()), cam_norm, cam_raw, spec_db, min_idx


def main():
    args = parse_args()
    device = torch.device(args.device)

    with open(args.config, "r") as f:
        config = json.load(f)
    sample_rate = config["sample_rate"]
    num_samples = config["num_samples"]
    data_root = args.data_root or config["data_root"]

    model = TwoBranchEmbeddingNet(num_samples=num_samples).to(device)
    model.load_state_dict(torch.load(args.model_path, map_location=device))
    model.eval()
    centers = np.load(args.centers_path)

    target_file = args.file or pick_random_abnormal_file(data_root, seed=args.seed)
    print(f"선택된 파일: {target_file}")

    wav_np = load_wav_fixed_length(target_file, sample_rate, num_samples, random_crop=False)
    wav = torch.from_numpy(wav_np).float().unsqueeze(0)  # (1, L)

    score, cam_norm, cam_raw, spec_db, center_idx = compute_gradcam(model, wav, centers, device)

    threshold = compute_train_threshold(
        model, device, centers, data_root, sample_rate, num_samples,
        args.threshold_percentile, batch_size=64)
    decision = "비정상 (anomalous)" if (threshold is not None and score > threshold) else "정상 (normal)"

    print(f"anomaly score = {score:.4f}   임계값(P{args.threshold_percentile:.0f}) = "
          f"{threshold:.4f}" if threshold is not None else f"anomaly score = {score:.4f}")
    print(f"판정: {decision}")

    # 주파수/시간 축 계산 (compute_spectrogram의 기본 crop: f_min_bin=13, f_max_bin=None)
    n_fft, hop = model.n_fft, model.hop
    all_freqs = np.fft.rfftfreq(n_fft, d=1.0 / sample_rate)
    freqs = all_freqs[13:13 + spec_db.shape[0]]
    times = np.arange(spec_db.shape[1]) * hop / sample_rate

    fig, axes = plt.subplots(1, 3, figsize=(19, 5))

    im0 = axes[0].imshow(spec_db, origin="lower", aspect="auto", cmap="magma",
                          extent=[times[0], times[-1], freqs[0], freqs[-1]])
    axes[0].set_title("입력 스펙트로그램 (모델 입력, dB)")
    axes[0].set_xlabel("시간 (초)")
    axes[0].set_ylabel("주파수 (Hz)")
    fig.colorbar(im0, ax=axes[0], fraction=0.046, pad=0.04, label="dB")

    axes[1].imshow(spec_db, origin="lower", aspect="auto", cmap="gray",
                    extent=[times[0], times[-1], freqs[0], freqs[-1]])
    im1 = axes[1].imshow(cam_norm, origin="lower", aspect="auto", cmap="jet", alpha=0.5,
                          vmin=0, vmax=1,
                          extent=[times[0], times[-1], freqs[0], freqs[-1]])
    axes[1].set_title("Grad-CAM (정규화, 0~1)\n파일 내 상대적 위치 비교용")
    axes[1].set_xlabel("시간 (초)")
    axes[1].set_ylabel("주파수 (Hz)")
    fig.colorbar(im1, ax=axes[1], fraction=0.046, pad=0.04, label="기여도 (정규화)")

    axes[2].imshow(spec_db, origin="lower", aspect="auto", cmap="gray",
                    extent=[times[0], times[-1], freqs[0], freqs[-1]])
    im2 = axes[2].imshow(cam_raw, origin="lower", aspect="auto", cmap="jet", alpha=0.5,
                          vmin=0, vmax=cam_raw.max(),
                          extent=[times[0], times[-1], freqs[0], freqs[-1]])
    axes[2].set_title(f"Grad-CAM (raw, 정규화 전)\n최대값={cam_raw.max():.3g} - 파일 간 반응 강도 비교용")
    axes[2].set_xlabel("시간 (초)")
    axes[2].set_ylabel("주파수 (Hz)")
    fig.colorbar(im2, ax=axes[2], fraction=0.046, pad=0.04, label="raw 기여도")

    fname = os.path.basename(target_file)
    thr_str = f"{threshold:.3f}" if threshold is not None else "N/A"
    fig.suptitle(f"파일: {fname}   |   판정: {decision}   |   "
                  f"score={score:.3f} (threshold={thr_str}, center#{center_idx})")
    fig.tight_layout(rect=[0, 0, 1, 0.92])
    fig.savefig(args.output, dpi=150)
    print(f"결과 저장: {args.output}")


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개의 댓글