重置所有文件与main一致,新增app.py(Gradio推理前端)
This commit is contained in:
@@ -0,0 +1,145 @@
|
||||
"""
|
||||
baseline/HOG_Baseline.py
|
||||
HOG + 颜色直方图特征提取 + LogisticRegression 四分类
|
||||
纯传统 CV/ML 基线,零神经网络依赖
|
||||
可独立运行,也可被 compare_models.py 导入
|
||||
author: yukun-hh
|
||||
date: 2026-5-14
|
||||
"""
|
||||
import sys, os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from tqdm import tqdm
|
||||
import matplotlib.pyplot as plt
|
||||
import matplotlib
|
||||
|
||||
from skimage.feature import hog
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
from sklearn.metrics import (
|
||||
accuracy_score, f1_score,
|
||||
confusion_matrix, ConfusionMatrixDisplay,
|
||||
classification_report,
|
||||
)
|
||||
|
||||
matplotlib.rcParams['font.sans-serif'] = ['SimHei', 'DejaVu Sans']
|
||||
matplotlib.rcParams['axes.unicode_minus'] = False
|
||||
|
||||
# ============================================================
|
||||
# ★★★ 可配置参数 ★★★
|
||||
# ============================================================
|
||||
DATA_ROOT = '../../trash_division_data/ultimate_4_class/'
|
||||
IMAGE_SIZE = 128
|
||||
HOG_ORIENTATIONS = 9
|
||||
HOG_PIXELS_PER_CELL = (8, 8)
|
||||
HOG_CELLS_PER_BLOCK = (2, 2)
|
||||
COLOR_BINS = 32
|
||||
# ============================================================
|
||||
|
||||
CLASS_NAMES = ['厨余垃圾', '可回收物', '其他垃圾', '有害垃圾']
|
||||
NUM_CLASSES = 4
|
||||
|
||||
|
||||
def extract_hog_color(image):
|
||||
img = image.convert('RGB').resize((IMAGE_SIZE, IMAGE_SIZE))
|
||||
arr = np.array(img, dtype=np.float64) / 255.0
|
||||
|
||||
hog_feat = hog(arr, orientations=HOG_ORIENTATIONS,
|
||||
pixels_per_cell=HOG_PIXELS_PER_CELL,
|
||||
cells_per_block=HOG_CELLS_PER_BLOCK,
|
||||
channel_axis=2, feature_vector=True)
|
||||
|
||||
color_feat = []
|
||||
for c in range(3):
|
||||
hist, _ = np.histogram(arr[:, :, c], bins=COLOR_BINS, range=(0, 1))
|
||||
color_feat.append(hist)
|
||||
color_feat = np.concatenate(color_feat)
|
||||
|
||||
return np.concatenate([hog_feat, color_feat])
|
||||
|
||||
|
||||
class HOGLRBaseline:
|
||||
def __init__(self, data_root=DATA_ROOT, image_size=IMAGE_SIZE):
|
||||
self.data_root = data_root
|
||||
self.image_size = image_size
|
||||
self.clf = LogisticRegression(
|
||||
C=1.0, max_iter=1000, solver='lbfgs', n_jobs=-1,
|
||||
)
|
||||
|
||||
def _load_data(self, split):
|
||||
dir_path = os.path.join(self.data_root, split)
|
||||
features, labels = [], []
|
||||
for class_id in range(1, NUM_CLASSES + 1):
|
||||
class_dir = os.path.join(dir_path, str(class_id))
|
||||
if not os.path.isdir(class_dir):
|
||||
continue
|
||||
files = sorted(os.listdir(class_dir))
|
||||
for fname in tqdm(files, desc=f'{split}/class_{class_id}'):
|
||||
fpath = os.path.join(class_dir, fname)
|
||||
try:
|
||||
with Image.open(fpath) as img:
|
||||
feat = extract_hog_color(img)
|
||||
features.append(feat)
|
||||
labels.append(class_id - 1)
|
||||
except Exception:
|
||||
pass
|
||||
print(f" {split}: {len(features)} 张")
|
||||
return np.array(features, dtype=np.float32), np.array(labels)
|
||||
|
||||
def fit(self, train_dir=None):
|
||||
if train_dir is None:
|
||||
train_dir = 'train'
|
||||
print(" 提取训练集 HOG 特征 ...")
|
||||
X, y = self._load_data(train_dir)
|
||||
self.clf.fit(X, y)
|
||||
|
||||
def predict(self, val_dir=None):
|
||||
if val_dir is None:
|
||||
val_dir = 'val'
|
||||
print(" 提取验证集 HOG 特征 ...")
|
||||
X, y = self._load_data(val_dir)
|
||||
preds = self.clf.predict(X)
|
||||
probs = self.clf.predict_proba(X)
|
||||
return y, preds, probs
|
||||
|
||||
|
||||
# ============================================================
|
||||
# compare_models.py 导入接口
|
||||
# ============================================================
|
||||
|
||||
def get_hog_lr_preds(train_loader, val_loader, device):
|
||||
baseline = HOGLRBaseline()
|
||||
baseline.fit('train')
|
||||
return baseline.predict('val')
|
||||
|
||||
|
||||
# ============================================================
|
||||
# 独立运行入口
|
||||
# ============================================================
|
||||
|
||||
if __name__ == '__main__':
|
||||
out_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
print("HOG + LogisticRegression 基线")
|
||||
baseline = HOGLRBaseline()
|
||||
|
||||
baseline.fit('train')
|
||||
y_true, y_preds, y_probs = baseline.predict('val')
|
||||
|
||||
acc = accuracy_score(y_true, y_preds)
|
||||
macro_f1 = f1_score(y_true, y_preds, average='macro')
|
||||
print(f"\n验证集 Accuracy: {acc:.4f}")
|
||||
print(f"验证集 Macro-F1: {macro_f1:.4f}")
|
||||
print(f"\n分类报告:\n{classification_report(y_true, y_preds, target_names=CLASS_NAMES)}")
|
||||
|
||||
cm = confusion_matrix(y_true, y_preds)
|
||||
fig, ax = plt.subplots(figsize=(8, 7))
|
||||
ConfusionMatrixDisplay(cm, display_labels=CLASS_NAMES).plot(
|
||||
ax=ax, cmap='Blues', values_format='d', xticks_rotation=30)
|
||||
ax.set_title('HOG + LogisticRegression 混淆矩阵', fontsize=14)
|
||||
plt.tight_layout()
|
||||
cm_path = os.path.join(out_dir, 'hog_lr_confusion_matrix.png')
|
||||
plt.savefig(cm_path, dpi=150, bbox_inches='tight')
|
||||
plt.show()
|
||||
print(f"混淆矩阵已保存: {cm_path}")
|
||||
@@ -0,0 +1,278 @@
|
||||
"""
|
||||
baseline/ResNet34_Pretrained_10pct.py
|
||||
ResNet-34 ImageNet 预训练权重 + 10% 训练集微调
|
||||
可独立运行训练,也可被 compare_models.py 导入
|
||||
author: yukun-hh
|
||||
date: 2026-5-14
|
||||
"""
|
||||
import sys, os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import random
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.optim as optim
|
||||
from torch.utils.data import DataLoader, Subset
|
||||
from torchvision import models, transforms
|
||||
from tqdm import tqdm
|
||||
import csv
|
||||
import matplotlib.pyplot as plt
|
||||
import matplotlib
|
||||
|
||||
from Dataloader import RobustImageFolder
|
||||
|
||||
matplotlib.rcParams['font.sans-serif'] = ['SimHei', 'DejaVu Sans']
|
||||
matplotlib.rcParams['axes.unicode_minus'] = False
|
||||
|
||||
# ============================================================
|
||||
# ★★★ 可配置参数 ★★★
|
||||
# ============================================================
|
||||
DATA_ROOT = '../../trash_division_data/ultimate_4_class/'
|
||||
BATCH_SIZE = 32
|
||||
IMAGE_SIZE = 256
|
||||
NUM_WORKERS = 4
|
||||
EPOCHS = 30
|
||||
LR = 0.001
|
||||
TRAIN_PCT = 0.1
|
||||
SEED = 42
|
||||
DROPOUT = 0.3
|
||||
MODEL_SAVE_PATH = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'resnet34_10pct.pth')
|
||||
LOG_PATH = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'resnet34_10pct_log.csv')
|
||||
# ============================================================
|
||||
|
||||
NUM_CLASSES = 4
|
||||
CLASS_NAMES = ['厨余垃圾', '可回收物', '其他垃圾', '有害垃圾']
|
||||
|
||||
|
||||
class PretrainedResNet34(nn.Module):
|
||||
def __init__(self, num_classes=NUM_CLASSES, dropout=DROPOUT):
|
||||
super().__init__()
|
||||
self.backbone = models.resnet34(weights='IMAGENET1K_V1')
|
||||
in_features = self.backbone.fc.in_features
|
||||
self.backbone.fc = nn.Identity()
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.fc = nn.Linear(in_features, num_classes)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.backbone(x)
|
||||
x = self.dropout(x)
|
||||
x = self.fc(x)
|
||||
return x
|
||||
|
||||
def freeze_early_layers(self):
|
||||
for param in self.backbone.conv1.parameters():
|
||||
param.requires_grad = False
|
||||
for param in self.backbone.bn1.parameters():
|
||||
param.requires_grad = False
|
||||
for param in self.backbone.layer1.parameters():
|
||||
param.requires_grad = False
|
||||
for param in self.backbone.layer2.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def print_trainable_info(self):
|
||||
frozen = sum(p.numel() for p in self.parameters() if not p.requires_grad)
|
||||
trainable = sum(p.numel() for p in self.parameters() if p.requires_grad)
|
||||
total = frozen + trainable
|
||||
print(f" 冻结参数: {frozen:,} 可训练参数: {trainable:,} ({100.*trainable/total:.1f}%)")
|
||||
|
||||
|
||||
def compute_macro_f1(predicted, targets, num_classes=NUM_CLASSES):
|
||||
tp = torch.zeros(num_classes, device=predicted.device)
|
||||
fp = torch.zeros(num_classes, device=predicted.device)
|
||||
fn = torch.zeros(num_classes, device=predicted.device)
|
||||
for c in range(num_classes):
|
||||
tp[c] = ((predicted == c) & (targets == c)).sum()
|
||||
fp[c] = ((predicted == c) & (targets != c)).sum()
|
||||
fn[c] = ((predicted != c) & (targets == c)).sum()
|
||||
precision = tp / (tp + fp + 1e-8)
|
||||
recall = tp / (tp + fn + 1e-8)
|
||||
f1 = 2 * precision * recall / (precision + recall + 1e-8)
|
||||
return f1.mean().item()
|
||||
|
||||
|
||||
def train_one_epoch(model, loader, criterion, optimizer, device, epoch):
|
||||
model.train()
|
||||
running_loss, correct, total = 0.0, 0, 0
|
||||
all_preds, all_labels = [], []
|
||||
pbar = tqdm(loader, desc=f'Epoch {epoch+1} [Train]')
|
||||
for images, labels in pbar:
|
||||
images, labels = images.to(device), labels.to(device)
|
||||
outputs = model(images)
|
||||
loss = criterion(outputs, labels)
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
running_loss += loss.item() * images.size(0)
|
||||
_, predicted = outputs.max(1)
|
||||
total += labels.size(0)
|
||||
correct += predicted.eq(labels).sum().item()
|
||||
all_preds.append(predicted)
|
||||
all_labels.append(labels)
|
||||
batch_f1 = compute_macro_f1(predicted, labels)
|
||||
pbar.set_postfix({'loss': loss.item(), 'F1': f'{batch_f1:.4f}',
|
||||
'Acc': f'{100.*correct/total:.2f}%'})
|
||||
epoch_loss = running_loss / total
|
||||
epoch_f1 = compute_macro_f1(torch.cat(all_preds), torch.cat(all_labels))
|
||||
epoch_acc = 100. * correct / total
|
||||
return epoch_loss, epoch_f1, epoch_acc
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def validate(model, loader, criterion, device):
|
||||
model.eval()
|
||||
running_loss, correct, total = 0.0, 0, 0
|
||||
all_preds, all_labels = [], []
|
||||
for images, labels in tqdm(loader, desc='[Validate]'):
|
||||
images, labels = images.to(device), labels.to(device)
|
||||
outputs = model(images)
|
||||
loss = criterion(outputs, labels)
|
||||
running_loss += loss.item() * images.size(0)
|
||||
_, predicted = outputs.max(1)
|
||||
total += labels.size(0)
|
||||
correct += predicted.eq(labels).sum().item()
|
||||
all_preds.append(predicted)
|
||||
all_labels.append(labels)
|
||||
epoch_loss = running_loss / total
|
||||
epoch_f1 = compute_macro_f1(torch.cat(all_preds), torch.cat(all_labels))
|
||||
epoch_acc = 100. * correct / total
|
||||
return epoch_loss, epoch_f1, epoch_acc
|
||||
|
||||
|
||||
def train_model(model, train_loader, val_loader, device, epochs=EPOCHS, lr=LR):
|
||||
criterion = nn.CrossEntropyLoss()
|
||||
optimizer = optim.SGD(filter(lambda p: p.requires_grad, model.parameters()),
|
||||
lr=lr, momentum=0.9, weight_decay=1e-4)
|
||||
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)
|
||||
|
||||
history = {'train_loss': [], 'train_f1': [], 'train_acc': [],
|
||||
'val_loss': [], 'val_f1': [], 'val_acc': []}
|
||||
best_val_f1 = 0.0
|
||||
|
||||
log_file = open(LOG_PATH, 'w', newline='')
|
||||
log_writer = csv.writer(log_file)
|
||||
log_writer.writerow(['epoch', 'train_loss', 'train_f1', 'train_acc',
|
||||
'val_loss', 'val_f1', 'val_acc', 'lr', 'best'])
|
||||
|
||||
for epoch in range(epochs):
|
||||
print(f'\n{"="*50}')
|
||||
print(f'Epoch {epoch+1}/{epochs}')
|
||||
|
||||
train_loss, train_f1, train_acc = train_one_epoch(
|
||||
model, train_loader, criterion, optimizer, device, epoch)
|
||||
val_loss, val_f1, val_acc = validate(model, val_loader, criterion, device)
|
||||
scheduler.step()
|
||||
|
||||
history['train_loss'].append(train_loss)
|
||||
history['train_f1'].append(train_f1)
|
||||
history['train_acc'].append(train_acc)
|
||||
history['val_loss'].append(val_loss)
|
||||
history['val_f1'].append(val_f1)
|
||||
history['val_acc'].append(val_acc)
|
||||
|
||||
print(f'Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.2f}% | Train Macro-F1: {train_f1:.4f}')
|
||||
print(f'Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.2f}% | Val Macro-F1: {val_f1:.4f}')
|
||||
print(f'Learning Rate: {optimizer.param_groups[0]["lr"]:.6f}')
|
||||
|
||||
best_mark = ''
|
||||
if val_f1 > best_val_f1:
|
||||
best_val_f1 = val_f1
|
||||
torch.save(model.state_dict(), MODEL_SAVE_PATH)
|
||||
best_mark = 'best'
|
||||
print(f'✓ 保存最佳模型 (Macro-F1: {val_f1:.4f})')
|
||||
|
||||
lr_val = optimizer.param_groups[0]['lr']
|
||||
log_writer.writerow([epoch+1, train_loss, train_f1, train_acc,
|
||||
val_loss, val_f1, val_acc, lr_val, best_mark])
|
||||
log_file.flush()
|
||||
|
||||
log_file.close()
|
||||
print(f'\n训练完成!最佳验证 Macro-F1: {best_val_f1:.4f}')
|
||||
return history
|
||||
|
||||
|
||||
# ============================================================
|
||||
# compare_models.py 导入接口
|
||||
# ============================================================
|
||||
|
||||
def get_resnet34_10pct_preds(train_loader, val_loader, device):
|
||||
model = PretrainedResNet34(num_classes=NUM_CLASSES)
|
||||
model.load_state_dict(torch.load(MODEL_SAVE_PATH, map_location='cpu'))
|
||||
model = model.to(device).eval()
|
||||
|
||||
y_true, y_preds, y_probs = [], [], []
|
||||
with torch.no_grad():
|
||||
for images, labels in tqdm(val_loader, desc='ResNet-34 (10%)'):
|
||||
images, labels = images.to(device), labels
|
||||
logits = model(images)
|
||||
probs = torch.softmax(logits, dim=1)
|
||||
preds = probs.argmax(dim=1)
|
||||
y_true.append(labels.numpy())
|
||||
y_preds.append(preds.cpu().numpy())
|
||||
y_probs.append(probs.cpu().numpy())
|
||||
return np.concatenate(y_true), np.concatenate(y_preds), np.concatenate(y_probs)
|
||||
|
||||
|
||||
# ============================================================
|
||||
# 独立训练入口
|
||||
# ============================================================
|
||||
|
||||
if __name__ == '__main__':
|
||||
random.seed(SEED)
|
||||
np.random.seed(SEED)
|
||||
torch.manual_seed(SEED)
|
||||
|
||||
device = torch.device('cuda' if torch.cuda.is_available()
|
||||
else 'xpu' if hasattr(torch, 'xpu') and torch.xpu.is_available()
|
||||
else 'cpu')
|
||||
print(f"Device: {device}")
|
||||
|
||||
val_transform = transforms.Compose([
|
||||
transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=[0.485, 0.456, 0.406],
|
||||
std=[0.229, 0.224, 0.225]),
|
||||
])
|
||||
train_transform = transforms.Compose([
|
||||
transforms.RandomResizedCrop(IMAGE_SIZE, scale=(0.8, 1.0)),
|
||||
transforms.RandomHorizontalFlip(p=0.5),
|
||||
transforms.RandomRotation(degrees=15),
|
||||
transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=[0.485, 0.456, 0.406],
|
||||
std=[0.229, 0.224, 0.225]),
|
||||
])
|
||||
|
||||
full_train_dataset = RobustImageFolder(
|
||||
root=os.path.join(DATA_ROOT, 'train'),
|
||||
transform=train_transform,
|
||||
)
|
||||
val_dataset = RobustImageFolder(
|
||||
root=os.path.join(DATA_ROOT, 'val'),
|
||||
transform=val_transform,
|
||||
)
|
||||
|
||||
n_train = len(full_train_dataset)
|
||||
n_subset = max(1, int(n_train * TRAIN_PCT))
|
||||
indices = random.sample(range(n_train), n_subset)
|
||||
train_dataset = Subset(full_train_dataset, indices)
|
||||
print(f"训练集: {len(train_dataset)} / {n_train} ({TRAIN_PCT*100:.0f}%)")
|
||||
print(f"验证集: {len(val_dataset)}")
|
||||
|
||||
train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE,
|
||||
shuffle=True, num_workers=NUM_WORKERS,
|
||||
pin_memory=True, drop_last=True)
|
||||
val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE,
|
||||
shuffle=False, num_workers=NUM_WORKERS,
|
||||
pin_memory=True, drop_last=False)
|
||||
|
||||
model = PretrainedResNet34(num_classes=NUM_CLASSES, dropout=DROPOUT)
|
||||
model.freeze_early_layers()
|
||||
model.print_trainable_info()
|
||||
model = model.to(device)
|
||||
|
||||
history = train_model(model, train_loader, val_loader, device, epochs=EPOCHS, lr=LR)
|
||||
|
||||
model.load_state_dict(torch.load(MODEL_SAVE_PATH, map_location='cpu'))
|
||||
print(f"模型已保存: {MODEL_SAVE_PATH}")
|
||||
print(f"训练日志已保存: {LOG_PATH}")
|
||||
@@ -0,0 +1,145 @@
|
||||
"""
|
||||
baseline/VGG_KNN.py
|
||||
VGG16 预训练模型特征提取 + KNN 四分类基线
|
||||
可独立运行,也可被 compare_models.py 导入复用
|
||||
author: yukun-hh
|
||||
date: 2026-5-14
|
||||
"""
|
||||
import sys, os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import matplotlib
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.utils.data import DataLoader
|
||||
from torchvision import models, transforms
|
||||
from tqdm import tqdm
|
||||
|
||||
from sklearn.neighbors import KNeighborsClassifier
|
||||
from sklearn.metrics import (
|
||||
accuracy_score, f1_score,
|
||||
confusion_matrix, ConfusionMatrixDisplay,
|
||||
classification_report,
|
||||
)
|
||||
|
||||
from Dataloader import RobustImageFolder
|
||||
|
||||
matplotlib.rcParams['font.sans-serif'] = ['SimHei', 'DejaVu Sans']
|
||||
matplotlib.rcParams['axes.unicode_minus'] = False
|
||||
|
||||
|
||||
CLASS_NAMES = ['厨余垃圾', '可回收物', '其他垃圾', '有害垃圾']
|
||||
|
||||
|
||||
def load_vgg16_extractor(device):
|
||||
try:
|
||||
model = models.vgg16(weights='IMAGENET1K_V1')
|
||||
except TypeError:
|
||||
model = models.vgg16(pretrained=True)
|
||||
model.classifier = nn.Identity()
|
||||
model = model.to(device).eval()
|
||||
for param in model.parameters():
|
||||
param.requires_grad = False
|
||||
return model
|
||||
|
||||
|
||||
def extract_features(model, loader, device):
|
||||
model.eval()
|
||||
all_features = []
|
||||
all_labels = []
|
||||
with torch.no_grad():
|
||||
for images, labels in tqdm(loader, desc='Extracting features'):
|
||||
images = images.to(device)
|
||||
feats = model(images)
|
||||
all_features.append(feats.cpu().numpy())
|
||||
all_labels.append(labels.numpy())
|
||||
return np.concatenate(all_features), np.concatenate(all_labels)
|
||||
|
||||
|
||||
class VGGKNNBaseline:
|
||||
def __init__(self, k=5, device='cpu',
|
||||
data_root='../trash_division_data/ultimate_4_class/',
|
||||
image_size=256, batch_size=32, num_workers=4):
|
||||
self.k = k
|
||||
self.device = device
|
||||
self.data_root = data_root
|
||||
self.image_size = image_size
|
||||
self.batch_size = batch_size
|
||||
self.num_workers = num_workers
|
||||
self.extractor = load_vgg16_extractor(device)
|
||||
self.knn = KNeighborsClassifier(n_neighbors=k, n_jobs=-1)
|
||||
|
||||
def _get_loader(self, split):
|
||||
transform = transforms.Compose([
|
||||
transforms.Resize((self.image_size, self.image_size)),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=[0.485, 0.456, 0.406],
|
||||
std=[0.229, 0.224, 0.225]),
|
||||
])
|
||||
dataset = RobustImageFolder(
|
||||
root=os.path.join(self.data_root, split),
|
||||
transform=transform,
|
||||
)
|
||||
print(f" {split}: {len(dataset)} 张")
|
||||
return DataLoader(dataset, batch_size=self.batch_size,
|
||||
shuffle=False, num_workers=self.num_workers,
|
||||
pin_memory=True, drop_last=False)
|
||||
|
||||
def fit(self, train_loader=None):
|
||||
if train_loader is None:
|
||||
train_loader = self._get_loader('train')
|
||||
print(" 提取训练集特征 ...")
|
||||
train_feats, train_labels = extract_features(self.extractor, train_loader, self.device)
|
||||
self.knn.fit(train_feats, train_labels)
|
||||
|
||||
def predict(self, val_loader=None):
|
||||
if val_loader is None:
|
||||
val_loader = self._get_loader('val')
|
||||
print(" 提取验证集特征 ...")
|
||||
val_feats, val_labels = extract_features(self.extractor, val_loader, self.device)
|
||||
preds = self.knn.predict(val_feats)
|
||||
probs = self.knn.predict_proba(val_feats)
|
||||
return val_labels, preds, probs
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
DATA_ROOT = '../trash_division_data/ultimate_4_class/'
|
||||
BATCH_SIZE = 32
|
||||
IMAGE_SIZE = 256
|
||||
NUM_WORKERS = 4
|
||||
K = 5
|
||||
|
||||
device = torch.device('cuda' if torch.cuda.is_available()
|
||||
else 'xpu' if hasattr(torch, 'xpu') and torch.xpu.is_available()
|
||||
else 'cpu')
|
||||
print(f"Device: {device}")
|
||||
|
||||
baseline = VGGKNNBaseline(k=K, device=device,
|
||||
data_root=DATA_ROOT, image_size=IMAGE_SIZE,
|
||||
batch_size=BATCH_SIZE, num_workers=NUM_WORKERS)
|
||||
|
||||
train_loader = baseline._get_loader('train')
|
||||
val_loader = baseline._get_loader('val')
|
||||
|
||||
baseline.fit(train_loader)
|
||||
y_true, y_preds, y_probs = baseline.predict(val_loader)
|
||||
|
||||
acc = accuracy_score(y_true, y_preds)
|
||||
macro_f1 = f1_score(y_true, y_preds, average='macro')
|
||||
print(f"\n验证集 Accuracy: {acc:.4f}")
|
||||
print(f"验证集 Macro-F1: {macro_f1:.4f}")
|
||||
print(f"\n分类报告:\n{classification_report(y_true, y_preds, target_names=CLASS_NAMES)}")
|
||||
|
||||
cm = confusion_matrix(y_true, y_preds)
|
||||
fig, ax = plt.subplots(figsize=(8, 7))
|
||||
ConfusionMatrixDisplay(cm, display_labels=CLASS_NAMES).plot(
|
||||
ax=ax, cmap='Blues', values_format='d', xticks_rotation=30)
|
||||
ax.set_title(f'Baseline Confusion Matrix (VGG16 + KNN, K={K})', fontsize=14)
|
||||
plt.tight_layout()
|
||||
out_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'vgg_knn_confusion_matrix.png')
|
||||
plt.savefig(out_path, dpi=150, bbox_inches='tight')
|
||||
plt.show()
|
||||
print(f"混淆矩阵已保存: {out_path}")
|
||||
@@ -0,0 +1 @@
|
||||
# baseline package
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 45 KiB |
@@ -0,0 +1,249 @@
|
||||
"""
|
||||
baseline/compare_models.py
|
||||
多模型对比:ROC 曲线 + 准确率柱状图
|
||||
添加新模型只需在 MODELS 列表加一行,无需修改绘图代码
|
||||
author: yukun-hh
|
||||
date: 2026-5-14
|
||||
"""
|
||||
import sys, os, re
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import matplotlib
|
||||
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
from torchvision import transforms
|
||||
from tqdm import tqdm
|
||||
|
||||
from sklearn.metrics import (
|
||||
roc_curve, auc, accuracy_score,
|
||||
precision_recall_curve, average_precision_score,
|
||||
)
|
||||
|
||||
from Model import Net
|
||||
from Dataloader import RobustImageFolder
|
||||
from baseline.VGG_KNN import VGGKNNBaseline
|
||||
from baseline.ResNet34_Pretrained_10pct import get_resnet34_10pct_preds
|
||||
from baseline.HOG_Baseline import get_hog_lr_preds
|
||||
|
||||
matplotlib.rcParams['font.sans-serif'] = ['SimHei', 'DejaVu Sans']
|
||||
matplotlib.rcParams['axes.unicode_minus'] = False
|
||||
|
||||
# ============================================================
|
||||
# ★★★ 可配置参数 ★★★
|
||||
# ============================================================
|
||||
DATA_ROOT = '../../trash_division_data/ultimate_4_class/'
|
||||
BATCH_SIZE = 32
|
||||
IMAGE_SIZE = 256
|
||||
NUM_WORKERS = 4
|
||||
K_KNN = 5
|
||||
# ============================================================
|
||||
|
||||
CLASS_NAMES = ['厨余垃圾', '可回收物', '其他垃圾', '有害垃圾']
|
||||
NUM_CLASSES = 4
|
||||
|
||||
# ============================================================
|
||||
# 预测函数 — 每个函数签名: (train_loader, val_loader, device) -> (y_true, y_preds, y_probs)
|
||||
# ============================================================
|
||||
|
||||
def get_resnet34_preds(train_loader, val_loader, device):
|
||||
model = Net(num_classes=NUM_CLASSES)
|
||||
state_dict = torch.load('../best_model.pth', map_location='cpu')
|
||||
if 'model_state_dict' in state_dict:
|
||||
state_dict = state_dict['model_state_dict']
|
||||
elif 'model' in state_dict:
|
||||
state_dict = state_dict['model']
|
||||
model.load_state_dict(state_dict)
|
||||
model = model.to(device).eval()
|
||||
|
||||
y_true, y_preds, y_probs = [], [], []
|
||||
with torch.no_grad():
|
||||
for images, labels in tqdm(val_loader, desc='ResNet-34'):
|
||||
images, labels = images.to(device), labels
|
||||
logits = model(images)
|
||||
probs = torch.softmax(logits, dim=1)
|
||||
preds = probs.argmax(dim=1)
|
||||
y_true.append(labels.numpy())
|
||||
y_preds.append(preds.cpu().numpy())
|
||||
y_probs.append(probs.cpu().numpy())
|
||||
return np.concatenate(y_true), np.concatenate(y_preds), np.concatenate(y_probs)
|
||||
|
||||
|
||||
def get_vgg_knn_preds(train_loader, val_loader, device):
|
||||
baseline = VGGKNNBaseline(k=K_KNN, device=device)
|
||||
baseline.fit(train_loader)
|
||||
return baseline.predict(val_loader)
|
||||
|
||||
|
||||
# ============================================================
|
||||
# ★ 模型注册表 — 添加新模型只需在这里加一行 ★
|
||||
# ============================================================
|
||||
|
||||
MODELS = [
|
||||
('ResNet-34', get_resnet34_preds),
|
||||
('ResNet-34 (10% Fine-tune)', get_resnet34_10pct_preds),
|
||||
('VGG16 + KNN (K=5)', get_vgg_knn_preds),
|
||||
('HOG + LogisticRegression', get_hog_lr_preds),
|
||||
# 未来轻松扩展示例:
|
||||
# ('ResNet-18 (pretrained)', get_resnet18_preds),
|
||||
# ('ResNet-50 (pretrained)', get_resnet50_preds),
|
||||
# ('ResNet-34 (finetuned)', get_finetuned_preds),
|
||||
]
|
||||
|
||||
# ============================================================
|
||||
# 调色板 (扩展时无需修改)
|
||||
# ============================================================
|
||||
COLORS = ['#1f77b4', '#ff7f0e', '#2ca02c', '#d62728', '#9467bd', '#8c564b',
|
||||
'#e377c2', '#7f7f7f', '#bcbd22', '#17becf']
|
||||
|
||||
|
||||
def compute_macro_roc(y_true, y_probs):
|
||||
one_hot = np.eye(NUM_CLASSES)[y_true]
|
||||
fpr_dict, tpr_dict = {}, {}
|
||||
for c in range(NUM_CLASSES):
|
||||
fpr_dict[c], tpr_dict[c], _ = roc_curve(one_hot[:, c], y_probs[:, c])
|
||||
all_fpr = np.unique(np.concatenate([fpr_dict[c] for c in range(NUM_CLASSES)]))
|
||||
mean_tpr = np.zeros_like(all_fpr)
|
||||
for c in range(NUM_CLASSES):
|
||||
mean_tpr += np.interp(all_fpr, fpr_dict[c], tpr_dict[c])
|
||||
mean_tpr /= NUM_CLASSES
|
||||
macro_auc = auc(all_fpr, mean_tpr)
|
||||
return all_fpr, mean_tpr, macro_auc
|
||||
|
||||
|
||||
def compute_macro_pr(y_true, y_probs):
|
||||
one_hot = np.eye(NUM_CLASSES)[y_true]
|
||||
prec_dict, rec_dict = {}, {}
|
||||
for c in range(NUM_CLASSES):
|
||||
prec_dict[c], rec_dict[c], _ = precision_recall_curve(one_hot[:, c], y_probs[:, c])
|
||||
all_rec = np.linspace(0, 1, 200)
|
||||
mean_prec = np.zeros_like(all_rec)
|
||||
for c in range(NUM_CLASSES):
|
||||
mean_prec += np.interp(all_rec, rec_dict[c][::-1], prec_dict[c][::-1])
|
||||
mean_prec /= NUM_CLASSES
|
||||
macro_ap = average_precision_score(one_hot, y_probs, average='macro')
|
||||
return all_rec, mean_prec, macro_ap
|
||||
|
||||
|
||||
def sanitize_filename(name):
|
||||
return re.sub(r'[^\w\-_]', '_', name).strip('_')
|
||||
|
||||
|
||||
def preds_csv_path(out_dir, model_name):
|
||||
safe = sanitize_filename(model_name)
|
||||
return os.path.join(out_dir, f'{safe}_preds.csv')
|
||||
|
||||
|
||||
def save_preds_csv(path, y_true, y_preds, y_probs):
|
||||
header = 'y_true,y_pred,' + ','.join(f'prob_{c}' for c in range(NUM_CLASSES))
|
||||
data = np.column_stack([y_true.astype(float), y_preds.astype(float), y_probs])
|
||||
np.savetxt(path, data, delimiter=',', header=header, comments='', fmt='%.6f')
|
||||
|
||||
|
||||
def load_preds_csv(path):
|
||||
data = np.loadtxt(path, delimiter=',', skiprows=1)
|
||||
y_true = data[:, 0].astype(int)
|
||||
y_preds = data[:, 1].astype(int)
|
||||
y_probs = data[:, 2:2 + NUM_CLASSES]
|
||||
return y_true, y_preds, y_probs
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
out_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
device = torch.device('cuda' if torch.cuda.is_available()
|
||||
else 'xpu' if hasattr(torch, 'xpu') and torch.xpu.is_available()
|
||||
else 'cpu')
|
||||
print(f"Device: {device}")
|
||||
|
||||
val_transform = transforms.Compose([
|
||||
transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=[0.485, 0.456, 0.406],
|
||||
std=[0.229, 0.224, 0.225]),
|
||||
])
|
||||
|
||||
train_dataset = RobustImageFolder(root=os.path.join(DATA_ROOT, 'train'),
|
||||
transform=val_transform)
|
||||
val_dataset = RobustImageFolder(root=os.path.join(DATA_ROOT, 'val'),
|
||||
transform=val_transform)
|
||||
print(f"训练集: {len(train_dataset)} 验证集: {len(val_dataset)}")
|
||||
|
||||
train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=False,
|
||||
num_workers=NUM_WORKERS, pin_memory=True, drop_last=False)
|
||||
val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False,
|
||||
num_workers=NUM_WORKERS, pin_memory=True, drop_last=False)
|
||||
|
||||
# ———— 评估所有模型(有缓存则跳过)————
|
||||
results = {}
|
||||
for name, func in MODELS:
|
||||
print(f"\n{'='*50}")
|
||||
csv_path = preds_csv_path(out_dir, name)
|
||||
if os.path.exists(csv_path):
|
||||
print(f"加载缓存: {os.path.basename(csv_path)}")
|
||||
y_true, y_preds, y_probs = load_preds_csv(csv_path)
|
||||
else:
|
||||
print(f"评估: {name}")
|
||||
y_true, y_preds, y_probs = func(train_loader, val_loader, device)
|
||||
save_preds_csv(csv_path, y_true, y_preds, y_probs)
|
||||
print(f" 预测数据已保存: {os.path.basename(csv_path)}")
|
||||
acc = accuracy_score(y_true, y_preds)
|
||||
fpr, tpr, roc_auc = compute_macro_roc(y_true, y_probs)
|
||||
rec, prec, macro_ap = compute_macro_pr(y_true, y_probs)
|
||||
results[name] = {'y_true': y_true, 'y_preds': y_preds, 'y_probs': y_probs,
|
||||
'acc': acc, 'fpr': fpr, 'tpr': tpr, 'auc': roc_auc,
|
||||
'rec': rec, 'prec': prec, 'ap': macro_ap}
|
||||
print(f" Accuracy: {acc:.4f} | Macro-AUC: {roc_auc:.4f} | Macro-AP: {macro_ap:.4f}")
|
||||
|
||||
# ———— ROC 对比图 ————
|
||||
fig, ax = plt.subplots(figsize=(8, 7))
|
||||
for i, (name, r) in enumerate(results.items()):
|
||||
color = COLORS[i % len(COLORS)]
|
||||
ax.plot(r['fpr'], r['tpr'], color=color, lw=2,
|
||||
label=f"{name} (AUC={r['auc']:.4f})")
|
||||
ax.plot([0, 1], [0, 1], 'k--', lw=1, alpha=0.5)
|
||||
ax.set_xlim(0, 1); ax.set_ylim(0, 1.05)
|
||||
ax.set_xlabel('False Positive Rate'); ax.set_ylabel('True Positive Rate')
|
||||
ax.set_title('ROC Curve Comparison (Macro-Average)', fontsize=14)
|
||||
ax.legend(loc='lower right'); ax.grid(True, alpha=0.3)
|
||||
plt.tight_layout()
|
||||
roc_path = os.path.join(out_dir, 'roc_comparison.png')
|
||||
plt.savefig(roc_path, dpi=150, bbox_inches='tight')
|
||||
plt.show()
|
||||
print(f"\nROC 对比图已保存: {roc_path}")
|
||||
|
||||
# ———— PR 对比图 ————
|
||||
fig, ax = plt.subplots(figsize=(8, 7))
|
||||
for i, (name, r) in enumerate(results.items()):
|
||||
color = COLORS[i % len(COLORS)]
|
||||
ax.plot(r['rec'], r['prec'], color=color, lw=2,
|
||||
label=f"{name} (AP={r['ap']:.4f})")
|
||||
ax.set_xlim(0, 1); ax.set_ylim(0, 1.05)
|
||||
ax.set_xlabel('Recall'); ax.set_ylabel('Precision')
|
||||
ax.set_title('PR Curve Comparison (Macro-Average)', fontsize=14)
|
||||
ax.legend(loc='lower left'); ax.grid(True, alpha=0.3)
|
||||
plt.tight_layout()
|
||||
pr_path = os.path.join(out_dir, 'pr_comparison.png')
|
||||
plt.savefig(pr_path, dpi=150, bbox_inches='tight')
|
||||
plt.show()
|
||||
print(f"PR 对比图已保存: {pr_path}")
|
||||
|
||||
# ———— 准确率柱状图 ————
|
||||
names = list(results.keys())
|
||||
accs = [results[n]['acc'] for n in names]
|
||||
fig, ax = plt.subplots(figsize=(8, 5))
|
||||
bar_colors = [COLORS[i % len(COLORS)] for i in range(len(names))]
|
||||
bars = ax.bar(names, accs, color=bar_colors, edgecolor='white', linewidth=1.2)
|
||||
for bar, acc in zip(bars, accs):
|
||||
ax.text(bar.get_x() + bar.get_width() / 2, bar.get_height() + 0.005,
|
||||
f'{acc:.4f}', ha='center', va='bottom', fontsize=12, fontweight='bold')
|
||||
ax.set_ylim(min(accs) - 0.03, max(accs) * 1.08)
|
||||
ax.set_ylabel('Accuracy'); ax.set_title('Accuracy Comparison', fontsize=14)
|
||||
ax.grid(True, alpha=0.3, axis='y')
|
||||
plt.tight_layout()
|
||||
bar_path = os.path.join(out_dir, 'accuracy_bar.png')
|
||||
plt.savefig(bar_path, dpi=150, bbox_inches='tight')
|
||||
plt.show()
|
||||
print(f"准确率柱状图已保存: {bar_path}")
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 122 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 125 KiB |
Reference in New Issue
Block a user