Initial commit

This commit is contained in:
xixu-me committed 2025-04-20 21:02:44 +08:00
commit 80fa1d394d
77 files changed
+5568

No files matched your search

+7
View File
@@ -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

+343
View File
@@ -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()
+286
View File
@@ -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.")
+37
View File
@@ -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
+66
View File
@@ -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()
+286
View File
@@ -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)
+171
View File
@@ -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