-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathplotscatter.py
More file actions
126 lines (103 loc) · 4.44 KB
/
Copy pathplotscatter.py
File metadata and controls
126 lines (103 loc) · 4.44 KB
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
import torch
import torch.nn as nn
import torchvision.models as models
from torch.utils.data import DataLoader
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import os
from tqdm import tqdm
import albumentations as A
from albumentations.pytorch import ToTensorV2
# [중요] dataset_B.py의 SarHeightDataset 클래스를 import
try:
from dataset_B import SarHeightDataset
except ImportError:
print("Error: 'dataset_B.py' 파일을 찾을 수 없습니다. (SarHeightDataset 클래스 필요)")
exit()
# --- 1. 설정 (train_B.py와 동일해야 함) ---
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {DEVICE}")
BASE_DIR = "./data/training_dataset_B"
LABELS_CSV = "Y_labels.csv"
MODEL_PATH = "./pth_folder/best_model_height.pth"
MAX_SAMPLES_FOR_PLOT = 10000 # 10000개 샘플링
# ------------------------------------
# --- 2. [ ★★★ 수정된 부분 ★★★ ] ---
# 모델 아키텍처를 train_B.py와 100% 동일하게 정의
model = models.resnet34(weights='IMAGENET1K_V1') # 1. ImageNet 가중치로 시작
# 2. train_B.py와 동일하게 conv1 레이어 가중치 평균
with torch.no_grad():
original_weights = model.conv1.weight.clone()
new_conv1 = nn.Conv2d(2, 64, kernel_size=7, stride=2, padding=3, bias=False)
# 2개 채널에 대해 (R+G+B)/3 가중치 평균을 적용
new_conv1.weight[:, 0, :, :] = original_weights.mean(dim=1)
new_conv1.weight[:, 1, :, :] = original_weights.mean(dim=1)
model.conv1 = new_conv1
# 3. fc 레이어 수정
num_ftrs = model.fc.in_features
model.fc = nn.Linear(num_ftrs, 1)
print(f"Model: Modified ResNet34 (Input: 2-ch, Output: 1-ch Regression)")
# --- [ ★★★ 수정 끝 ★★★ ] ---
try:
# 4. 학습된 가중치 덮어쓰기
model.load_state_dict(torch.load(MODEL_PATH, map_location=DEVICE))
print(f"'{MODEL_PATH}' 모델 로드 성공.")
except Exception as e:
print(f"Error: 모델 로드 실패: {e}")
exit()
model = model.to(DEVICE)
model.eval()
# ---------------------------------
# --- 3. 데이터 로더 준비 (검증용) ---
IMG_SIZE = 256
val_transform = A.Compose([
A.Resize(IMG_SIZE, IMG_SIZE),
ToTensorV2()
])
try:
labels_df = pd.read_csv(os.path.join(BASE_DIR, LABELS_CSV))
if len(labels_df) == 0: raise FileNotFoundError("Y_labels.csv is empty.")
except FileNotFoundError as e:
print(f"Error: {os.path.join(BASE_DIR, LABELS_CSV)} 파일을 찾을 수 없습니다: {e}")
exit()
if len(labels_df) > MAX_SAMPLES_FOR_PLOT:
print(f"총 {len(labels_df)}개 칩 중 {MAX_SAMPLES_FOR_PLOT}개만 무작위 샘플링합니다...")
labels_df = labels_df.sample(n=MAX_SAMPLES_FOR_PLOT, random_state=42)
else:
print(f"총 {len(labels_df)}개 칩으로 예측을 수행합니다...")
labels_df = labels_df.reset_index(drop=True)
val_dataset = SarHeightDataset(data_dir=BASE_DIR, labels_df=labels_df, transform=val_transform)
val_loader = DataLoader(val_dataset, batch_size=16, shuffle=False)
# ---------------------------------
# --- 4. 모델 예측 수행 ---
all_predictions = []
all_actuals = []
with torch.no_grad():
for images, heights in tqdm(val_loader, desc="Predicting", total=len(val_loader)):
images = images.to(DEVICE)
outputs = model(images)
all_predictions.extend(outputs.cpu().numpy())
all_actuals.extend(heights.numpy())
all_predictions = np.array(all_predictions).flatten()
all_actuals = np.array(all_actuals).flatten()
# ---------------------------------
# --- 5. 산점도 (Scatter Plot) 그리기 ---
print("예측 완료. 산점도를 생성합니다...")
plt.figure(figsize=(10, 10))
plt.scatter(all_actuals, all_predictions, alpha=0.3, label='Predictions (Sampled)')
max_val = max(all_actuals.max(), all_predictions.max())
min_val = min(all_actuals.min(), all_predictions.min())
plt.plot([min_val, max_val], [min_val, max_val], 'r--', lw=2, label='Perfect Fit (Y=X)')
plt.title(f'모듈 B (높이 분석) 예측 결과 (샘플 {len(all_actuals)}개)', fontsize=16)
plt.xlabel('실제 높이 (m) - Ground Truth', fontsize=12)
plt.ylabel('AI 예측 높이 (m) - Predicted', fontsize=12)
plt.legend()
plt.grid(True, linestyle='--', alpha=0.6)
plt.axis('equal')
plt.tight_layout()
output_filename = "module_B_scatter_plot_FIXED.png" # 새 이름으로 저장
plt.savefig(output_filename)
print("="*30)
print(f"✅ 산점도 저장 완료: {output_filename}")
print("이 이미지를 논문의 '그림 8'로 사용하실 수 있습니다.")