Initial commit
This commit is contained in:
commit
80fa1d394d
77 files changed
+5568
No files matched your search
@@ -0,0 +1,7 @@
|
||||
{
|
||||
"0": "daisy",
|
||||
"1": "dandelion",
|
||||
"2": "roses",
|
||||
"3": "sunflowers",
|
||||
"4": "tulips"
|
||||
}
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 192 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 95 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 131 KiB |
@@ -0,0 +1,343 @@
|
||||
import json
|
||||
import os
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import seaborn as sns
|
||||
import torch
|
||||
from model import densenet121
|
||||
from my_dataset import MyDataSet
|
||||
from sklearn.metrics import (
|
||||
accuracy_score,
|
||||
confusion_matrix,
|
||||
f1_score,
|
||||
precision_score,
|
||||
recall_score,
|
||||
)
|
||||
from torchvision import transforms
|
||||
from tqdm import tqdm
|
||||
from utils import read_split_data
|
||||
|
||||
|
||||
def main():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
print(f"Using {device} device for evaluation")
|
||||
|
||||
# Data transformation for validation
|
||||
data_transform = transforms.Compose(
|
||||
[
|
||||
transforms.Resize(256),
|
||||
transforms.CenterCrop(224),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
||||
]
|
||||
)
|
||||
|
||||
# Load validation dataset
|
||||
_, _, val_images_path, val_images_label = read_split_data("../data/flower_photos")
|
||||
|
||||
# Create validation dataset
|
||||
val_dataset = MyDataSet(
|
||||
images_path=val_images_path,
|
||||
images_class=val_images_label,
|
||||
transform=data_transform,
|
||||
)
|
||||
|
||||
batch_size = 4
|
||||
nw = min([os.cpu_count(), batch_size if batch_size > 1 else 0, 8])
|
||||
validate_loader = torch.utils.data.DataLoader(
|
||||
val_dataset,
|
||||
batch_size=batch_size,
|
||||
shuffle=False,
|
||||
num_workers=nw,
|
||||
collate_fn=val_dataset.collate_fn,
|
||||
)
|
||||
|
||||
# Load class indices
|
||||
json_path = "./class_indices.json"
|
||||
assert os.path.exists(json_path), f"File: '{json_path}' does not exist."
|
||||
with open(json_path, "r") as f:
|
||||
class_indict = json.load(f)
|
||||
|
||||
# Create model
|
||||
model = densenet121(num_classes=5).to(device)
|
||||
|
||||
# Load model weights
|
||||
weights_path = "./weights/best_model.pth"
|
||||
assert os.path.exists(weights_path), f"File: '{weights_path}' does not exist."
|
||||
model.load_state_dict(torch.load(weights_path, map_location=device))
|
||||
|
||||
# Evaluation
|
||||
model.eval()
|
||||
|
||||
# Lists to store predictions and ground truth
|
||||
all_preds = []
|
||||
all_labels = []
|
||||
|
||||
# Store per-image predictions and confidences for error analysis
|
||||
image_paths = []
|
||||
image_preds = []
|
||||
image_confidences = []
|
||||
image_true_labels = []
|
||||
|
||||
with torch.no_grad():
|
||||
for i, (images, labels) in enumerate(tqdm(validate_loader, desc="Evaluating")):
|
||||
outputs = model(images.to(device))
|
||||
|
||||
# Get predictions and confidences
|
||||
probabilities = torch.softmax(outputs, dim=1)
|
||||
confidence, predict_y = torch.max(probabilities, dim=1)
|
||||
|
||||
# Store predictions and labels
|
||||
all_preds.extend(predict_y.cpu().numpy())
|
||||
all_labels.extend(labels.numpy())
|
||||
|
||||
# Store image details for error analysis
|
||||
for j in range(len(images)):
|
||||
if i * batch_size + j < len(val_images_path):
|
||||
image_paths.append(val_images_path[i * batch_size + j])
|
||||
image_preds.append(predict_y[j].item())
|
||||
image_confidences.append(confidence[j].item())
|
||||
image_true_labels.append(labels[j].item())
|
||||
|
||||
# Convert to numpy arrays
|
||||
all_preds = np.array(all_preds)
|
||||
all_labels = np.array(all_labels)
|
||||
|
||||
# Calculate metrics
|
||||
accuracy = accuracy_score(all_labels, all_preds)
|
||||
precision = precision_score(all_labels, all_preds, average="macro")
|
||||
recall = recall_score(all_labels, all_preds, average="macro")
|
||||
f1 = f1_score(all_labels, all_preds, average="macro")
|
||||
|
||||
# Calculate per-class metrics
|
||||
class_precision = precision_score(all_labels, all_preds, average=None)
|
||||
class_recall = recall_score(all_labels, all_preds, average=None)
|
||||
class_f1 = f1_score(all_labels, all_preds, average=None)
|
||||
|
||||
# Print overall metrics
|
||||
print("\nOverall Metrics:")
|
||||
print(f"Accuracy: {accuracy:.4f}")
|
||||
print(f"Precision: {precision:.4f}")
|
||||
print(f"Recall: {recall:.4f}")
|
||||
print(f"F1 Score: {f1:.4f}")
|
||||
|
||||
# Print per-class metrics
|
||||
print("\nPer-class Metrics:")
|
||||
for i in range(len(class_indict)):
|
||||
class_name = class_indict[str(i)]
|
||||
print(f"Class: {class_name}")
|
||||
print(f" Precision: {class_precision[i]:.4f}")
|
||||
print(f" Recall: {class_recall[i]:.4f}")
|
||||
print(f" F1 Score: {class_f1[i]:.4f}")
|
||||
|
||||
# Create confusion matrix
|
||||
cm = confusion_matrix(all_labels, all_preds)
|
||||
|
||||
# Analyze misclassified images
|
||||
misclassified = []
|
||||
for path, pred, conf, true_label in zip(
|
||||
image_paths, image_preds, image_confidences, image_true_labels
|
||||
):
|
||||
if pred != true_label:
|
||||
misclassified.append(
|
||||
{
|
||||
"path": path,
|
||||
"predicted": class_indict[str(pred)],
|
||||
"true": class_indict[str(true_label)],
|
||||
"confidence": conf,
|
||||
}
|
||||
)
|
||||
|
||||
# Sort misclassified by confidence (high to low)
|
||||
misclassified.sort(key=lambda x: x["confidence"], reverse=True)
|
||||
|
||||
# Print most confident misclassifications
|
||||
print("\nTop 5 Most Confident Misclassifications:")
|
||||
for i, item in enumerate(misclassified[:5]):
|
||||
print(f"{i+1}. Path: {os.path.basename(item['path'])}")
|
||||
print(
|
||||
f" True: {item['true']}, Predicted: {item['predicted']}, Confidence: {item['confidence']:.4f}"
|
||||
)
|
||||
|
||||
# Visualize results
|
||||
visualize_results(
|
||||
cm,
|
||||
class_indict,
|
||||
accuracy,
|
||||
precision,
|
||||
recall,
|
||||
f1,
|
||||
class_precision,
|
||||
class_recall,
|
||||
class_f1,
|
||||
misclassified[:10], # Pass top 10 misclassified for visualization
|
||||
)
|
||||
|
||||
|
||||
def visualize_results(
|
||||
cm,
|
||||
class_indict,
|
||||
accuracy,
|
||||
precision,
|
||||
recall,
|
||||
f1,
|
||||
class_precision,
|
||||
class_recall,
|
||||
class_f1,
|
||||
misclassified=None,
|
||||
):
|
||||
"""
|
||||
Visualize evaluation metrics and confusion matrix
|
||||
"""
|
||||
plt.figure(figsize=(20, 15))
|
||||
|
||||
# Plot confusion matrix with percentages
|
||||
plt.subplot(2, 2, 1)
|
||||
cm_norm = cm.astype("float") / cm.sum(axis=1)[:, np.newaxis]
|
||||
sns.heatmap(
|
||||
cm_norm,
|
||||
annot=True,
|
||||
fmt=".2f",
|
||||
cmap="Blues",
|
||||
xticklabels=[class_indict[str(i)] for i in range(len(class_indict))],
|
||||
yticklabels=[class_indict[str(i)] for i in range(len(class_indict))],
|
||||
)
|
||||
plt.xlabel("Predicted")
|
||||
plt.ylabel("True")
|
||||
plt.title("Normalized Confusion Matrix")
|
||||
|
||||
# Plot absolute confusion matrix
|
||||
plt.subplot(2, 2, 2)
|
||||
sns.heatmap(
|
||||
cm,
|
||||
annot=True,
|
||||
fmt="d",
|
||||
cmap="Blues",
|
||||
xticklabels=[class_indict[str(i)] for i in range(len(class_indict))],
|
||||
yticklabels=[class_indict[str(i)] for i in range(len(class_indict))],
|
||||
)
|
||||
plt.xlabel("Predicted")
|
||||
plt.ylabel("True")
|
||||
plt.title("Confusion Matrix (Counts)")
|
||||
|
||||
# Plot overall metrics
|
||||
plt.subplot(2, 2, 3)
|
||||
metrics = ["Accuracy", "Precision", "Recall", "F1 Score"]
|
||||
values = [accuracy, precision, recall, f1]
|
||||
colors = ["#4CAF50", "#2196F3", "#FFC107", "#F44336"] # Green, Blue, Yellow, Red
|
||||
|
||||
bars = plt.bar(metrics, values, color=colors)
|
||||
plt.ylim(0, 1.0)
|
||||
plt.title("Overall Performance Metrics")
|
||||
|
||||
# Add value labels above bars
|
||||
for bar in bars:
|
||||
height = bar.get_height()
|
||||
plt.text(
|
||||
bar.get_x() + bar.get_width() / 2.0,
|
||||
height + 0.02,
|
||||
f"{height:.4f}",
|
||||
ha="center",
|
||||
va="bottom",
|
||||
)
|
||||
|
||||
# Plot per-class metrics
|
||||
plt.subplot(2, 2, 4)
|
||||
x = np.arange(len(class_indict))
|
||||
width = 0.25
|
||||
|
||||
fig, ax = plt.gcf(), plt.gca()
|
||||
rects1 = ax.bar(
|
||||
x - width, class_precision, width, label="Precision", color="#2196F3"
|
||||
)
|
||||
rects2 = ax.bar(x, class_recall, width, label="Recall", color="#FFC107")
|
||||
rects3 = ax.bar(x + width, class_f1, width, label="F1 Score", color="#F44336")
|
||||
|
||||
ax.set_xlabel("Class")
|
||||
ax.set_ylabel("Score")
|
||||
ax.set_title("Per-class Performance Metrics")
|
||||
ax.set_xticks(x)
|
||||
ax.set_xticklabels([class_indict[str(i)] for i in range(len(class_indict))])
|
||||
ax.legend()
|
||||
ax.set_ylim(0, 1.0)
|
||||
|
||||
# Save main evaluation visualization
|
||||
plt.tight_layout()
|
||||
plt.savefig("densenet_evaluation_metrics.png")
|
||||
|
||||
# If we have misclassified examples, create an additional visualization
|
||||
if misclassified:
|
||||
# Create a radar chart for performance metrics
|
||||
plt.figure(figsize=(15, 10))
|
||||
|
||||
# Create a radar chart for per-class F1 scores
|
||||
plt.subplot(1, 2, 1, polar=True)
|
||||
categories = [class_indict[str(i)] for i in range(len(class_indict))]
|
||||
N = len(categories)
|
||||
|
||||
# Create angles for each metric
|
||||
angles = [n / float(N) * 2 * np.pi for n in range(N)]
|
||||
angles += angles[:1] # Close the loop
|
||||
|
||||
# Add metrics
|
||||
metrics_data = [
|
||||
class_precision.tolist() + [class_precision[0]], # Close the loop
|
||||
class_recall.tolist() + [class_recall[0]], # Close the loop
|
||||
class_f1.tolist() + [class_f1[0]], # Close the loop
|
||||
]
|
||||
|
||||
labels = ["Precision", "Recall", "F1 Score"]
|
||||
colors = ["#2196F3", "#FFC107", "#F44336"]
|
||||
|
||||
ax = plt.subplot(1, 2, 1, polar=True)
|
||||
|
||||
for i, data in enumerate(metrics_data):
|
||||
ax.plot(angles, data, linewidth=2, label=labels[i], color=colors[i])
|
||||
ax.fill(angles, data, alpha=0.1, color=colors[i])
|
||||
|
||||
plt.xticks(angles[:-1], categories)
|
||||
ax.set_rlabel_position(0)
|
||||
plt.yticks(
|
||||
[0.2, 0.4, 0.6, 0.8, 1.0], ["0.2", "0.4", "0.6", "0.8", "1.0"], color="gray"
|
||||
)
|
||||
plt.ylim(0, 1)
|
||||
plt.title("Class Performance Metrics")
|
||||
plt.legend(loc="upper right", bbox_to_anchor=(0.1, 0.1))
|
||||
|
||||
# Create error analysis pie chart
|
||||
plt.subplot(1, 2, 2)
|
||||
|
||||
# Count misclassifications by class
|
||||
true_class_counts = {}
|
||||
for item in misclassified:
|
||||
true_class = item["true"]
|
||||
if true_class in true_class_counts:
|
||||
true_class_counts[true_class] += 1
|
||||
else:
|
||||
true_class_counts[true_class] = 1
|
||||
|
||||
# Create pie chart of misclassified true classes
|
||||
labels = list(true_class_counts.keys())
|
||||
sizes = list(true_class_counts.values())
|
||||
|
||||
plt.pie(
|
||||
sizes,
|
||||
labels=labels,
|
||||
autopct="%1.1f%%",
|
||||
startangle=90,
|
||||
shadow=True,
|
||||
explode=[0.05] * len(labels),
|
||||
colors=plt.cm.tab10.colors[: len(labels)],
|
||||
)
|
||||
plt.axis("equal")
|
||||
plt.title("Distribution of Misclassified Images by True Class")
|
||||
|
||||
plt.tight_layout()
|
||||
plt.savefig("densenet_error_analysis.png")
|
||||
|
||||
plt.show()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,286 @@
|
||||
import re
|
||||
from collections import OrderedDict
|
||||
from typing import Any, List, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.utils.checkpoint as cp
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class _DenseLayer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_c: int,
|
||||
growth_rate: int,
|
||||
bn_size: int,
|
||||
drop_rate: float,
|
||||
memory_efficient: bool = False,
|
||||
):
|
||||
super(_DenseLayer, self).__init__()
|
||||
|
||||
self.add_module("norm1", nn.BatchNorm2d(input_c))
|
||||
self.add_module("relu1", nn.ReLU(inplace=True))
|
||||
self.add_module(
|
||||
"conv1",
|
||||
nn.Conv2d(
|
||||
in_channels=input_c,
|
||||
out_channels=bn_size * growth_rate,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
bias=False,
|
||||
),
|
||||
)
|
||||
self.add_module("norm2", nn.BatchNorm2d(bn_size * growth_rate))
|
||||
self.add_module("relu2", nn.ReLU(inplace=True))
|
||||
self.add_module(
|
||||
"conv2",
|
||||
nn.Conv2d(
|
||||
bn_size * growth_rate,
|
||||
growth_rate,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1,
|
||||
bias=False,
|
||||
),
|
||||
)
|
||||
self.drop_rate = drop_rate
|
||||
self.memory_efficient = memory_efficient
|
||||
|
||||
def bn_function(self, inputs: List[Tensor]) -> Tensor:
|
||||
concat_features = torch.cat(inputs, 1)
|
||||
bottleneck_output = self.conv1(self.relu1(self.norm1(concat_features)))
|
||||
return bottleneck_output
|
||||
|
||||
@staticmethod
|
||||
def any_requires_grad(inputs: List[Tensor]) -> bool:
|
||||
for tensor in inputs:
|
||||
if tensor.requires_grad:
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
@torch.jit.unused
|
||||
def call_checkpoint_bottleneck(self, inputs: List[Tensor]) -> Tensor:
|
||||
def closure(*inp):
|
||||
return self.bn_function(inp)
|
||||
|
||||
return cp.checkpoint(closure, *inputs)
|
||||
|
||||
def forward(self, inputs: Tensor) -> Tensor:
|
||||
if isinstance(inputs, Tensor):
|
||||
prev_features = [inputs]
|
||||
else:
|
||||
prev_features = inputs
|
||||
|
||||
if self.memory_efficient and self.any_requires_grad(prev_features):
|
||||
if torch.jit.is_scripting():
|
||||
raise Exception("memory efficient not supported in JIT")
|
||||
|
||||
bottleneck_output = self.call_checkpoint_bottleneck(prev_features)
|
||||
else:
|
||||
bottleneck_output = self.bn_function(prev_features)
|
||||
|
||||
new_features = self.conv2(self.relu2(self.norm2(bottleneck_output)))
|
||||
if self.drop_rate > 0:
|
||||
new_features = F.dropout(
|
||||
new_features, p=self.drop_rate, training=self.training
|
||||
)
|
||||
|
||||
return new_features
|
||||
|
||||
|
||||
class _DenseBlock(nn.ModuleDict):
|
||||
_version = 2
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_layers: int,
|
||||
input_c: int,
|
||||
bn_size: int,
|
||||
growth_rate: int,
|
||||
drop_rate: float,
|
||||
memory_efficient: bool = False,
|
||||
):
|
||||
super(_DenseBlock, self).__init__()
|
||||
for i in range(num_layers):
|
||||
layer = _DenseLayer(
|
||||
input_c + i * growth_rate,
|
||||
growth_rate=growth_rate,
|
||||
bn_size=bn_size,
|
||||
drop_rate=drop_rate,
|
||||
memory_efficient=memory_efficient,
|
||||
)
|
||||
self.add_module("denselayer%d" % (i + 1), layer)
|
||||
|
||||
def forward(self, init_features: Tensor) -> Tensor:
|
||||
features = [init_features]
|
||||
for name, layer in self.items():
|
||||
new_features = layer(features)
|
||||
features.append(new_features)
|
||||
return torch.cat(features, 1)
|
||||
|
||||
|
||||
class _Transition(nn.Sequential):
|
||||
def __init__(self, input_c: int, output_c: int):
|
||||
super(_Transition, self).__init__()
|
||||
self.add_module("norm", nn.BatchNorm2d(input_c))
|
||||
self.add_module("relu", nn.ReLU(inplace=True))
|
||||
self.add_module(
|
||||
"conv", nn.Conv2d(input_c, output_c, kernel_size=1, stride=1, bias=False)
|
||||
)
|
||||
self.add_module("pool", nn.AvgPool2d(kernel_size=2, stride=2))
|
||||
|
||||
|
||||
class DenseNet(nn.Module):
|
||||
"""
|
||||
Densenet-BC model class for imagenet
|
||||
|
||||
Args:
|
||||
growth_rate (int) - how many filters to add each layer (`k` in paper)
|
||||
block_config (list of 4 ints) - how many layers in each pooling block
|
||||
num_init_features (int) - the number of filters to learn in the first convolution layer
|
||||
bn_size (int) - multiplicative factor for number of bottle neck layers
|
||||
(i.e. bn_size * k features in the bottleneck layer)
|
||||
drop_rate (float) - dropout rate after each dense layer
|
||||
num_classes (int) - number of classification classes
|
||||
memory_efficient (bool) - If True, uses checkpointing. Much more memory efficient
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
growth_rate: int = 32,
|
||||
block_config: Tuple[int, int, int, int] = (6, 12, 24, 16),
|
||||
num_init_features: int = 64,
|
||||
bn_size: int = 4,
|
||||
drop_rate: float = 0,
|
||||
num_classes: int = 1000,
|
||||
memory_efficient: bool = False,
|
||||
):
|
||||
super(DenseNet, self).__init__()
|
||||
|
||||
# first conv+bn+relu+pool
|
||||
self.features = nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
(
|
||||
"conv0",
|
||||
nn.Conv2d(
|
||||
3,
|
||||
num_init_features,
|
||||
kernel_size=7,
|
||||
stride=2,
|
||||
padding=3,
|
||||
bias=False,
|
||||
),
|
||||
),
|
||||
("norm0", nn.BatchNorm2d(num_init_features)),
|
||||
("relu0", nn.ReLU(inplace=True)),
|
||||
("pool0", nn.MaxPool2d(kernel_size=3, stride=2, padding=1)),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
# each dense block
|
||||
num_features = num_init_features
|
||||
for i, num_layers in enumerate(block_config):
|
||||
block = _DenseBlock(
|
||||
num_layers=num_layers,
|
||||
input_c=num_features,
|
||||
bn_size=bn_size,
|
||||
growth_rate=growth_rate,
|
||||
drop_rate=drop_rate,
|
||||
memory_efficient=memory_efficient,
|
||||
)
|
||||
self.features.add_module("denseblock%d" % (i + 1), block)
|
||||
num_features = num_features + num_layers * growth_rate
|
||||
|
||||
if i != len(block_config) - 1:
|
||||
trans = _Transition(input_c=num_features, output_c=num_features // 2)
|
||||
self.features.add_module("transition%d" % (i + 1), trans)
|
||||
num_features = num_features // 2
|
||||
|
||||
# finnal batch norm
|
||||
self.features.add_module("norm5", nn.BatchNorm2d(num_features))
|
||||
|
||||
# fc layer
|
||||
self.classifier = nn.Linear(num_features, num_classes)
|
||||
|
||||
# init weights
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
nn.init.kaiming_normal_(m.weight)
|
||||
elif isinstance(m, nn.BatchNorm2d):
|
||||
nn.init.constant_(m.weight, 1)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.Linear):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
features = self.features(x)
|
||||
out = F.relu(features, inplace=True)
|
||||
out = F.adaptive_avg_pool2d(out, (1, 1))
|
||||
out = torch.flatten(out, 1)
|
||||
out = self.classifier(out)
|
||||
return out
|
||||
|
||||
|
||||
def densenet121(**kwargs: Any) -> DenseNet:
|
||||
# Top-1 error: 25.35%
|
||||
# 'densenet121': 'https://download.pytorch.org/models/densenet121-a639ec97.pth'
|
||||
return DenseNet(
|
||||
growth_rate=32, block_config=(6, 12, 24, 16), num_init_features=64, **kwargs
|
||||
)
|
||||
|
||||
|
||||
def densenet169(**kwargs: Any) -> DenseNet:
|
||||
# Top-1 error: 24.00%
|
||||
# 'densenet169': 'https://download.pytorch.org/models/densenet169-b2777c0a.pth'
|
||||
return DenseNet(
|
||||
growth_rate=32, block_config=(6, 12, 32, 32), num_init_features=64, **kwargs
|
||||
)
|
||||
|
||||
|
||||
def densenet201(**kwargs: Any) -> DenseNet:
|
||||
# Top-1 error: 22.80%
|
||||
# 'densenet201': 'https://download.pytorch.org/models/densenet201-c1103571.pth'
|
||||
return DenseNet(
|
||||
growth_rate=32, block_config=(6, 12, 48, 32), num_init_features=64, **kwargs
|
||||
)
|
||||
|
||||
|
||||
def densenet161(**kwargs: Any) -> DenseNet:
|
||||
# Top-1 error: 22.35%
|
||||
# 'densenet161': 'https://download.pytorch.org/models/densenet161-8d451a50.pth'
|
||||
return DenseNet(
|
||||
growth_rate=48, block_config=(6, 12, 36, 24), num_init_features=96, **kwargs
|
||||
)
|
||||
|
||||
|
||||
def load_state_dict(model: nn.Module, weights_path: str) -> None:
|
||||
# '.'s are no longer allowed in module names, but previous _DenseLayer
|
||||
# has keys 'norm.1', 'relu.1', 'conv.1', 'norm.2', 'relu.2', 'conv.2'.
|
||||
# They are also in the checkpoints in model_urls. This pattern is used
|
||||
# to find such keys.
|
||||
pattern = re.compile(
|
||||
r"^(.*denselayer\d+\.(?:norm|relu|conv))\.((?:[12])\.(?:weight|bias|running_mean|running_var))$"
|
||||
)
|
||||
|
||||
state_dict = torch.load(weights_path)
|
||||
|
||||
num_classes = model.classifier.out_features
|
||||
load_fc = num_classes == 1000
|
||||
|
||||
for key in list(state_dict.keys()):
|
||||
if load_fc is False:
|
||||
if "classifier" in key:
|
||||
del state_dict[key]
|
||||
|
||||
res = pattern.match(key)
|
||||
if res:
|
||||
new_key = res.group(1) + res.group(2)
|
||||
state_dict[new_key] = state_dict[key]
|
||||
del state_dict[key]
|
||||
model.load_state_dict(state_dict, strict=load_fc)
|
||||
print("successfully load pretrain-weights.")
|
||||
@@ -0,0 +1,37 @@
|
||||
import torch
|
||||
from PIL import Image
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
|
||||
class MyDataSet(Dataset):
|
||||
"""自定义数据集"""
|
||||
|
||||
def __init__(self, images_path: list, images_class: list, transform=None):
|
||||
self.images_path = images_path
|
||||
self.images_class = images_class
|
||||
self.transform = transform
|
||||
|
||||
def __len__(self):
|
||||
return len(self.images_path)
|
||||
|
||||
def __getitem__(self, item):
|
||||
img = Image.open(self.images_path[item])
|
||||
# RGB为彩色图片,L为灰度图片
|
||||
if img.mode != "RGB":
|
||||
raise ValueError("image: {} isn't RGB mode.".format(self.images_path[item]))
|
||||
label = self.images_class[item]
|
||||
|
||||
if self.transform is not None:
|
||||
img = self.transform(img)
|
||||
|
||||
return img, label
|
||||
|
||||
@staticmethod
|
||||
def collate_fn(batch):
|
||||
# 官方实现的default_collate可以参考
|
||||
# https://github.com/pytorch/pytorch/blob/67b7e751e6b5931a9f45274653f4f653a4e6cdf6/torch/utils/data/_utils/collate.py
|
||||
images, labels = tuple(zip(*batch))
|
||||
|
||||
images = torch.stack(images, dim=0)
|
||||
labels = torch.as_tensor(labels)
|
||||
return images, labels
|
||||
@@ -0,0 +1,66 @@
|
||||
import json
|
||||
import os
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import torch
|
||||
from model import densenet121
|
||||
from PIL import Image
|
||||
from torchvision import transforms
|
||||
|
||||
|
||||
def main():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
data_transform = transforms.Compose(
|
||||
[
|
||||
transforms.Resize(256),
|
||||
transforms.CenterCrop(224),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
||||
]
|
||||
)
|
||||
|
||||
# load image
|
||||
img_path = "5794839_200acd910c_n.jpg"
|
||||
assert os.path.exists(img_path), "file: '{}' dose not exist.".format(img_path)
|
||||
img = Image.open(img_path)
|
||||
plt.imshow(img)
|
||||
# [N, C, H, W]
|
||||
img = data_transform(img)
|
||||
# expand batch dimension
|
||||
img = torch.unsqueeze(img, dim=0)
|
||||
|
||||
# read class_indict
|
||||
json_path = "./class_indices.json"
|
||||
assert os.path.exists(json_path), "file: '{}' dose not exist.".format(json_path)
|
||||
|
||||
json_file = open(json_path, "r")
|
||||
class_indict = json.load(json_file)
|
||||
|
||||
# create model
|
||||
model = densenet121(num_classes=5).to(device)
|
||||
# load model weights
|
||||
model_weight_path = "weights/best_model.pth"
|
||||
model.load_state_dict(torch.load(model_weight_path, map_location=device))
|
||||
model.eval()
|
||||
with torch.no_grad():
|
||||
# predict class
|
||||
output = torch.squeeze(model(img.to(device))).cpu()
|
||||
predict = torch.softmax(output, dim=0)
|
||||
predict_cla = torch.argmax(predict).numpy()
|
||||
|
||||
print_res = "class: {} prob: {:.3}".format(
|
||||
class_indict[str(predict_cla)], predict[predict_cla].numpy()
|
||||
)
|
||||
plt.title(print_res)
|
||||
for i in range(len(predict)):
|
||||
print(
|
||||
"class: {:10} prob: {:.3}".format(
|
||||
class_indict[str(i)], predict[i].numpy()
|
||||
)
|
||||
)
|
||||
plt.show()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,286 @@
|
||||
import argparse
|
||||
import math
|
||||
import os
|
||||
import urllib.request
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.optim as optim
|
||||
import torch.optim.lr_scheduler as lr_scheduler
|
||||
from model import densenet121, load_state_dict
|
||||
from my_dataset import MyDataSet
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
from torchvision import transforms
|
||||
from utils import evaluate, read_split_data, train_one_epoch
|
||||
|
||||
|
||||
def download_weights(url, filename):
|
||||
"""
|
||||
Download weights file if it doesn't exist locally
|
||||
"""
|
||||
if not os.path.exists(filename):
|
||||
print(f"Downloading weights file from {url}...")
|
||||
try:
|
||||
# Create directory if it doesn't exist
|
||||
os.makedirs(
|
||||
os.path.dirname(filename) if os.path.dirname(filename) else ".",
|
||||
exist_ok=True,
|
||||
)
|
||||
# Download the file
|
||||
urllib.request.urlretrieve(url, filename)
|
||||
print(f"Downloaded weights file to {filename}")
|
||||
return True
|
||||
except Exception as e:
|
||||
print(f"Error downloading weights file: {e}")
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def main(args):
|
||||
device = torch.device(args.device if torch.cuda.is_available() else "cpu")
|
||||
|
||||
print(args)
|
||||
print(
|
||||
'Start Tensorboard with "tensorboard --logdir=runs", view at http://localhost:6006/'
|
||||
)
|
||||
tb_writer = SummaryWriter()
|
||||
if os.path.exists("./weights") is False:
|
||||
os.makedirs("./weights")
|
||||
|
||||
train_images_path, train_images_label, val_images_path, val_images_label = (
|
||||
read_split_data(args.data_path)
|
||||
)
|
||||
|
||||
data_transform = {
|
||||
"train": transforms.Compose(
|
||||
[
|
||||
transforms.RandomResizedCrop(224),
|
||||
transforms.RandomHorizontalFlip(),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
||||
]
|
||||
),
|
||||
"val": transforms.Compose(
|
||||
[
|
||||
transforms.Resize(256),
|
||||
transforms.CenterCrop(224),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
||||
]
|
||||
),
|
||||
}
|
||||
|
||||
# 实例化训练数据集
|
||||
train_dataset = MyDataSet(
|
||||
images_path=train_images_path,
|
||||
images_class=train_images_label,
|
||||
transform=data_transform["train"],
|
||||
)
|
||||
|
||||
# 实例化验证数据集
|
||||
val_dataset = MyDataSet(
|
||||
images_path=val_images_path,
|
||||
images_class=val_images_label,
|
||||
transform=data_transform["val"],
|
||||
)
|
||||
|
||||
batch_size = args.batch_size
|
||||
nw = min(
|
||||
[os.cpu_count(), batch_size if batch_size > 1 else 0, 8]
|
||||
) # number of workers
|
||||
print("Using {} dataloader workers every process".format(nw))
|
||||
train_loader = torch.utils.data.DataLoader(
|
||||
train_dataset,
|
||||
batch_size=batch_size,
|
||||
shuffle=True,
|
||||
pin_memory=True,
|
||||
num_workers=nw,
|
||||
collate_fn=train_dataset.collate_fn,
|
||||
)
|
||||
|
||||
val_loader = torch.utils.data.DataLoader(
|
||||
val_dataset,
|
||||
batch_size=batch_size,
|
||||
shuffle=False,
|
||||
pin_memory=True,
|
||||
num_workers=nw,
|
||||
collate_fn=val_dataset.collate_fn,
|
||||
)
|
||||
|
||||
# 如果存在预训练权重则载入
|
||||
model = densenet121(num_classes=args.num_classes).to(device)
|
||||
if args.weights != "":
|
||||
if os.path.exists(args.weights):
|
||||
load_state_dict(model, args.weights)
|
||||
else:
|
||||
# Try to download the weights file
|
||||
weights_url = "https://download.pytorch.org/models/densenet121-a639ec97.pth"
|
||||
if download_weights(weights_url, args.weights):
|
||||
load_state_dict(model, args.weights)
|
||||
else:
|
||||
print(
|
||||
f"Warning: Could not download weights file. Training from scratch."
|
||||
)
|
||||
|
||||
# 是否冻结权重
|
||||
if args.freeze_layers:
|
||||
for name, para in model.named_parameters():
|
||||
# 除最后的全连接层外,其他权重全部冻结
|
||||
if "classifier" not in name:
|
||||
para.requires_grad_(False)
|
||||
|
||||
pg = [p for p in model.parameters() if p.requires_grad]
|
||||
optimizer = optim.SGD(
|
||||
pg, lr=args.lr, momentum=0.9, weight_decay=1e-4, nesterov=True
|
||||
)
|
||||
# Scheduler https://arxiv.org/pdf/1812.01187.pdf
|
||||
lf = (
|
||||
lambda x: ((1 + math.cos(x * math.pi / args.epochs)) / 2) * (1 - args.lrf)
|
||||
+ args.lrf
|
||||
) # cosine
|
||||
scheduler = lr_scheduler.LambdaLR(optimizer, lr_lambda=lf)
|
||||
|
||||
# Lists to store metrics for plotting
|
||||
train_losses = []
|
||||
val_accuracies = []
|
||||
learning_rates = []
|
||||
|
||||
best_acc = 0.0
|
||||
for epoch in range(args.epochs):
|
||||
# train
|
||||
mean_loss = train_one_epoch(
|
||||
model=model,
|
||||
optimizer=optimizer,
|
||||
data_loader=train_loader,
|
||||
device=device,
|
||||
epoch=epoch,
|
||||
)
|
||||
|
||||
scheduler.step()
|
||||
current_lr = optimizer.param_groups[0]["lr"]
|
||||
learning_rates.append(current_lr)
|
||||
|
||||
# validate
|
||||
acc = evaluate(model=model, data_loader=val_loader, device=device)
|
||||
|
||||
# Store metrics for plotting
|
||||
train_losses.append(mean_loss)
|
||||
val_accuracies.append(acc)
|
||||
|
||||
print(
|
||||
"[epoch %d] train_loss: %.3f val_accuracy: %.3f lr: %.6f"
|
||||
% (epoch + 1, mean_loss, acc, current_lr)
|
||||
)
|
||||
|
||||
tags = ["loss", "accuracy", "learning_rate"]
|
||||
tb_writer.add_scalar(tags[0], mean_loss, epoch)
|
||||
tb_writer.add_scalar(tags[1], acc, epoch)
|
||||
tb_writer.add_scalar(tags[2], current_lr, epoch)
|
||||
|
||||
torch.save(model.state_dict(), "./weights/model-{}.pth".format(epoch))
|
||||
if acc > best_acc:
|
||||
best_acc = acc
|
||||
# Save the model with the best validation accuracy
|
||||
torch.save(model.state_dict(), "./weights/best_model.pth")
|
||||
|
||||
# When training is complete, visualize the training process
|
||||
visualize_training(args.epochs, train_losses, val_accuracies, learning_rates)
|
||||
|
||||
|
||||
def visualize_training(epochs, train_losses, val_accuracies, learning_rates):
|
||||
"""
|
||||
Create and save visualizations of the training process
|
||||
"""
|
||||
plt.figure(figsize=(15, 10))
|
||||
|
||||
# Plot training loss
|
||||
plt.subplot(2, 2, 1)
|
||||
plt.plot(
|
||||
range(1, epochs + 1), train_losses, "b-", marker="o", label="Training Loss"
|
||||
)
|
||||
plt.title("Training Loss vs. Epochs")
|
||||
plt.xlabel("Epochs")
|
||||
plt.ylabel("Loss")
|
||||
plt.grid(True)
|
||||
plt.legend()
|
||||
|
||||
# Plot validation accuracy
|
||||
plt.subplot(2, 2, 2)
|
||||
plt.plot(
|
||||
range(1, epochs + 1),
|
||||
val_accuracies,
|
||||
"r-",
|
||||
marker="o",
|
||||
label="Validation Accuracy",
|
||||
)
|
||||
plt.title("Validation Accuracy vs. Epochs")
|
||||
plt.xlabel("Epochs")
|
||||
plt.ylabel("Accuracy")
|
||||
plt.grid(True)
|
||||
plt.legend()
|
||||
|
||||
# Plot learning rate
|
||||
plt.subplot(2, 2, 3)
|
||||
plt.plot(
|
||||
range(1, epochs + 1), learning_rates, "m-", marker="o", label="Learning Rate"
|
||||
)
|
||||
plt.title("Learning Rate vs. Epochs")
|
||||
plt.xlabel("Epochs")
|
||||
plt.ylabel("Learning Rate")
|
||||
plt.grid(True)
|
||||
plt.legend()
|
||||
|
||||
# Plot loss vs. accuracy
|
||||
plt.subplot(2, 2, 4)
|
||||
plt.scatter(
|
||||
train_losses,
|
||||
val_accuracies,
|
||||
c=range(epochs),
|
||||
cmap="viridis",
|
||||
s=50,
|
||||
alpha=0.7,
|
||||
edgecolors="k",
|
||||
linewidths=0.5,
|
||||
)
|
||||
plt.colorbar(label="Epoch")
|
||||
plt.title("Validation Accuracy vs. Training Loss")
|
||||
plt.xlabel("Training Loss")
|
||||
plt.ylabel("Validation Accuracy")
|
||||
plt.grid(True)
|
||||
|
||||
plt.tight_layout()
|
||||
plt.savefig("densenet_training_visualization.png")
|
||||
plt.show()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--num_classes", type=int, default=5)
|
||||
parser.add_argument("--epochs", type=int, default=30)
|
||||
parser.add_argument(
|
||||
"--batch-size", type=int, default=4
|
||||
) # 如果cuda报超出显存可以改小一点
|
||||
parser.add_argument("--lr", type=float, default=0.001)
|
||||
parser.add_argument("--lrf", type=float, default=0.1)
|
||||
|
||||
# 数据集所在根目录
|
||||
# http://download.tensorflow.org/example_images/flower_photos.tgz
|
||||
parser.add_argument("--data-path", type=str, default="../data/flower_photos")
|
||||
|
||||
# densenet121 官方权重下载地址
|
||||
# https://download.pytorch.org/models/densenet121-a639ec97.pth
|
||||
parser.add_argument(
|
||||
"--weights",
|
||||
type=str,
|
||||
default="densenet121-a639ec97.pth",
|
||||
help="initial weights path",
|
||||
)
|
||||
parser.add_argument("--freeze-layers", type=bool, default=False)
|
||||
parser.add_argument(
|
||||
"--device", default="cuda:0", help="device id (i.e. 0 or 0,1 or cpu)"
|
||||
)
|
||||
|
||||
opt = parser.parse_args()
|
||||
|
||||
main(opt)
|
||||
@@ -0,0 +1,171 @@
|
||||
import json
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
import sys
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def read_split_data(root: str, val_rate: float = 0.2):
|
||||
random.seed(0) # 保证随机结果可复现
|
||||
assert os.path.exists(root), "dataset root: {} does not exist.".format(root)
|
||||
|
||||
# 遍历文件夹,一个文件夹对应一个类别
|
||||
flower_class = [
|
||||
cla for cla in os.listdir(root) if os.path.isdir(os.path.join(root, cla))
|
||||
]
|
||||
# 排序,保证顺序一致
|
||||
flower_class.sort()
|
||||
# 生成类别名称以及对应的数字索引
|
||||
class_indices = dict((k, v) for v, k in enumerate(flower_class))
|
||||
json_str = json.dumps(
|
||||
dict((val, key) for key, val in class_indices.items()), indent=4
|
||||
)
|
||||
with open("class_indices.json", "w") as json_file:
|
||||
json_file.write(json_str)
|
||||
|
||||
train_images_path = [] # 存储训练集的所有图片路径
|
||||
train_images_label = [] # 存储训练集图片对应索引信息
|
||||
val_images_path = [] # 存储验证集的所有图片路径
|
||||
val_images_label = [] # 存储验证集图片对应索引信息
|
||||
every_class_num = [] # 存储每个类别的样本总数
|
||||
supported = [".jpg", ".JPG", ".png", ".PNG"] # 支持的文件后缀类型
|
||||
# 遍历每个文件夹下的文件
|
||||
for cla in flower_class:
|
||||
cla_path = os.path.join(root, cla)
|
||||
# 遍历获取supported支持的所有文件路径
|
||||
images = [
|
||||
os.path.join(root, cla, i)
|
||||
for i in os.listdir(cla_path)
|
||||
if os.path.splitext(i)[-1] in supported
|
||||
]
|
||||
# 获取该类别对应的索引
|
||||
image_class = class_indices[cla]
|
||||
# 记录该类别的样本数量
|
||||
every_class_num.append(len(images))
|
||||
# 按比例随机采样验证样本
|
||||
val_path = random.sample(images, k=int(len(images) * val_rate))
|
||||
|
||||
for img_path in images:
|
||||
if img_path in val_path: # 如果该路径在采样的验证集样本中则存入验证集
|
||||
val_images_path.append(img_path)
|
||||
val_images_label.append(image_class)
|
||||
else: # 否则存入训练集
|
||||
train_images_path.append(img_path)
|
||||
train_images_label.append(image_class)
|
||||
|
||||
print("{} images were found in the dataset.".format(sum(every_class_num)))
|
||||
print("{} images for training.".format(len(train_images_path)))
|
||||
print("{} images for validation.".format(len(val_images_path)))
|
||||
|
||||
plot_image = False
|
||||
if plot_image:
|
||||
# 绘制每种类别个数柱状图
|
||||
plt.bar(range(len(flower_class)), every_class_num, align="center")
|
||||
# 将横坐标0,1,2,3,4替换为相应的类别名称
|
||||
plt.xticks(range(len(flower_class)), flower_class)
|
||||
# 在柱状图上添加数值标签
|
||||
for i, v in enumerate(every_class_num):
|
||||
plt.text(x=i, y=v + 5, s=str(v), ha="center")
|
||||
# 设置x坐标
|
||||
plt.xlabel("image class")
|
||||
# 设置y坐标
|
||||
plt.ylabel("number of images")
|
||||
# 设置柱状图的标题
|
||||
plt.title("flower class distribution")
|
||||
plt.show()
|
||||
|
||||
return train_images_path, train_images_label, val_images_path, val_images_label
|
||||
|
||||
|
||||
def plot_data_loader_image(data_loader):
|
||||
batch_size = data_loader.batch_size
|
||||
plot_num = min(batch_size, 4)
|
||||
|
||||
json_path = "./class_indices.json"
|
||||
assert os.path.exists(json_path), json_path + " does not exist."
|
||||
json_file = open(json_path, "r")
|
||||
class_indices = json.load(json_file)
|
||||
|
||||
for data in data_loader:
|
||||
images, labels = data
|
||||
for i in range(plot_num):
|
||||
# [C, H, W] -> [H, W, C]
|
||||
img = images[i].numpy().transpose(1, 2, 0)
|
||||
# 反Normalize操作
|
||||
img = (img * [0.229, 0.224, 0.225] + [0.485, 0.456, 0.406]) * 255
|
||||
label = labels[i].item()
|
||||
plt.subplot(1, plot_num, i + 1)
|
||||
plt.xlabel(class_indices[str(label)])
|
||||
plt.xticks([]) # 去掉x轴的刻度
|
||||
plt.yticks([]) # 去掉y轴的刻度
|
||||
plt.imshow(img.astype("uint8"))
|
||||
plt.show()
|
||||
|
||||
|
||||
def write_pickle(list_info: list, file_name: str):
|
||||
with open(file_name, "wb") as f:
|
||||
pickle.dump(list_info, f)
|
||||
|
||||
|
||||
def read_pickle(file_name: str) -> list:
|
||||
with open(file_name, "rb") as f:
|
||||
info_list = pickle.load(f)
|
||||
return info_list
|
||||
|
||||
|
||||
def train_one_epoch(model, optimizer, data_loader, device, epoch):
|
||||
model.train()
|
||||
loss_function = torch.nn.CrossEntropyLoss()
|
||||
mean_loss = torch.zeros(1).to(device)
|
||||
optimizer.zero_grad()
|
||||
|
||||
data_loader = tqdm(data_loader)
|
||||
|
||||
for step, data in enumerate(data_loader):
|
||||
images, labels = data
|
||||
|
||||
pred = model(images.to(device))
|
||||
|
||||
loss = loss_function(pred, labels.to(device))
|
||||
loss.backward()
|
||||
mean_loss = (mean_loss * step + loss.detach()) / (
|
||||
step + 1
|
||||
) # update mean losses
|
||||
|
||||
data_loader.desc = "[epoch {}] mean loss {}".format(
|
||||
epoch, round(mean_loss.item(), 3)
|
||||
)
|
||||
|
||||
if not torch.isfinite(loss):
|
||||
print("WARNING: non-finite loss, ending training ", loss)
|
||||
sys.exit(1)
|
||||
|
||||
optimizer.step()
|
||||
optimizer.zero_grad()
|
||||
|
||||
return mean_loss.item()
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def evaluate(model, data_loader, device):
|
||||
model.eval()
|
||||
|
||||
# 验证样本总个数
|
||||
total_num = len(data_loader.dataset)
|
||||
|
||||
# 用于存储预测正确的样本个数
|
||||
sum_num = torch.zeros(1).to(device)
|
||||
|
||||
data_loader = tqdm(data_loader)
|
||||
|
||||
for step, data in enumerate(data_loader):
|
||||
images, labels = data
|
||||
pred = model(images.to(device))
|
||||
pred = torch.max(pred, dim=1)[1]
|
||||
sum_num += torch.eq(pred, labels.to(device)).sum()
|
||||
|
||||
return sum_num.item() / total_num
|
||||
Reference in new issue
Block a user