commit 80fa1d394db2375b09c0658869f10fa36c96b466 Author: Xi Xu Date: Sun Apr 20 21:02:44 2025 +0800 Initial commit diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..15c6ac2 --- /dev/null +++ b/.gitignore @@ -0,0 +1,159 @@ +# Byte-compiled / optimized / DLL files +__pycache__/ +*.py[cod] +*$py.class + +# C extensions +*.so + +# Distribution / packaging +.Python +build/ +develop-eggs/ +dist/ +downloads/ +eggs/ +.eggs/ +lib/ +lib64/ +parts/ +sdist/ +var/ +wheels/ +share/python-wheels/ +*.egg-info/ +.installed.cfg +*.egg +MANIFEST + +# PyInstaller +# Usually these files are written by a python script from a template +# before PyInstaller builds the exe, so as to inject date/other infos into it. +*.manifest +*.spec + +# Installer logs +pip-log.txt +pip-delete-this-directory.txt + +# Unit test / coverage reports +htmlcov/ +.tox/ +.nox/ +.coverage +.coverage.* +.cache +nosetests.xml +coverage.xml +*.cover +*.py,cover +.hypothesis/ +.pytest_cache/ +cover/ + +# Translations +*.mo +*.pot + +# Django stuff: +*.log +local_settings.py +db.sqlite3 +db.sqlite3-journal + +# Flask stuff: +instance/ +.webassets-cache + +# Scrapy stuff: +.scrapy + +# Sphinx documentation +docs/_build/ + +# PyBuilder +.pybuilder/ +target/ + +# Jupyter Notebook +.ipynb_checkpoints + +# IPython +profile_default/ +ipython_config.py + +# pyenv +# For a library or package, you might want to ignore these files since the code is +# intended to run in multiple environments; otherwise, check them in: +# .python-version + +# pipenv +# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control. +# However, in case of collaboration, if having platform-specific dependencies or dependencies +# having no cross-platform support, pipenv may install dependencies that don't work, or not +# install all needed dependencies. +#Pipfile.lock + +# poetry +# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control. +# This is especially recommended for binary packages to ensure reproducibility, and is more +# commonly ignored for libraries. +# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control +#poetry.lock + +# pdm +# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control. +#pdm.lock +# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it +# in version control. +# https://pdm.fming.dev/#use-with-ide +.pdm.toml + +# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm +__pypackages__/ + +# Celery stuff +celerybeat-schedule +celerybeat.pid + +# SageMath parsed files +*.sage.py + +# Environments +.env +.venv +env/ +venv/ +ENV/ +env.bak/ +venv.bak/ + +# Spyder project settings +.spyderproject +.spyproject + +# Rope project settings +.ropeproject + +# mkdocs documentation +/site + +# mypy +.mypy_cache/ +.dmypy.json +dmypy.json + +# Pyre type checker +.pyre/ + +# pytype static type analyzer +.pytype/ + +# Cython debug symbols +cython_debug/ + +# Data files +data/ +weights/ +EfficientNet/efficientnetb?.pth +Transformer/jx_vit_*.pth diff --git a/DenseNet/class_indices.json b/DenseNet/class_indices.json new file mode 100644 index 0000000..4849997 --- /dev/null +++ b/DenseNet/class_indices.json @@ -0,0 +1,7 @@ +{ + "0": "daisy", + "1": "dandelion", + "2": "roses", + "3": "sunflowers", + "4": "tulips" +} \ No newline at end of file diff --git a/DenseNet/densenet_error_analysis.png b/DenseNet/densenet_error_analysis.png new file mode 100644 index 0000000..9efdcf0 Binary files /dev/null and b/DenseNet/densenet_error_analysis.png differ diff --git a/DenseNet/densenet_evaluation_metrics.png b/DenseNet/densenet_evaluation_metrics.png new file mode 100644 index 0000000..fefa912 Binary files /dev/null and b/DenseNet/densenet_evaluation_metrics.png differ diff --git a/DenseNet/densenet_training_visualization.png b/DenseNet/densenet_training_visualization.png new file mode 100644 index 0000000..1ad8a78 Binary files /dev/null and b/DenseNet/densenet_training_visualization.png differ diff --git a/DenseNet/evaluate.py b/DenseNet/evaluate.py new file mode 100644 index 0000000..8390afc --- /dev/null +++ b/DenseNet/evaluate.py @@ -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() diff --git a/DenseNet/model.py b/DenseNet/model.py new file mode 100644 index 0000000..b09c591 --- /dev/null +++ b/DenseNet/model.py @@ -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.") diff --git a/DenseNet/my_dataset.py b/DenseNet/my_dataset.py new file mode 100644 index 0000000..1dabf48 --- /dev/null +++ b/DenseNet/my_dataset.py @@ -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 diff --git a/DenseNet/predict.py b/DenseNet/predict.py new file mode 100644 index 0000000..a0f81b3 --- /dev/null +++ b/DenseNet/predict.py @@ -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() diff --git a/DenseNet/train.py b/DenseNet/train.py new file mode 100644 index 0000000..4112f90 --- /dev/null +++ b/DenseNet/train.py @@ -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) diff --git a/DenseNet/utils.py b/DenseNet/utils.py new file mode 100644 index 0000000..359d43e --- /dev/null +++ b/DenseNet/utils.py @@ -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 diff --git a/EfficientNet/class_indices.json b/EfficientNet/class_indices.json new file mode 100644 index 0000000..4849997 --- /dev/null +++ b/EfficientNet/class_indices.json @@ -0,0 +1,7 @@ +{ + "0": "daisy", + "1": "dandelion", + "2": "roses", + "3": "sunflowers", + "4": "tulips" +} \ No newline at end of file diff --git a/EfficientNet/evaluate_models.py b/EfficientNet/evaluate_models.py new file mode 100644 index 0000000..0339168 --- /dev/null +++ b/EfficientNet/evaluate_models.py @@ -0,0 +1,342 @@ +import argparse +import json +import os +import time + +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +import seaborn as sns +import torch +import torchvision.transforms as transforms + +# Import from project files +from model import ( + efficientnet_b0, + efficientnet_b1, + efficientnet_b2, + efficientnet_b3, + efficientnet_b4, + efficientnet_b5, +) +from my_dataset import MyDataSet +from sklearn.metrics import ( + classification_report, + confusion_matrix, + f1_score, + precision_score, + recall_score, +) +from torch.utils.data import DataLoader +from utils import read_split_data + + +def load_model(model_name, num_classes, device): + """Load a specific EfficientNet model with its best weights""" + model_functions = { + "B0": efficientnet_b0, + "B1": efficientnet_b1, + "B2": efficientnet_b2, + "B3": efficientnet_b3, + "B4": efficientnet_b4, + "B5": efficientnet_b5, + } + + # Create the model + model = model_functions[model_name](num_classes=num_classes).to(device) + + # Load weights + weights_path = f"./weights/{model_name}/best_model.pth" + if os.path.exists(weights_path): + model.load_state_dict(torch.load(weights_path, map_location=device)) + print(f"Loaded weights from {weights_path}") + else: + print(f"WARNING: No weights found at {weights_path}") + + return model + + +def evaluate_model(model, data_loader, device, class_names): + """Evaluate model performance with multiple metrics""" + model.eval() + + all_preds = [] + all_labels = [] + inference_times = [] + + with torch.no_grad(): + for images, labels in data_loader: + images, labels = images.to(device), labels.to(device) + + # Measure inference time + start_time = time.time() + outputs = model(images) + end_time = time.time() + + # Record batch inference time + inference_times.append(end_time - start_time) + + # Get predictions + _, preds = torch.max(outputs, 1) + + # Store predictions and labels + all_preds.extend(preds.cpu().numpy()) + all_labels.extend(labels.cpu().numpy()) + + # Convert to numpy arrays for sklearn metrics + all_preds = np.array(all_preds) + all_labels = np.array(all_labels) + + # Calculate metrics + results = { + "accuracy": np.mean(all_preds == all_labels), + "precision_macro": precision_score(all_labels, all_preds, average="macro"), + "recall_macro": recall_score(all_labels, all_preds, average="macro"), + "f1_macro": f1_score(all_labels, all_preds, average="macro"), + "precision_weighted": precision_score( + all_labels, all_preds, average="weighted" + ), + "recall_weighted": recall_score(all_labels, all_preds, average="weighted"), + "f1_weighted": f1_score(all_labels, all_preds, average="weighted"), + "avg_inference_time": np.mean(inference_times), + "total_inference_time": np.sum(inference_times), + "images_per_second": len(all_labels) / np.sum(inference_times), + "confusion_matrix": confusion_matrix(all_labels, all_preds), + } + + # Per-class metrics + class_report = classification_report( + all_labels, all_preds, target_names=class_names, output_dict=True + ) + + # Add per-class metrics to results + for class_name in class_names: + results[f"precision_{class_name}"] = class_report[class_name]["precision"] + results[f"recall_{class_name}"] = class_report[class_name]["recall"] + results[f"f1_{class_name}"] = class_report[class_name]["f1-score"] + + return results + + +def plot_confusion_matrix(cm, class_names, model_name, save_dir): + """Plot confusion matrix for a model""" + plt.figure(figsize=(10, 8)) + sns.heatmap( + cm, + annot=True, + fmt="d", + cmap="Blues", + xticklabels=class_names, + yticklabels=class_names, + ) + plt.title(f"Confusion Matrix - EfficientNet-{model_name}") + plt.ylabel("True Label") + plt.xlabel("Predicted Label") + plt.tight_layout() + + # Save figure + if not os.path.exists(save_dir): + os.makedirs(save_dir) + plt.savefig(f"{save_dir}/confusion_matrix_{model_name}.png") + plt.close() + + +def plot_metrics_comparison(metrics_df, save_dir): + """Plot comparison of metrics across models""" + # Create plots directory if it doesn't exist + if not os.path.exists(save_dir): + os.makedirs(save_dir) + + # Selected metrics to plot + metrics_to_plot = [ + "accuracy", + "precision_macro", + "recall_macro", + "f1_macro", + "images_per_second", + "avg_inference_time", + ] + + # Plot each metric + for metric in metrics_to_plot: + plt.figure(figsize=(10, 6)) + plt.bar(metrics_df["Model"], metrics_df[metric]) + plt.title(f'{metric.replace("_", " ").title()} Comparison') + plt.xlabel("Model") + plt.ylabel(metric.replace("_", " ").title()) + plt.xticks(rotation=45) + plt.tight_layout() + plt.savefig(f"{save_dir}/{metric}_comparison.png") + plt.close() + + # Create radar chart for comparing models + metrics_for_radar = ["accuracy", "precision_macro", "recall_macro", "f1_macro"] + + # Normalize data for radar chart + radar_data = metrics_df[metrics_for_radar].values + radar_data_normalized = radar_data / radar_data.max(axis=0) + + # Plot radar chart + plt.figure(figsize=(12, 10)) + angles = np.linspace(0, 2 * np.pi, len(metrics_for_radar), endpoint=False) + angles = np.concatenate((angles, [angles[0]])) + + for i, model in enumerate(metrics_df["Model"]): + values = radar_data_normalized[i] + values = np.concatenate((values, [values[0]])) + plt.polar(angles, values, marker="o", label=f"EfficientNet-{model}") + + plt.xticks(angles[:-1], metrics_for_radar) + plt.title("Model Performance Comparison (Normalized)") + plt.legend(loc="upper right") + plt.tight_layout() + plt.savefig(f"{save_dir}/radar_chart_comparison.png") + plt.close() + + +def main(args): + device = torch.device(args.device if torch.cuda.is_available() else "cpu") + print(f"Using device: {device}") + + # Load class names + with open("class_indices.json") as f: + class_mapping = json.load(f) + class_names = [class_mapping[str(i)] for i in range(len(class_mapping))] + + # Load dataset + _, _, val_images_path, val_images_label = read_split_data(args.data_path) + + # Dictionary to store results for all models + all_results = [] + + # Models to evaluate + models_to_evaluate = ["B0", "B1", "B2", "B3", "B4", "B5"] + + # Create results directory + results_dir = "./evaluation_results" + if not os.path.exists(results_dir): + os.makedirs(results_dir) + + for model_name in models_to_evaluate: + print(f"\n{'='*50}") + print(f"Evaluating EfficientNet-{model_name}") + print(f"{'='*50}") + + # Set up data transformations for this model + img_size = { + "B0": 224, + "B1": 240, + "B2": 260, + "B3": 300, + "B4": 380, + "B5": 456, + } + + data_transform = transforms.Compose( + [ + transforms.Resize(img_size[model_name]), + transforms.CenterCrop(img_size[model_name]), + transforms.ToTensor(), + transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), + ] + ) + + # Create validation dataset + val_dataset = MyDataSet( + images_path=val_images_path, + images_class=val_images_label, + transform=data_transform, + ) + + # Create validation dataloader + val_loader = DataLoader( + val_dataset, + batch_size=args.batch_size, + shuffle=False, + num_workers=min( + [os.cpu_count(), args.batch_size if args.batch_size > 1 else 0, 8] + ), + pin_memory=True, + collate_fn=val_dataset.collate_fn, + ) + + # Load model + model = load_model(model_name, args.num_classes, device) + + # Evaluate model + results = evaluate_model(model, val_loader, device, class_names) + + # Add model name to results + results["model"] = f"EfficientNet-{model_name}" + + # Print results summary + print(f"\nResults for EfficientNet-{model_name}:") + print(f"Accuracy: {results['accuracy']*100:.2f}%") + print(f"Precision (macro): {results['precision_macro']*100:.2f}%") + print(f"Recall (macro): {results['recall_macro']*100:.2f}%") + print(f"F1 Score (macro): {results['f1_macro']*100:.2f}%") + print( + f"Average inference time: {results['avg_inference_time']*1000:.2f} ms per batch" + ) + print(f"Inference speed: {results['images_per_second']:.2f} images/second") + + # Plot confusion matrix + plot_confusion_matrix( + results["confusion_matrix"], + class_names, + model_name, + os.path.join(results_dir, "confusion_matrices"), + ) + + # Store results excluding confusion matrix + results_for_df = {k: v for k, v in results.items() if k != "confusion_matrix"} + all_results.append(results_for_df) + + # Create DataFrame from all results + results_df = pd.DataFrame(all_results) + + # Prepare DataFrame for plotting + plot_df = results_df.rename(columns={"model": "Model"}) + plot_df["Model"] = plot_df["Model"].str.replace("EfficientNet-", "") + + # Plot comparison metrics + plot_metrics_comparison(plot_df, os.path.join(results_dir, "comparisons")) + + # Save results to CSV + results_df.to_csv(os.path.join(results_dir, "evaluation_results.csv"), index=False) + print(f"\nResults saved to {os.path.join(results_dir, 'evaluation_results.csv')}") + + # Display comparison table + comparison_cols = [ + "model", + "accuracy", + "precision_macro", + "recall_macro", + "f1_macro", + "avg_inference_time", + "images_per_second", + ] + + comparison_df = results_df[comparison_cols].sort_values("accuracy", ascending=False) + print("\nModel Performance Comparison (sorted by accuracy):") + print(comparison_df.to_string(index=False)) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description="Evaluate EfficientNet models") + parser.add_argument( + "--data-path", + type=str, + default=r"..\\data\\flower_photos", + help="dataset root path", + ) + parser.add_argument("--num_classes", type=int, default=5, help="number of classes") + parser.add_argument( + "--batch-size", type=int, default=16, help="batch size for evaluation" + ) + parser.add_argument( + "--device", default="cuda:0", help="device id (i.e. 0 or 0,1 or cpu)" + ) + + args = parser.parse_args() + main(args) diff --git a/EfficientNet/evaluation_results/comparisons/accuracy_comparison.png b/EfficientNet/evaluation_results/comparisons/accuracy_comparison.png new file mode 100644 index 0000000..d69bb92 Binary files /dev/null and b/EfficientNet/evaluation_results/comparisons/accuracy_comparison.png differ diff --git a/EfficientNet/evaluation_results/comparisons/avg_inference_time_comparison.png b/EfficientNet/evaluation_results/comparisons/avg_inference_time_comparison.png new file mode 100644 index 0000000..c8d43e2 Binary files /dev/null and b/EfficientNet/evaluation_results/comparisons/avg_inference_time_comparison.png differ diff --git a/EfficientNet/evaluation_results/comparisons/f1_macro_comparison.png b/EfficientNet/evaluation_results/comparisons/f1_macro_comparison.png new file mode 100644 index 0000000..254f4ee Binary files /dev/null and b/EfficientNet/evaluation_results/comparisons/f1_macro_comparison.png differ diff --git a/EfficientNet/evaluation_results/comparisons/images_per_second_comparison.png b/EfficientNet/evaluation_results/comparisons/images_per_second_comparison.png new file mode 100644 index 0000000..91c6335 Binary files /dev/null and b/EfficientNet/evaluation_results/comparisons/images_per_second_comparison.png differ diff --git a/EfficientNet/evaluation_results/comparisons/precision_macro_comparison.png b/EfficientNet/evaluation_results/comparisons/precision_macro_comparison.png new file mode 100644 index 0000000..39b9acd Binary files /dev/null and b/EfficientNet/evaluation_results/comparisons/precision_macro_comparison.png differ diff --git a/EfficientNet/evaluation_results/comparisons/radar_chart_comparison.png b/EfficientNet/evaluation_results/comparisons/radar_chart_comparison.png new file mode 100644 index 0000000..e8082e9 Binary files /dev/null and b/EfficientNet/evaluation_results/comparisons/radar_chart_comparison.png differ diff --git a/EfficientNet/evaluation_results/comparisons/recall_macro_comparison.png b/EfficientNet/evaluation_results/comparisons/recall_macro_comparison.png new file mode 100644 index 0000000..38780e9 Binary files /dev/null and b/EfficientNet/evaluation_results/comparisons/recall_macro_comparison.png differ diff --git a/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B0.png b/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B0.png new file mode 100644 index 0000000..e5f97c5 Binary files /dev/null and b/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B0.png differ diff --git a/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B1.png b/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B1.png new file mode 100644 index 0000000..9737ae2 Binary files /dev/null and b/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B1.png differ diff --git a/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B2.png b/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B2.png new file mode 100644 index 0000000..2aa6751 Binary files /dev/null and b/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B2.png differ diff --git a/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B3.png b/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B3.png new file mode 100644 index 0000000..fd5564a Binary files /dev/null and b/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B3.png differ diff --git a/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B4.png b/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B4.png new file mode 100644 index 0000000..8464594 Binary files /dev/null and b/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B4.png differ diff --git a/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B5.png b/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B5.png new file mode 100644 index 0000000..1ccb398 Binary files /dev/null and b/EfficientNet/evaluation_results/confusion_matrices/confusion_matrix_B5.png differ diff --git a/EfficientNet/evaluation_results/evaluation_results.csv b/EfficientNet/evaluation_results/evaluation_results.csv new file mode 100644 index 0000000..364bcec --- /dev/null +++ b/EfficientNet/evaluation_results/evaluation_results.csv @@ -0,0 +1,7 @@ +accuracy,precision_macro,recall_macro,f1_macro,precision_weighted,recall_weighted,f1_weighted,avg_inference_time,total_inference_time,images_per_second,precision_daisy,recall_daisy,f1_daisy,precision_dandelion,recall_dandelion,f1_dandelion,precision_roses,recall_roses,f1_roses,precision_sunflowers,recall_sunflowers,f1_sunflowers,precision_tulips,recall_tulips,f1_tulips,model +0.957592339261286,0.9577966282225473,0.9562733310129442,0.9567672214182226,0.9583364479556171,0.957592339261286,0.9576985362250258,0.016896361890046493,0.7772326469421387,940.5163342996058,0.975609756097561,0.9523809523809523,0.963855421686747,0.9567567567567568,0.9888268156424581,0.9725274725274725,0.9029850746268657,0.9453125,0.9236641221374046,0.9925925925925926,0.9640287769784173,0.9781021897810219,0.961038961038961,0.9308176100628931,0.9456869009584664,EfficientNet-B0 +0.9712722298221614,0.9696336624400939,0.9710926049211679,0.9702721673943255,0.9715562886900023,0.9712722298221614,0.9713231175356396,0.014946569567141325,0.687542200088501,1063.2074655285232,0.9534883720930233,0.9761904761904762,0.9647058823529412,0.9831460674157303,0.9776536312849162,0.9803921568627451,0.9534883720930233,0.9609375,0.9571984435797666,0.9645390070921985,0.9784172661870504,0.9714285714285714,0.9935064935064936,0.9622641509433962,0.9776357827476039,EfficientNet-B1 +0.9767441860465116,0.9758506980950198,0.9759251028569246,0.975863977685071,0.9767211148025102,0.9767441860465116,0.9767118713443993,0.014863568803538446,0.6837241649627686,1069.1446016681389,0.96875,0.9841269841269841,0.9763779527559056,0.9888268156424581,0.9888268156424581,0.9888268156424581,0.9603174603174603,0.9453125,0.952755905511811,0.9928057553956835,0.9928057553956835,0.9928057553956835,0.9685534591194969,0.9685534591194969,0.9685534591194969,EfficientNet-B2 +0.9740082079343365,0.9741130496386882,0.9735993455382512,0.9737579622397536,0.9742066830281588,0.9740082079343365,0.9740063993250383,0.019874458727629288,0.9142251014709473,799.5842586512377,0.96875,0.9841269841269841,0.9763779527559056,0.9672131147540983,0.9888268156424581,0.9779005524861878,0.9612403100775194,0.96875,0.9649805447470817,0.9925925925925926,0.9640287769784173,0.9781021897810219,0.9807692307692307,0.9622641509433962,0.9714285714285714,EfficientNet-B3 +0.9767441860465116,0.9752607491378678,0.9772841995155099,0.9761568002196187,0.9769907043693243,0.9767441860465116,0.9767535240605703,0.01837624674258025,0.8453073501586914,864.7742148021873,0.96875,0.9841269841269841,0.9763779527559056,0.9943502824858758,0.9832402234636871,0.9887640449438202,0.946969696969697,0.9765625,0.9615384615384616,0.9857142857142858,0.9928057553956835,0.989247311827957,0.9805194805194806,0.949685534591195,0.9648562300319489,EfficientNet-B4 +0.9753761969904241,0.9741550080328338,0.9756800058078792,0.9747398058300302,0.9757384563373198,0.9753761969904241,0.9753865089344038,0.022696059683094853,1.0440187454223633,700.1789988974481,0.968503937007874,0.9761904761904762,0.9723320158102767,0.9888268156424581,0.9888268156424581,0.9888268156424581,0.9402985074626866,0.984375,0.9618320610687023,0.9927536231884058,0.9856115107913669,0.9891696750902527,0.9803921568627451,0.9433962264150944,0.9615384615384616,EfficientNet-B5 diff --git a/EfficientNet/model.py b/EfficientNet/model.py new file mode 100644 index 0000000..8eb0374 --- /dev/null +++ b/EfficientNet/model.py @@ -0,0 +1,360 @@ +import copy +import math +from collections import OrderedDict +from functools import partial +from typing import Callable, Optional + +import torch +import torch.nn as nn +from torch import Tensor +from torch.nn import functional as F + + +def _make_divisible(ch, divisor=8, min_ch=None): + + if min_ch is None: + min_ch = divisor + new_ch = max(min_ch, int(ch + divisor / 2) // divisor * divisor) + if new_ch < 0.9 * ch: + new_ch += divisor + return new_ch + + +def drop_path(x, drop_prob: float = 0.0, training: bool = False): + if drop_prob == 0.0 or not training: + return x + keep_prob = 1 - drop_prob + shape = (x.shape[0],) + (1,) * ( + x.ndim - 1 + ) # work with diff dim tensors, not just 2D ConvNets + random_tensor = keep_prob + torch.rand(shape, dtype=x.dtype, device=x.device) + random_tensor.floor_() # binarize + output = x.div(keep_prob) * random_tensor + return output + + +class DropPath(nn.Module): + def __init__(self, drop_prob=None): + super(DropPath, self).__init__() + self.drop_prob = drop_prob + + def forward(self, x): + return drop_path(x, self.drop_prob, self.training) + + +class ConvBNActivation(nn.Sequential): + def __init__( + self, + in_planes: int, + out_planes: int, + kernel_size: int = 3, + stride: int = 1, + groups: int = 1, + norm_layer: Optional[Callable[..., nn.Module]] = None, + activation_layer: Optional[Callable[..., nn.Module]] = None, + ): + padding = (kernel_size - 1) // 2 + if norm_layer is None: + norm_layer = nn.BatchNorm2d + if activation_layer is None: + activation_layer = nn.SiLU # alias Swish (torch>=1.7) + super(ConvBNActivation, self).__init__( + nn.Conv2d( + in_channels=in_planes, + out_channels=out_planes, + kernel_size=kernel_size, + stride=stride, + padding=padding, + groups=groups, + bias=False, + ), + norm_layer(out_planes), + activation_layer(), + ) + + +class SqueezeExcitation(nn.Module): + def __init__(self, input_c: int, expand_c: int, squeeze_factor: int = 4): + super(SqueezeExcitation, self).__init__() + squeeze_c = input_c // squeeze_factor + self.fc1 = nn.Conv2d(expand_c, squeeze_c, 1) + self.ac1 = nn.SiLU() + self.fc2 = nn.Conv2d(squeeze_c, expand_c, 1) + self.ac2 = nn.Sigmoid() + + def forward(self, x: Tensor) -> Tensor: + scale = F.adaptive_avg_pool2d(x, output_size=(1, 1)) + scale = self.fc1(scale) + scale = self.ac1(scale) + scale = self.fc2(scale) + scale = self.ac2(scale) + return scale * x + + +class InvertedResidualConfig: + def __init__( + self, + kernel: int, + input_c: int, + out_c: int, + expanded_ratio: int, + stride: int, + use_se: bool, + drop_rate: float, + index: str, + width_coefficient: float, + ): + self.input_c = self.adjust_channels(input_c, width_coefficient) + self.kernel = kernel + self.expanded_c = self.input_c * expanded_ratio + self.out_c = self.adjust_channels(out_c, width_coefficient) + self.use_se = use_se + self.stride = stride + self.drop_rate = drop_rate + self.index = index + + @staticmethod + def adjust_channels(channels: int, width_coefficient: float): + return _make_divisible(channels * width_coefficient, 8) + + +class InvertedResidual(nn.Module): + def __init__( + self, cnf: InvertedResidualConfig, norm_layer: Callable[..., nn.Module] + ): + super(InvertedResidual, self).__init__() + if cnf.stride not in [1, 2]: + raise ValueError("illegal stride value.") + self.use_res_connect = cnf.stride == 1 and cnf.input_c == cnf.out_c + layers = OrderedDict() + activation_layer = nn.SiLU + if cnf.expanded_c != cnf.input_c: + layers.update( + { + "expand_conv": ConvBNActivation( + cnf.input_c, + cnf.expanded_c, + kernel_size=1, + norm_layer=norm_layer, + activation_layer=activation_layer, + ) + } + ) + layers.update( + { + "dwconv": ConvBNActivation( + cnf.expanded_c, + cnf.expanded_c, + kernel_size=cnf.kernel, + stride=cnf.stride, + groups=cnf.expanded_c, + norm_layer=norm_layer, + activation_layer=activation_layer, + ) + } + ) + if cnf.use_se: + layers.update({"se": SqueezeExcitation(cnf.input_c, cnf.expanded_c)}) + layers.update( + { + "project_conv": ConvBNActivation( + cnf.expanded_c, + cnf.out_c, + kernel_size=1, + norm_layer=norm_layer, + activation_layer=nn.Identity, + ) + } + ) + self.block = nn.Sequential(layers) + self.out_channels = cnf.out_c + self.is_strided = cnf.stride > 1 + if self.use_res_connect and cnf.drop_rate > 0: + self.dropout = DropPath(cnf.drop_rate) + else: + self.dropout = nn.Identity() + + def forward(self, x: Tensor) -> Tensor: + result = self.block(x) + result = self.dropout(result) + if self.use_res_connect: + result += x + return result + + +class EfficientNet(nn.Module): + def __init__( + self, + width_coefficient: float, + depth_coefficient: float, + num_classes: int = 1000, + dropout_rate: float = 0.2, + drop_connect_rate: float = 0.2, + block: Optional[Callable[..., nn.Module]] = None, + norm_layer: Optional[Callable[..., nn.Module]] = None, + ): + super(EfficientNet, self).__init__() + default_cnf = [ + [3, 32, 16, 1, 1, True, drop_connect_rate, 1], + [3, 16, 24, 6, 2, True, drop_connect_rate, 2], + [5, 24, 40, 6, 2, True, drop_connect_rate, 2], + [3, 40, 80, 6, 2, True, drop_connect_rate, 3], + [5, 80, 112, 6, 1, True, drop_connect_rate, 3], + [5, 112, 192, 6, 2, True, drop_connect_rate, 4], + [3, 192, 320, 6, 1, True, drop_connect_rate, 1], + ] + + def round_repeats(repeats): + return int(math.ceil(depth_coefficient * repeats)) + + if block is None: + block = InvertedResidual + if norm_layer is None: + norm_layer = partial(nn.BatchNorm2d, eps=1e-3, momentum=0.1) + + adjust_channels = partial( + InvertedResidualConfig.adjust_channels, width_coefficient=width_coefficient + ) + bneck_conf = partial( + InvertedResidualConfig, width_coefficient=width_coefficient + ) + b = 0 + num_blocks = float(sum(round_repeats(i[-1]) for i in default_cnf)) + inverted_residual_setting = [] + for stage, args in enumerate(default_cnf): + cnf = copy.copy(args) + for i in range(round_repeats(cnf.pop(-1))): + if i > 0: + cnf[-3] = 1 + cnf[1] = cnf[2] + cnf[-1] = args[-2] * b / num_blocks + index = str(stage + 1) + chr(i + 97) + inverted_residual_setting.append(bneck_conf(*cnf, index)) + b += 1 + layers = OrderedDict() + layers.update( + { + "stem_conv": ConvBNActivation( + in_planes=3, + out_planes=adjust_channels(32), + kernel_size=3, + stride=2, + norm_layer=norm_layer, + ) + } + ) + for cnf in inverted_residual_setting: + layers.update({cnf.index: block(cnf, norm_layer)}) + last_conv_input_c = inverted_residual_setting[-1].out_c + last_conv_output_c = adjust_channels(1280) + layers.update( + { + "top": ConvBNActivation( + in_planes=last_conv_input_c, + out_planes=last_conv_output_c, + kernel_size=1, + norm_layer=norm_layer, + ) + } + ) + self.features = nn.Sequential(layers) + self.avgpool = nn.AdaptiveAvgPool2d(1) + classifier = [] + if dropout_rate > 0: + classifier.append(nn.Dropout(p=dropout_rate, inplace=True)) + classifier.append(nn.Linear(last_conv_output_c, num_classes)) + self.classifier = nn.Sequential(*classifier) + + for m in self.modules(): + if isinstance(m, nn.Conv2d): + nn.init.kaiming_normal_(m.weight, mode="fan_out") + if m.bias is not None: + nn.init.zeros_(m.bias) + elif isinstance(m, nn.BatchNorm2d): + nn.init.ones_(m.weight) + nn.init.zeros_(m.bias) + elif isinstance(m, nn.Linear): + nn.init.normal_(m.weight, 0, 0.01) + nn.init.zeros_(m.bias) + + def _forward_impl(self, x: Tensor) -> Tensor: + x = self.features(x) + x = self.avgpool(x) + x = torch.flatten(x, 1) + x = self.classifier(x) + return x + + def forward(self, x: Tensor) -> Tensor: + return self._forward_impl(x) + + +def efficientnet_b0(num_classes=1000): + return EfficientNet( + width_coefficient=1.0, + depth_coefficient=1.0, + dropout_rate=0.2, + num_classes=num_classes, + ) + + +def efficientnet_b1(num_classes=1000): + return EfficientNet( + width_coefficient=1.0, + depth_coefficient=1.1, + dropout_rate=0.2, + num_classes=num_classes, + ) + + +def efficientnet_b2(num_classes=1000): + return EfficientNet( + width_coefficient=1.1, + depth_coefficient=1.2, + dropout_rate=0.3, + num_classes=num_classes, + ) + + +def efficientnet_b3(num_classes=1000): + return EfficientNet( + width_coefficient=1.2, + depth_coefficient=1.4, + dropout_rate=0.3, + num_classes=num_classes, + ) + + +def efficientnet_b4(num_classes=1000): + return EfficientNet( + width_coefficient=1.4, + depth_coefficient=1.8, + dropout_rate=0.4, + num_classes=num_classes, + ) + + +def efficientnet_b5(num_classes=1000): + return EfficientNet( + width_coefficient=1.6, + depth_coefficient=2.2, + dropout_rate=0.4, + num_classes=num_classes, + ) + + +def efficientnet_b6(num_classes=1000): + return EfficientNet( + width_coefficient=1.8, + depth_coefficient=2.6, + dropout_rate=0.5, + num_classes=num_classes, + ) + + +def efficientnet_b7(num_classes=1000): + return EfficientNet( + width_coefficient=2.0, + depth_coefficient=3.1, + dropout_rate=0.5, + num_classes=num_classes, + ) diff --git a/EfficientNet/my_dataset.py b/EfficientNet/my_dataset.py new file mode 100644 index 0000000..3e16f09 --- /dev/null +++ b/EfficientNet/my_dataset.py @@ -0,0 +1,29 @@ +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]) + if img.mode != "RGB": + img = img.convert("RGB") + label = self.images_class[item] + if self.transform is not None: + img = self.transform(img) + return img, label + + @staticmethod + def collate_fn(batch): + images, labels = tuple(zip(*batch)) + images = torch.stack(images, dim=0) + labels = torch.as_tensor(labels) + return images, labels diff --git a/EfficientNet/plots/B0/learning_rate.png b/EfficientNet/plots/B0/learning_rate.png new file mode 100644 index 0000000..87062ef Binary files /dev/null and b/EfficientNet/plots/B0/learning_rate.png differ diff --git a/EfficientNet/plots/B0/training_loss.png b/EfficientNet/plots/B0/training_loss.png new file mode 100644 index 0000000..a294a05 Binary files /dev/null and b/EfficientNet/plots/B0/training_loss.png differ diff --git a/EfficientNet/plots/B0/training_summary.png b/EfficientNet/plots/B0/training_summary.png new file mode 100644 index 0000000..50ae786 Binary files /dev/null and b/EfficientNet/plots/B0/training_summary.png differ diff --git a/EfficientNet/plots/B0/validation_accuracy.png b/EfficientNet/plots/B0/validation_accuracy.png new file mode 100644 index 0000000..04ec82d Binary files /dev/null and b/EfficientNet/plots/B0/validation_accuracy.png differ diff --git a/EfficientNet/plots/B1/learning_rate.png b/EfficientNet/plots/B1/learning_rate.png new file mode 100644 index 0000000..87062ef Binary files /dev/null and b/EfficientNet/plots/B1/learning_rate.png differ diff --git a/EfficientNet/plots/B1/training_loss.png b/EfficientNet/plots/B1/training_loss.png new file mode 100644 index 0000000..3f53f04 Binary files /dev/null and b/EfficientNet/plots/B1/training_loss.png differ diff --git a/EfficientNet/plots/B1/training_summary.png b/EfficientNet/plots/B1/training_summary.png new file mode 100644 index 0000000..2c06e6d Binary files /dev/null and b/EfficientNet/plots/B1/training_summary.png differ diff --git a/EfficientNet/plots/B1/validation_accuracy.png b/EfficientNet/plots/B1/validation_accuracy.png new file mode 100644 index 0000000..e403cda Binary files /dev/null and b/EfficientNet/plots/B1/validation_accuracy.png differ diff --git a/EfficientNet/plots/B2/learning_rate.png b/EfficientNet/plots/B2/learning_rate.png new file mode 100644 index 0000000..87062ef Binary files /dev/null and b/EfficientNet/plots/B2/learning_rate.png differ diff --git a/EfficientNet/plots/B2/training_loss.png b/EfficientNet/plots/B2/training_loss.png new file mode 100644 index 0000000..2a2fd94 Binary files /dev/null and b/EfficientNet/plots/B2/training_loss.png differ diff --git a/EfficientNet/plots/B2/training_summary.png b/EfficientNet/plots/B2/training_summary.png new file mode 100644 index 0000000..6e7ce03 Binary files /dev/null and b/EfficientNet/plots/B2/training_summary.png differ diff --git a/EfficientNet/plots/B2/validation_accuracy.png b/EfficientNet/plots/B2/validation_accuracy.png new file mode 100644 index 0000000..7ba4937 Binary files /dev/null and b/EfficientNet/plots/B2/validation_accuracy.png differ diff --git a/EfficientNet/plots/B3/learning_rate.png b/EfficientNet/plots/B3/learning_rate.png new file mode 100644 index 0000000..87062ef Binary files /dev/null and b/EfficientNet/plots/B3/learning_rate.png differ diff --git a/EfficientNet/plots/B3/training_loss.png b/EfficientNet/plots/B3/training_loss.png new file mode 100644 index 0000000..6fe7eaa Binary files /dev/null and b/EfficientNet/plots/B3/training_loss.png differ diff --git a/EfficientNet/plots/B3/training_summary.png b/EfficientNet/plots/B3/training_summary.png new file mode 100644 index 0000000..1c616bf Binary files /dev/null and b/EfficientNet/plots/B3/training_summary.png differ diff --git a/EfficientNet/plots/B3/validation_accuracy.png b/EfficientNet/plots/B3/validation_accuracy.png new file mode 100644 index 0000000..91892e5 Binary files /dev/null and b/EfficientNet/plots/B3/validation_accuracy.png differ diff --git a/EfficientNet/plots/B4/learning_rate.png b/EfficientNet/plots/B4/learning_rate.png new file mode 100644 index 0000000..87062ef Binary files /dev/null and b/EfficientNet/plots/B4/learning_rate.png differ diff --git a/EfficientNet/plots/B4/training_loss.png b/EfficientNet/plots/B4/training_loss.png new file mode 100644 index 0000000..e8889f3 Binary files /dev/null and b/EfficientNet/plots/B4/training_loss.png differ diff --git a/EfficientNet/plots/B4/training_summary.png b/EfficientNet/plots/B4/training_summary.png new file mode 100644 index 0000000..4a35d74 Binary files /dev/null and b/EfficientNet/plots/B4/training_summary.png differ diff --git a/EfficientNet/plots/B4/validation_accuracy.png b/EfficientNet/plots/B4/validation_accuracy.png new file mode 100644 index 0000000..ce9724c Binary files /dev/null and b/EfficientNet/plots/B4/validation_accuracy.png differ diff --git a/EfficientNet/plots/B5/learning_rate.png b/EfficientNet/plots/B5/learning_rate.png new file mode 100644 index 0000000..87062ef Binary files /dev/null and b/EfficientNet/plots/B5/learning_rate.png differ diff --git a/EfficientNet/plots/B5/training_loss.png b/EfficientNet/plots/B5/training_loss.png new file mode 100644 index 0000000..8e69551 Binary files /dev/null and b/EfficientNet/plots/B5/training_loss.png differ diff --git a/EfficientNet/plots/B5/training_summary.png b/EfficientNet/plots/B5/training_summary.png new file mode 100644 index 0000000..85d3094 Binary files /dev/null and b/EfficientNet/plots/B5/training_summary.png differ diff --git a/EfficientNet/plots/B5/validation_accuracy.png b/EfficientNet/plots/B5/validation_accuracy.png new file mode 100644 index 0000000..2e3a3a2 Binary files /dev/null and b/EfficientNet/plots/B5/validation_accuracy.png differ diff --git a/EfficientNet/predict.py b/EfficientNet/predict.py new file mode 100644 index 0000000..775b3f4 --- /dev/null +++ b/EfficientNet/predict.py @@ -0,0 +1,74 @@ +import json +import os + +import matplotlib.pyplot as plt +import torch +from model import efficientnet_b0 as create_model +from PIL import Image +from torchvision import transforms + + +def main(): + device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") + + img_size = { + "B0": 224, + "B1": 240, + "B2": 260, + "B3": 300, + "B4": 380, + "B5": 456, + "B6": 528, + "B7": 600, + } + num_model = "B0" + + data_transform = transforms.Compose( + [ + transforms.Resize(img_size[num_model]), + transforms.CenterCrop(img_size[num_model]), + transforms.ToTensor(), + transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), + ] + ) + + img_path = "../dress.png" + assert os.path.exists(img_path), "file: '{}' dose not exist.".format(img_path) + img = Image.open(img_path) + img = img.convert("RGB") + plt.imshow(img) + + img = data_transform(img) + + img = torch.unsqueeze(img, dim=0) + + 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) + + model = create_model(num_classes=5).to(device) + model_weight_path = "./weights/model-1.pth" + model.load_state_dict(torch.load(model_weight_path, map_location=device)) + model.eval() + with torch.no_grad(): + 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() diff --git a/EfficientNet/train.py b/EfficientNet/train.py new file mode 100644 index 0000000..32c9f9b --- /dev/null +++ b/EfficientNet/train.py @@ -0,0 +1,359 @@ +import argparse +import math +import os + +import matplotlib.pyplot as plt +import numpy as np +import torch +import torch.optim as optim +import torch.optim.lr_scheduler as lr_scheduler + +# Import all EfficientNet models instead of just B0 +from model import ( + efficientnet_b0, + efficientnet_b1, + efficientnet_b2, + efficientnet_b3, + efficientnet_b4, + efficientnet_b5, + efficientnet_b6, + efficientnet_b7, +) +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 + +# Dictionary mapping model names to their respective functions +model_functions = { + "B0": efficientnet_b0, + "B1": efficientnet_b1, + "B2": efficientnet_b2, + "B3": efficientnet_b3, + "B4": efficientnet_b4, + "B5": efficientnet_b5, + "B6": efficientnet_b6, + "B7": efficientnet_b7, +} + + +def plot_training_process( + train_losses, val_accuracies, learning_rates, save_dir="./plots" +): + """ + Plot training metrics and save the figures + """ + if not os.path.exists(save_dir): + os.makedirs(save_dir) + + # Plot training loss + plt.figure(figsize=(10, 5)) + plt.plot(train_losses, label="Training Loss") + plt.xlabel("Epochs") + plt.ylabel("Loss") + plt.title("Training Loss over Epochs") + plt.legend() + plt.grid(True) + plt.savefig(f"{save_dir}/training_loss.png") + plt.close() + + # Plot validation accuracy + plt.figure(figsize=(10, 5)) + plt.plot(val_accuracies, label="Validation Accuracy") + plt.xlabel("Epochs") + plt.ylabel("Accuracy") + plt.title("Validation Accuracy over Epochs") + plt.legend() + plt.grid(True) + plt.savefig(f"{save_dir}/validation_accuracy.png") + plt.close() + + # Plot learning rate + plt.figure(figsize=(10, 5)) + plt.plot(learning_rates, label="Learning Rate") + plt.xlabel("Epochs") + plt.ylabel("Learning Rate") + plt.title("Learning Rate over Epochs") + plt.legend() + plt.grid(True) + plt.savefig(f"{save_dir}/learning_rate.png") + plt.close() + + # Combined plot + plt.figure(figsize=(12, 8)) + + ax1 = plt.subplot(3, 1, 1) + ax1.plot(train_losses, "b-", label="Training Loss") + ax1.set_ylabel("Loss") + ax1.legend(loc="upper right") + ax1.grid(True) + + ax2 = plt.subplot(3, 1, 2) + ax2.plot(val_accuracies, "r-", label="Validation Accuracy") + ax2.set_ylabel("Accuracy") + ax2.legend(loc="lower right") + ax2.grid(True) + + ax3 = plt.subplot(3, 1, 3) + ax3.plot(learning_rates, "g-", label="Learning Rate") + ax3.set_xlabel("Epochs") + ax3.set_ylabel("Learning Rate") + ax3.legend(loc="upper right") + ax3.grid(True) + + plt.tight_layout() + plt.savefig(f"{save_dir}/training_summary.png") + plt.close() + + +def train_model( + model_name, + args, + train_images_path, + train_images_label, + val_images_path, + val_images_label, +): + """ + Function to train a specific EfficientNet model variant + """ + print(f"\n{'='*50}") + print(f"Starting training for EfficientNet-{model_name}") + print(f"{'='*50}\n") + + device = torch.device(args.device if torch.cuda.is_available() else "cpu") + + # Create model-specific directory for weights and plots + model_weights_dir = f"./weights/{model_name}" + model_plots_dir = f"./plots/{model_name}" + if not os.path.exists(model_weights_dir): + os.makedirs(model_weights_dir) + if not os.path.exists(model_plots_dir): + os.makedirs(model_plots_dir) + + # Create a new SummaryWriter for this model + tb_writer = SummaryWriter(log_dir=f"runs/efficientnet_{model_name.lower()}") + + img_size = { + "B0": 224, + "B1": 240, + "B2": 260, + "B3": 300, + "B4": 380, + "B5": 456, + "B6": 528, + "B7": 600, + } + + data_transform = { + "train": transforms.Compose( + [ + transforms.RandomResizedCrop(img_size[model_name]), + transforms.RandomHorizontalFlip(), + transforms.ToTensor(), + transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), + ] + ), + "val": transforms.Compose( + [ + transforms.Resize(img_size[model_name]), + transforms.CenterCrop(img_size[model_name]), + transforms.ToTensor(), + transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), + ] + ), + } + + # Create datasets with model-specific transforms + 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]) + print(f"Using {nw} dataloader workers every process") + + 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, + ) + + # Create model using the appropriate function from the dictionary + create_model = model_functions[model_name] + model = create_model(num_classes=args.num_classes).to(device) + + # Model-specific weights path + weights_path = f"./efficientnet{model_name.lower()}.pth" + + # Load pretrained weights if available + if args.weights != "" and model_name == "B0": + # Only use the specified weights for B0 + weights_path = args.weights + + if os.path.exists(weights_path): + print(f"Loading pretrained weights from: {weights_path}") + weights_dict = torch.load(weights_path, map_location=device) + load_weights_dict = { + k: v + for k, v in weights_dict.items() + if model.state_dict()[k].numel() == v.numel() + } + print(model.load_state_dict(load_weights_dict, strict=False)) + else: + print(f"No pretrained weights found at: {weights_path}, starting from scratch") + + # Handle layer freezing + if args.freeze_layers: + for name, para in model.named_parameters(): + if ("features.top" not in name) and ("classifier" not in name): + para.requires_grad_(False) + else: + print(f"Training {name}") + + 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) + 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) + + best_acc = 0.0 + train_losses = [] + val_accuracies = [] + learning_rates = [] + + for epoch in range(args.epochs): + mean_loss = train_one_epoch( + model=model, + optimizer=optimizer, + data_loader=train_loader, + device=device, + epoch=epoch, + ) + scheduler.step() + + acc = evaluate(model=model, data_loader=val_loader, device=device) + print(f"[epoch {epoch}] accuracy: {round(acc, 3)}") + + train_losses.append(mean_loss) + val_accuracies.append(acc) + learning_rates.append(optimizer.param_groups[0]["lr"]) + + # Log metrics + 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], optimizer.param_groups[0]["lr"], epoch) + + # Save model weights periodically + if (epoch + 1) % 5 == 0: + torch.save(model.state_dict(), f"{model_weights_dir}/model-{epoch}.pth") + + # Save best model + if acc > best_acc: + best_acc = acc + torch.save(model.state_dict(), f"{model_weights_dir}/best_model.pth") + + # Plot training metrics for this model + plot_training_process( + train_losses, val_accuracies, learning_rates, save_dir=model_plots_dir + ) + + print(f"Best accuracy for EfficientNet-{model_name}: {best_acc:.4f}") + print(f"Training visualizations saved to {model_plots_dir}") + + # Close the tensorboard writer + tb_writer.close() + + return best_acc + + +def main(args): + print(args) + print( + 'Start Tensorboard with "tensorboard --logdir=runs", view at http://localhost:6006/' + ) + + if os.path.exists("./weights") is False: + os.makedirs("./weights") + + # Read dataset only once + train_images_path, train_images_label, val_images_path, val_images_label = ( + read_split_data(args.data_path) + ) + + # List of all EfficientNet models to train sequentially + model_variants = ["B0", "B1", "B2", "B3", "B4", "B5", "B6", "B7"] + + # Dictionary to store best accuracy for each model + best_accuracies = {} + + # Train each model sequentially + for model_name in model_variants: + best_acc = train_model( + model_name=model_name, + args=args, + train_images_path=train_images_path, + train_images_label=train_images_label, + val_images_path=val_images_path, + val_images_label=val_images_label, + ) + best_accuracies[model_name] = best_acc + + # Print summary of all models + print("\n" + "=" * 50) + print("Training complete for all EfficientNet models!") + print("Best accuracies for each model:") + for model_name, acc in best_accuracies.items(): + print(f"EfficientNet-{model_name}: {acc:.4f}") + print("=" * 50) + + +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=16) + parser.add_argument("--lr", type=float, default=0.01) + parser.add_argument("--lrf", type=float, default=0.01) + + # 数据集所在根目录 + # http://download.tensorflow.org/example_images/flower_photos.tgz + parser.add_argument("--data-path", type=str, default=r"..\\data\\flower_photos") + + parser.add_argument( + "--weights", + type=str, + default="./efficientnetb0.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) diff --git a/EfficientNet/trans_weights_to_pytorch.py b/EfficientNet/trans_weights_to_pytorch.py new file mode 100644 index 0000000..36abcab --- /dev/null +++ b/EfficientNet/trans_weights_to_pytorch.py @@ -0,0 +1,114 @@ +import numpy as np +import tensorflow as tf +import torch + +assert tf.version.VERSION >= "2.4.0", "version of tf must greater/equal than 2.4.0" + + +def main(): + # save pytorch weights path + save_path = "./efficientnetb0.pth" + + # create keras model and download weights + # EfficientNetB0, EfficientNetB1, EfficientNetB2, ... + m = tf.keras.applications.EfficientNetB0() + + weights_dict = dict() + weights = m.weights[3:] # delete norm weights + for weight in weights: + name = weight.name + data = weight.numpy() + + if "stem_conv/kernel:0" == name: + torch_name = "features.stem_conv.0.weight" + weights_dict[torch_name] = np.transpose(data, (3, 2, 0, 1)).astype( + np.float32 + ) + elif "stem_bn/gamma:0" == name: + torch_name = "features.stem_conv.1.weight" + weights_dict[torch_name] = data + elif "stem_bn/beta:0" == name: + torch_name = "features.stem_conv.1.bias" + weights_dict[torch_name] = data + elif "stem_bn/moving_mean:0" == name: + torch_name = "features.stem_conv.1.running_mean" + weights_dict[torch_name] = data + elif "stem_bn/moving_variance:0" == name: + torch_name = "features.stem_conv.1.running_var" + weights_dict[torch_name] = data + elif "block" in name: + name = name[5:] # delete "block" word + block_index = name[:2] # 1a, 2a, ... + name = name[3:] # delete block_index and "_" + torch_prefix = "features.{}.block.".format(block_index) + + trans_dict = { + "expand_conv/kernel:0": "expand_conv.0.weight", + "expand_bn/gamma:0": "expand_conv.1.weight", + "expand_bn/beta:0": "expand_conv.1.bias", + "expand_bn/moving_mean:0": "expand_conv.1.running_mean", + "expand_bn/moving_variance:0": "expand_conv.1.running_var", + "dwconv/depthwise_kernel:0": "dwconv.0.weight", + "bn/gamma:0": "dwconv.1.weight", + "bn/beta:0": "dwconv.1.bias", + "bn/moving_mean:0": "dwconv.1.running_mean", + "bn/moving_variance:0": "dwconv.1.running_var", + "se_reduce/kernel:0": "se.fc1.weight", + "se_reduce/bias:0": "se.fc1.bias", + "se_expand/kernel:0": "se.fc2.weight", + "se_expand/bias:0": "se.fc2.bias", + "project_conv/kernel:0": "project_conv.0.weight", + "project_bn/gamma:0": "project_conv.1.weight", + "project_bn/beta:0": "project_conv.1.bias", + "project_bn/moving_mean:0": "project_conv.1.running_mean", + "project_bn/moving_variance:0": "project_conv.1.running_var", + } + + assert name in trans_dict, "key '{}' not in trans_dict".format(name) + torch_postfix = trans_dict[name] + torch_name = torch_prefix + torch_postfix + if torch_postfix in [ + "expand_conv.0.weight", + "se.fc1.weight", + "se.fc2.weight", + "project_conv.0.weight", + ]: + data = np.transpose(data, (3, 2, 0, 1)).astype(np.float32) + elif torch_postfix == "dwconv.0.weight": + data = np.transpose(data, (2, 3, 0, 1)).astype(np.float32) + weights_dict[torch_name] = data + elif "top_conv/kernel:0" == name: + torch_name = "features.top.0.weight" + weights_dict[torch_name] = np.transpose(data, (3, 2, 0, 1)).astype( + np.float32 + ) + elif "top_bn/gamma:0" == name: + torch_name = "features.top.1.weight" + weights_dict[torch_name] = data + elif "top_bn/beta:0" == name: + torch_name = "features.top.1.bias" + weights_dict[torch_name] = data + elif "top_bn/moving_mean:0" == name: + torch_name = "features.top.1.running_mean" + weights_dict[torch_name] = data + elif "top_bn/moving_variance:0" == name: + torch_name = "features.top.1.running_var" + weights_dict[torch_name] = data + elif "predictions/kernel:0" == name: + torch_name = "classifier.1.weight" + weights_dict[torch_name] = np.transpose(data, (1, 0)).astype(np.float32) + elif "predictions/bias:0" == name: + torch_name = "classifier.1.bias" + weights_dict[torch_name] = data + else: + raise KeyError("no match key '{}'".format(name)) + + for k, v in weights_dict.items(): + weights_dict[k] = torch.as_tensor(v) + + torch.save(weights_dict, save_path) + print("Conversion complete.") + + +if __name__ == "__main__": + main() diff --git a/EfficientNet/utils.py b/EfficientNet/utils.py new file mode 100644 index 0000000..974acff --- /dev/null +++ b/EfficientNet/utils.py @@ -0,0 +1,149 @@ +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) + 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") + 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") + plt.xlabel("image class") + 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): + img = images[i].numpy().transpose(1, 2, 0) + 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([]) + plt.yticks([]) + 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, file=sys.stdout) + + 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, file=sys.stdout) + 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 diff --git a/GoogLeNet/class_indices.json b/GoogLeNet/class_indices.json new file mode 100644 index 0000000..84f8a53 --- /dev/null +++ b/GoogLeNet/class_indices.json @@ -0,0 +1,7 @@ +{ + "0": "daisy", + "1": "dandelion", + "2": "roses", + "3": "sunflowers", + "4": "tulips" +} diff --git a/GoogLeNet/evaluate.py b/GoogLeNet/evaluate.py new file mode 100644 index 0000000..0c4d37d --- /dev/null +++ b/GoogLeNet/evaluate.py @@ -0,0 +1,223 @@ +import json +import os + +import matplotlib.pyplot as plt +import numpy as np +import seaborn as sns +import torch +from model import GoogLeNet +from sklearn.metrics import ( + accuracy_score, + confusion_matrix, + f1_score, + precision_score, + recall_score, +) +from torchvision import datasets, transforms +from tqdm import tqdm + + +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((224, 224)), + transforms.ToTensor(), + transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), + ] + ) + + # Load validation dataset + data_root = os.path.abspath(os.path.join(os.getcwd(), "..")) + image_path = os.path.join(data_root, "data", "flower_data") + assert os.path.exists(image_path), f"{image_path} path does not exist." + + validate_dataset = datasets.ImageFolder( + root=os.path.join(image_path, "val"), transform=data_transform + ) + val_num = len(validate_dataset) + + batch_size = 4 + nw = min([os.cpu_count(), batch_size if batch_size > 1 else 0, 4]) + validate_loader = torch.utils.data.DataLoader( + validate_dataset, batch_size=batch_size, shuffle=False, num_workers=nw + ) + + # 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 = GoogLeNet(num_classes=5, aux_logits=False).to(device) + + # Load model weights + weights_path = "./googleNet.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), strict=False) + + # Evaluation + model.eval() + + # Lists to store predictions and ground truth + all_preds = [] + all_labels = [] + + with torch.no_grad(): + for val_data in tqdm(validate_loader, desc="Evaluating"): + val_images, val_labels = val_data + outputs = model(val_images.to(device)) + + # Get predictions + predict_y = torch.max(outputs, dim=1)[1] + + # Store predictions and labels + all_preds.extend(predict_y.cpu().numpy()) + all_labels.extend(val_labels.numpy()) + + # 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) + + # Visualize results + visualize_results( + cm, + class_indict, + accuracy, + precision, + recall, + f1, + class_precision, + class_recall, + class_f1, + ) + + +def visualize_results( + cm, + class_indict, + accuracy, + precision, + recall, + f1, + class_precision, + class_recall, + class_f1, +): + """ + Visualize evaluation metrics and confusion matrix + """ + plt.figure(figsize=(16, 12)) + + # Plot confusion matrix + plt.subplot(2, 2, 1) + 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") + + # Plot overall metrics + plt.subplot(2, 2, 2) + metrics = ["Accuracy", "Precision", "Recall", "F1 Score"] + values = [accuracy, precision, recall, f1] + plt.bar(metrics, values, color=["blue", "green", "orange", "red"]) + plt.ylim(0, 1.0) + plt.title("Overall Metrics") + for i, v in enumerate(values): + plt.text(i, v + 0.02, f"{v:.4f}", ha="center") + + # Plot per-class precision, recall, and F1 score + plt.subplot(2, 2, 3) + x = np.arange(len(class_indict)) + width = 0.25 + + plt.bar(x - width, class_precision, width, label="Precision", color="green") + plt.bar(x, class_recall, width, label="Recall", color="orange") + plt.bar(x + width, class_f1, width, label="F1 Score", color="red") + + plt.xlabel("Class") + plt.ylabel("Score") + plt.title("Per-class Metrics") + plt.xticks(x, [class_indict[str(i)] for i in range(len(class_indict))]) + plt.legend() + plt.ylim(0, 1.0) + + # Plot ROC curve (simplified version for multiclass) + plt.subplot(2, 2, 4) + + # Create a radar chart for per-class F1 scores + 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 F1 scores + values = class_f1.tolist() + values += values[:1] # Close the loop + + # Draw the plot + ax = plt.subplot(2, 2, 4, polar=True) + 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="grey", + size=7, + ) + plt.ylim(0, 1) + + plt.plot(angles, values, linewidth=1, linestyle="solid") + plt.fill(angles, values, "b", alpha=0.1) + plt.title("F1 Score by Class") + + plt.tight_layout() + plt.savefig("evaluation_metrics.png") + plt.show() + + +if __name__ == "__main__": + main() diff --git a/GoogLeNet/model.py b/GoogLeNet/model.py new file mode 100644 index 0000000..53fe8ee --- /dev/null +++ b/GoogLeNet/model.py @@ -0,0 +1,179 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F +import torchvision.models + + +class GoogLeNet(nn.Module): + def __init__(self, num_classes=1000, aux_logits=True, init_weights=False): + super(GoogLeNet, self).__init__() + self.aux_logits = aux_logits + + self.conv1 = BasicConv2d(3, 64, kernel_size=7, stride=2, padding=3) + self.maxpool1 = nn.MaxPool2d(3, stride=2, ceil_mode=True) + + self.conv2 = BasicConv2d(64, 64, kernel_size=1) + self.conv3 = BasicConv2d(64, 192, kernel_size=3, padding=1) + self.maxpool2 = nn.MaxPool2d(3, stride=2, ceil_mode=True) + + self.inception3a = Inception(192, 64, 96, 128, 16, 32, 32) + self.inception3b = Inception(256, 128, 128, 192, 32, 96, 64) + self.maxpool3 = nn.MaxPool2d(3, stride=2, ceil_mode=True) + + self.inception4a = Inception(480, 192, 96, 208, 16, 48, 64) + self.inception4b = Inception(512, 160, 112, 224, 24, 64, 64) + self.inception4c = Inception(512, 128, 128, 256, 24, 64, 64) + self.inception4d = Inception(512, 112, 144, 288, 32, 64, 64) + self.inception4e = Inception(528, 256, 160, 320, 32, 128, 128) + self.maxpool4 = nn.MaxPool2d(3, stride=2, ceil_mode=True) + + self.inception5a = Inception(832, 256, 160, 320, 32, 128, 128) + self.inception5b = Inception(832, 384, 192, 384, 48, 128, 128) + + if self.aux_logits: + self.aux1 = InceptionAux(512, num_classes) + self.aux2 = InceptionAux(528, num_classes) + + self.avgpool = nn.AdaptiveAvgPool2d((1, 1)) + self.dropout = nn.Dropout(0.4) + self.fc = nn.Linear(1024, num_classes) + if init_weights: + self._initialize_weights() + + def forward(self, x): + # N x 3 x 224 x 224 + x = self.conv1(x) + # N x 64 x 112 x 112 + x = self.maxpool1(x) + # N x 64 x 56 x 56 + x = self.conv2(x) + # N x 64 x 56 x 56 + x = self.conv3(x) + # N x 192 x 56 x 56 + x = self.maxpool2(x) + + # N x 192 x 28 x 28 + x = self.inception3a(x) + # N x 256 x 28 x 28 + x = self.inception3b(x) + # N x 480 x 28 x 28 + x = self.maxpool3(x) + # N x 480 x 14 x 14 + x = self.inception4a(x) + # N x 512 x 14 x 14 + if self.training and self.aux_logits: # eval model lose this layer + aux1 = self.aux1(x) + + x = self.inception4b(x) + # N x 512 x 14 x 14 + x = self.inception4c(x) + # N x 512 x 14 x 14 + x = self.inception4d(x) + # N x 528 x 14 x 14 + if self.training and self.aux_logits: # eval model lose this layer + aux2 = self.aux2(x) + + x = self.inception4e(x) + # N x 832 x 14 x 14 + x = self.maxpool4(x) + # N x 832 x 7 x 7 + x = self.inception5a(x) + # N x 832 x 7 x 7 + x = self.inception5b(x) + # N x 1024 x 7 x 7 + + x = self.avgpool(x) + # N x 1024 x 1 x 1 + x = torch.flatten(x, 1) + # N x 1024 + x = self.dropout(x) + x = self.fc(x) + # N x 1000 (num_classes) + if self.training and self.aux_logits: # eval model lose this layer + return x, aux2, aux1 + return x + + def _initialize_weights(self): + for m in self.modules(): + if isinstance(m, nn.Conv2d): + nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu") + if m.bias is not None: + nn.init.constant_(m.bias, 0) + elif isinstance(m, nn.Linear): + nn.init.normal_(m.weight, 0, 0.01) + nn.init.constant_(m.bias, 0) + + +class Inception(nn.Module): + def __init__(self, in_channels, ch1x1, ch3x3red, ch3x3, ch5x5red, ch5x5, pool_proj): + super(Inception, self).__init__() + + self.branch1 = BasicConv2d(in_channels, ch1x1, kernel_size=1) + + self.branch2 = nn.Sequential( + BasicConv2d(in_channels, ch3x3red, kernel_size=1), + BasicConv2d( + ch3x3red, ch3x3, kernel_size=3, padding=1 + ), # 保证输出大小等于输入大小 + ) + + self.branch3 = nn.Sequential( + BasicConv2d(in_channels, ch5x5red, kernel_size=1), + BasicConv2d( + ch5x5red, ch5x5, kernel_size=5, padding=2 + ), # 保证输出大小等于输入大小 + ) + + self.branch4 = nn.Sequential( + nn.MaxPool2d(kernel_size=3, stride=1, padding=1), + BasicConv2d(in_channels, pool_proj, kernel_size=1), + ) + + def forward(self, x): + branch1 = self.branch1(x) + branch2 = self.branch2(x) + branch3 = self.branch3(x) + branch4 = self.branch4(x) + + outputs = [branch1, branch2, branch3, branch4] + return torch.cat(outputs, 1) + + +class InceptionAux(nn.Module): + def __init__(self, in_channels, num_classes): + super(InceptionAux, self).__init__() + self.averagePool = nn.AvgPool2d(kernel_size=5, stride=3) + self.conv = BasicConv2d( + in_channels, 128, kernel_size=1 + ) # output[batch, 128, 4, 4] + + self.fc1 = nn.Linear(2048, 1024) + self.fc2 = nn.Linear(1024, num_classes) + + def forward(self, x): + # aux1: N x 512 x 14 x 14, aux2: N x 528 x 14 x 14 + x = self.averagePool(x) + # aux1: N x 512 x 4 x 4, aux2: N x 528 x 4 x 4 + x = self.conv(x) + # N x 128 x 4 x 4 + x = torch.flatten(x, 1) + x = F.dropout(x, 0.5, training=self.training) + # N x 2048 + x = F.relu(self.fc1(x), inplace=True) + x = F.dropout(x, 0.5, training=self.training) + # N x 1024 + x = self.fc2(x) + # N x num_classes + return x + + +class BasicConv2d(nn.Module): + def __init__(self, in_channels, out_channels, **kwargs): + super(BasicConv2d, self).__init__() + self.conv = nn.Conv2d(in_channels, out_channels, **kwargs) + self.relu = nn.ReLU(inplace=True) + + def forward(self, x): + x = self.conv(x) + x = self.relu(x) + return x diff --git a/GoogLeNet/predict.py b/GoogLeNet/predict.py new file mode 100644 index 0000000..e3923eb --- /dev/null +++ b/GoogLeNet/predict.py @@ -0,0 +1,72 @@ +import json +import os + +import matplotlib.pyplot as plt +import torch +from model import GoogLeNet +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((224, 224)), + transforms.ToTensor(), + transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), + ] + ) + + # load image + img_path = "./rose.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 = GoogLeNet(num_classes=5, aux_logits=False).to(device) + + # load model weights + weights_path = "./googleNet.pth" + assert os.path.exists(weights_path), "file: '{}' dose not exist.".format( + weights_path + ) + missing_keys, unexpected_keys = model.load_state_dict( + torch.load(weights_path, map_location=device), strict=False + ) + + 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() diff --git a/GoogLeNet/train.py b/GoogLeNet/train.py new file mode 100644 index 0000000..458ec9e --- /dev/null +++ b/GoogLeNet/train.py @@ -0,0 +1,183 @@ +import json +import os + +import matplotlib.pyplot as plt +import torch +import torch.nn as nn +import torch.optim as optim +from model import GoogLeNet +from torchvision import datasets, transforms +from tqdm import tqdm + + +def main(): + device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") + print("using {} device.".format(device)) + + data_transform = { + "train": transforms.Compose( + [ + transforms.RandomResizedCrop(224), + transforms.RandomHorizontalFlip(), + transforms.ToTensor(), + transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), + ] + ), + "val": transforms.Compose( + [ + transforms.Resize((224, 224)), + transforms.ToTensor(), + transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), + ] + ), + } + + data_root = os.path.abspath(os.path.join(os.getcwd(), "..")) # get data root path + image_path = os.path.join(data_root, "data", "flower_data") # flower data set path + assert os.path.exists(image_path), "{} path does not exist.".format(image_path) + train_dataset = datasets.ImageFolder( + root=os.path.join(image_path, "train"), transform=data_transform["train"] + ) + train_num = len(train_dataset) + + # {'daisy':0, 'dandelion':1, 'roses':2, 'sunflower':3, 'tulips':4} + flower_list = train_dataset.class_to_idx + cla_dict = dict((val, key) for key, val in flower_list.items()) + # write dict into json file + json_str = json.dumps(cla_dict, indent=4) + with open("class_indices.json", "w") as json_file: + json_file.write(json_str) + + batch_size = 4 # 如果cuda报超出显存可以改小一点 + nw = min( + [os.cpu_count(), batch_size if batch_size > 1 else 0, 4] + ) # 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, num_workers=nw + ) + + validate_dataset = datasets.ImageFolder( + root=os.path.join(image_path, "val"), transform=data_transform["val"] + ) + val_num = len(validate_dataset) + validate_loader = torch.utils.data.DataLoader( + validate_dataset, batch_size=batch_size, shuffle=False, num_workers=nw + ) + + print( + "using {} images for training, {} images for validation.".format( + train_num, val_num + ) + ) + + # test_data_iter = iter(validate_loader) + # test_image, test_label = test_data_iter.next() + + # net = torchvision.models.googlenet(num_classes=5) + # model_dict = net.state_dict() + # pretrain_model = torch.load("googlenet.pth") + # del_list = ["aux1.fc2.weight", "aux1.fc2.bias", + # "aux2.fc2.weight", "aux2.fc2.bias", + # "fc.weight", "fc.bias"] + # pretrain_dict = {k: v for k, v in pretrain_model.items() if k not in del_list} + # model_dict.update(pretrain_dict) + # net.load_state_dict(model_dict) + net = GoogLeNet(num_classes=5, aux_logits=True, init_weights=True) + + net.to(device) + loss_function = nn.CrossEntropyLoss() + optimizer = optim.Adam(net.parameters(), lr=0.0003) + + epochs = 30 + best_acc = 0.0 + # googlenet 官方权重下载: https://download.pytorch.org/models/googlenet-1378be20.pth + save_path = "./googleNet.pth" + train_steps = len(train_loader) + + # Lists to store metrics for visualization + train_losses = [] + val_accuracies = [] + + for epoch in range(epochs): + # train + net.train() + running_loss = 0.0 + train_bar = tqdm(train_loader) + for step, data in enumerate(train_bar): + images, labels = data + optimizer.zero_grad() + logits, aux_logits2, aux_logits1 = net(images.to(device)) + loss0 = loss_function(logits, labels.to(device)) + loss1 = loss_function(aux_logits1, labels.to(device)) + loss2 = loss_function(aux_logits2, labels.to(device)) + loss = loss0 + loss1 * 0.3 + loss2 * 0.3 + loss.backward() + optimizer.step() + + # print statistics + running_loss += loss.item() + + train_bar.desc = "train epoch[{}/{}] loss:{:.3f}".format( + epoch + 1, epochs, loss + ) + + # validate + net.eval() + acc = 0.0 # accumulate accurate number / epoch + with torch.no_grad(): + val_bar = tqdm(validate_loader) + for val_data in val_bar: + val_images, val_labels = val_data + outputs = net( + val_images.to(device) + ) # eval model only have last output layer + predict_y = torch.max(outputs, dim=1)[1] + acc += torch.eq(predict_y, val_labels.to(device)).sum().item() + + val_accurate = acc / val_num + epoch_loss = running_loss / train_steps + + # Store metrics for visualization + train_losses.append(epoch_loss) + val_accuracies.append(val_accurate) + + print( + "[epoch %d] train_loss: %.3f val_accuracy: %.3f" + % (epoch + 1, epoch_loss, val_accurate) + ) + + if val_accurate > best_acc: + best_acc = val_accurate + torch.save(net.state_dict(), save_path) + + print("Finished Training") + + # Visualize training process + plt.figure(figsize=(12, 5)) + + # Plot training loss + plt.subplot(1, 2, 1) + plt.plot(range(1, epochs + 1), train_losses, "bo-", label="Training Loss") + plt.title("Training Loss per Epoch") + plt.xlabel("Epoch") + plt.ylabel("Loss") + plt.grid(True) + plt.legend() + + # Plot validation accuracy + plt.subplot(1, 2, 2) + plt.plot(range(1, epochs + 1), val_accuracies, "ro-", label="Validation Accuracy") + plt.title("Validation Accuracy per Epoch") + plt.xlabel("Epoch") + plt.ylabel("Accuracy") + plt.grid(True) + plt.legend() + + plt.tight_layout() + plt.savefig("training_visualization.png") + plt.show() + + +if __name__ == "__main__": + main() diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..e62ec04 --- /dev/null +++ b/LICENSE @@ -0,0 +1,674 @@ +GNU GENERAL PUBLIC LICENSE + Version 3, 29 June 2007 + + Copyright (C) 2007 Free Software Foundation, Inc. + Everyone is permitted to copy and distribute verbatim copies + of this license document, but changing it is not allowed. + + Preamble + + The GNU General Public License is a free, copyleft license for +software and other kinds of works. + + The licenses for most software and other practical works are designed +to take away your freedom to share and change the works. By contrast, +the GNU General Public License is intended to guarantee your freedom to +share and change all versions of a program--to make sure it remains free +software for all its users. We, the Free Software Foundation, use the +GNU General Public License for most of our software; it applies also to +any other work released this way by its authors. You can apply it to +your programs, too. + + When we speak of free software, we are referring to freedom, not +price. Our General Public Licenses are designed to make sure that you +have the freedom to distribute copies of free software (and charge for +them if you wish), that you receive source code or can get it if you +want it, that you can change the software or use pieces of it in new +free programs, and that you know you can do these things. + + To protect your rights, we need to prevent others from denying you +these rights or asking you to surrender the rights. Therefore, you have +certain responsibilities if you distribute copies of the software, or if +you modify it: responsibilities to respect the freedom of others. + + For example, if you distribute copies of such a program, whether +gratis or for a fee, you must pass on to the recipients the same +freedoms that you received. You must make sure that they, too, receive +or can get the source code. And you must show them these terms so they +know their rights. + + Developers that use the GNU GPL protect your rights with two steps: +(1) assert copyright on the software, and (2) offer you this License +giving you legal permission to copy, distribute and/or modify it. + + For the developers' and authors' protection, the GPL clearly explains +that there is no warranty for this free software. For both users' and +authors' sake, the GPL requires that modified versions be marked as +changed, so that their problems will not be attributed erroneously to +authors of previous versions. + + Some devices are designed to deny users access to install or run +modified versions of the software inside them, although the manufacturer +can do so. This is fundamentally incompatible with the aim of +protecting users' freedom to change the software. The systematic +pattern of such abuse occurs in the area of products for individuals to +use, which is precisely where it is most unacceptable. Therefore, we +have designed this version of the GPL to prohibit the practice for those +products. If such problems arise substantially in other domains, we +stand ready to extend this provision to those domains in future versions +of the GPL, as needed to protect the freedom of users. + + Finally, every program is threatened constantly by software patents. +States should not allow patents to restrict development and use of +software on general-purpose computers, but in those that do, we wish to +avoid the special danger that patents applied to a free program could +make it effectively proprietary. To prevent this, the GPL assures that +patents cannot be used to render the program non-free. + + The precise terms and conditions for copying, distribution and +modification follow. + + TERMS AND CONDITIONS + + 0. Definitions. + + "This License" refers to version 3 of the GNU General Public License. + + "Copyright" also means copyright-like laws that apply to other kinds of +works, such as semiconductor masks. + + "The Program" refers to any copyrightable work licensed under this +License. Each licensee is addressed as "you". "Licensees" and +"recipients" may be individuals or organizations. + + To "modify" a work means to copy from or adapt all or part of the work +in a fashion requiring copyright permission, other than the making of an +exact copy. The resulting work is called a "modified version" of the +earlier work or a work "based on" the earlier work. + + A "covered work" means either the unmodified Program or a work based +on the Program. + + To "propagate" a work means to do anything with it that, without +permission, would make you directly or secondarily liable for +infringement under applicable copyright law, except executing it on a +computer or modifying a private copy. Propagation includes copying, +distribution (with or without modification), making available to the +public, and in some countries other activities as well. + + To "convey" a work means any kind of propagation that enables other +parties to make or receive copies. Mere interaction with a user through +a computer network, with no transfer of a copy, is not conveying. + + An interactive user interface displays "Appropriate Legal Notices" +to the extent that it includes a convenient and prominently visible +feature that (1) displays an appropriate copyright notice, and (2) +tells the user that there is no warranty for the work (except to the +extent that warranties are provided), that licensees may convey the +work under this License, and how to view a copy of this License. If +the interface presents a list of user commands or options, such as a +menu, a prominent item in the list meets this criterion. + + 1. Source Code. + + The "source code" for a work means the preferred form of the work +for making modifications to it. "Object code" means any non-source +form of a work. + + A "Standard Interface" means an interface that either is an official +standard defined by a recognized standards body, or, in the case of +interfaces specified for a particular programming language, one that +is widely used among developers working in that language. + + The "System Libraries" of an executable work include anything, other +than the work as a whole, that (a) is included in the normal form of +packaging a Major Component, but which is not part of that Major +Component, and (b) serves only to enable use of the work with that +Major Component, or to implement a Standard Interface for which an +implementation is available to the public in source code form. A +"Major Component", in this context, means a major essential component +(kernel, window system, and so on) of the specific operating system +(if any) on which the executable work runs, or a compiler used to +produce the work, or an object code interpreter used to run it. + + The "Corresponding Source" for a work in object code form means all +the source code needed to generate, install, and (for an executable +work) run the object code and to modify the work, including scripts to +control those activities. However, it does not include the work's +System Libraries, or general-purpose tools or generally available free +programs which are used unmodified in performing those activities but +which are not part of the work. For example, Corresponding Source +includes interface definition files associated with source files for +the work, and the source code for shared libraries and dynamically +linked subprograms that the work is specifically designed to require, +such as by intimate data communication or control flow between those +subprograms and other parts of the work. + + The Corresponding Source need not include anything that users +can regenerate automatically from other parts of the Corresponding +Source. + + The Corresponding Source for a work in source code form is that +same work. + + 2. Basic Permissions. + + All rights granted under this License are granted for the term of +copyright on the Program, and are irrevocable provided the stated +conditions are met. This License explicitly affirms your unlimited +permission to run the unmodified Program. The output from running a +covered work is covered by this License only if the output, given its +content, constitutes a covered work. This License acknowledges your +rights of fair use or other equivalent, as provided by copyright law. + + You may make, run and propagate covered works that you do not +convey, without conditions so long as your license otherwise remains +in force. You may convey covered works to others for the sole purpose +of having them make modifications exclusively for you, or provide you +with facilities for running those works, provided that you comply with +the terms of this License in conveying all material for which you do +not control copyright. Those thus making or running the covered works +for you must do so exclusively on your behalf, under your direction +and control, on terms that prohibit them from making any copies of +your copyrighted material outside their relationship with you. + + Conveying under any other circumstances is permitted solely under +the conditions stated below. Sublicensing is not allowed; section 10 +makes it unnecessary. + + 3. Protecting Users' Legal Rights From Anti-Circumvention Law. + + No covered work shall be deemed part of an effective technological +measure under any applicable law fulfilling obligations under article +11 of the WIPO copyright treaty adopted on 20 December 1996, or +similar laws prohibiting or restricting circumvention of such +measures. + + When you convey a covered work, you waive any legal power to forbid +circumvention of technological measures to the extent such circumvention +is effected by exercising rights under this License with respect to +the covered work, and you disclaim any intention to limit operation or +modification of the work as a means of enforcing, against the work's +users, your or third parties' legal rights to forbid circumvention of +technological measures. + + 4. Conveying Verbatim Copies. + + You may convey verbatim copies of the Program's source code as you +receive it, in any medium, provided that you conspicuously and +appropriately publish on each copy an appropriate copyright notice; +keep intact all notices stating that this License and any +non-permissive terms added in accord with section 7 apply to the code; +keep intact all notices of the absence of any warranty; and give all +recipients a copy of this License along with the Program. + + You may charge any price or no price for each copy that you convey, +and you may offer support or warranty protection for a fee. + + 5. Conveying Modified Source Versions. + + You may convey a work based on the Program, or the modifications to +produce it from the Program, in the form of source code under the +terms of section 4, provided that you also meet all of these conditions: + + a) The work must carry prominent notices stating that you modified + it, and giving a relevant date. + + b) The work must carry prominent notices stating that it is + released under this License and any conditions added under section + 7. This requirement modifies the requirement in section 4 to + "keep intact all notices". + + c) You must license the entire work, as a whole, under this + License to anyone who comes into possession of a copy. This + License will therefore apply, along with any applicable section 7 + additional terms, to the whole of the work, and all its parts, + regardless of how they are packaged. This License gives no + permission to license the work in any other way, but it does not + invalidate such permission if you have separately received it. + + d) If the work has interactive user interfaces, each must display + Appropriate Legal Notices; however, if the Program has interactive + interfaces that do not display Appropriate Legal Notices, your + work need not make them do so. + + A compilation of a covered work with other separate and independent +works, which are not by their nature extensions of the covered work, +and which are not combined with it such as to form a larger program, +in or on a volume of a storage or distribution medium, is called an +"aggregate" if the compilation and its resulting copyright are not +used to limit the access or legal rights of the compilation's users +beyond what the individual works permit. Inclusion of a covered work +in an aggregate does not cause this License to apply to the other +parts of the aggregate. + + 6. Conveying Non-Source Forms. + + You may convey a covered work in object code form under the terms +of sections 4 and 5, provided that you also convey the +machine-readable Corresponding Source under the terms of this License, +in one of these ways: + + a) Convey the object code in, or embodied in, a physical product + (including a physical distribution medium), accompanied by the + Corresponding Source fixed on a durable physical medium + customarily used for software interchange. + + b) Convey the object code in, or embodied in, a physical product + (including a physical distribution medium), accompanied by a + written offer, valid for at least three years and valid for as + long as you offer spare parts or customer support for that product + model, to give anyone who possesses the object code either (1) a + copy of the Corresponding Source for all the software in the + product that is covered by this License, on a durable physical + medium customarily used for software interchange, for a price no + more than your reasonable cost of physically performing this + conveying of source, or (2) access to copy the + Corresponding Source from a network server at no charge. + + c) Convey individual copies of the object code with a copy of the + written offer to provide the Corresponding Source. This + alternative is allowed only occasionally and noncommercially, and + only if you received the object code with such an offer, in accord + with subsection 6b. + + d) Convey the object code by offering access from a designated + place (gratis or for a charge), and offer equivalent access to the + Corresponding Source in the same way through the same place at no + further charge. You need not require recipients to copy the + Corresponding Source along with the object code. If the place to + copy the object code is a network server, the Corresponding Source + may be on a different server (operated by you or a third party) + that supports equivalent copying facilities, provided you maintain + clear directions next to the object code saying where to find the + Corresponding Source. Regardless of what server hosts the + Corresponding Source, you remain obligated to ensure that it is + available for as long as needed to satisfy these requirements. + + e) Convey the object code using peer-to-peer transmission, provided + you inform other peers where the object code and Corresponding + Source of the work are being offered to the general public at no + charge under subsection 6d. + + A separable portion of the object code, whose source code is excluded +from the Corresponding Source as a System Library, need not be +included in conveying the object code work. + + A "User Product" is either (1) a "consumer product", which means any +tangible personal property which is normally used for personal, family, +or household purposes, or (2) anything designed or sold for incorporation +into a dwelling. In determining whether a product is a consumer product, +doubtful cases shall be resolved in favor of coverage. For a particular +product received by a particular user, "normally used" refers to a +typical or common use of that class of product, regardless of the status +of the particular user or of the way in which the particular user +actually uses, or expects or is expected to use, the product. A product +is a consumer product regardless of whether the product has substantial +commercial, industrial or non-consumer uses, unless such uses represent +the only significant mode of use of the product. + + "Installation Information" for a User Product means any methods, +procedures, authorization keys, or other information required to install +and execute modified versions of a covered work in that User Product from +a modified version of its Corresponding Source. The information must +suffice to ensure that the continued functioning of the modified object +code is in no case prevented or interfered with solely because +modification has been made. + + If you convey an object code work under this section in, or with, or +specifically for use in, a User Product, and the conveying occurs as +part of a transaction in which the right of possession and use of the +User Product is transferred to the recipient in perpetuity or for a +fixed term (regardless of how the transaction is characterized), the +Corresponding Source conveyed under this section must be accompanied +by the Installation Information. But this requirement does not apply +if neither you nor any third party retains the ability to install +modified object code on the User Product (for example, the work has +been installed in ROM). + + The requirement to provide Installation Information does not include a +requirement to continue to provide support service, warranty, or updates +for a work that has been modified or installed by the recipient, or for +the User Product in which it has been modified or installed. Access to a +network may be denied when the modification itself materially and +adversely affects the operation of the network or violates the rules and +protocols for communication across the network. + + Corresponding Source conveyed, and Installation Information provided, +in accord with this section must be in a format that is publicly +documented (and with an implementation available to the public in +source code form), and must require no special password or key for +unpacking, reading or copying. + + 7. Additional Terms. + + "Additional permissions" are terms that supplement the terms of this +License by making exceptions from one or more of its conditions. +Additional permissions that are applicable to the entire Program shall +be treated as though they were included in this License, to the extent +that they are valid under applicable law. If additional permissions +apply only to part of the Program, that part may be used separately +under those permissions, but the entire Program remains governed by +this License without regard to the additional permissions. + + When you convey a copy of a covered work, you may at your option +remove any additional permissions from that copy, or from any part of +it. (Additional permissions may be written to require their own +removal in certain cases when you modify the work.) You may place +additional permissions on material, added by you to a covered work, +for which you have or can give appropriate copyright permission. + + Notwithstanding any other provision of this License, for material you +add to a covered work, you may (if authorized by the copyright holders of +that material) supplement the terms of this License with terms: + + a) Disclaiming warranty or limiting liability differently from the + terms of sections 15 and 16 of this License; or + + b) Requiring preservation of specified reasonable legal notices or + author attributions in that material or in the Appropriate Legal + Notices displayed by works containing it; or + + c) Prohibiting misrepresentation of the origin of that material, or + requiring that modified versions of such material be marked in + reasonable ways as different from the original version; or + + d) Limiting the use for publicity purposes of names of licensors or + authors of the material; or + + e) Declining to grant rights under trademark law for use of some + trade names, trademarks, or service marks; or + + f) Requiring indemnification of licensors and authors of that + material by anyone who conveys the material (or modified versions of + it) with contractual assumptions of liability to the recipient, for + any liability that these contractual assumptions directly impose on + those licensors and authors. + + All other non-permissive additional terms are considered "further +restrictions" within the meaning of section 10. If the Program as you +received it, or any part of it, contains a notice stating that it is +governed by this License along with a term that is a further +restriction, you may remove that term. If a license document contains +a further restriction but permits relicensing or conveying under this +License, you may add to a covered work material governed by the terms +of that license document, provided that the further restriction does +not survive such relicensing or conveying. + + If you add terms to a covered work in accord with this section, you +must place, in the relevant source files, a statement of the +additional terms that apply to those files, or a notice indicating +where to find the applicable terms. + + Additional terms, permissive or non-permissive, may be stated in the +form of a separately written license, or stated as exceptions; +the above requirements apply either way. + + 8. Termination. + + You may not propagate or modify a covered work except as expressly +provided under this License. Any attempt otherwise to propagate or +modify it is void, and will automatically terminate your rights under +this License (including any patent licenses granted under the third +paragraph of section 11). + + However, if you cease all violation of this License, then your +license from a particular copyright holder is reinstated (a) +provisionally, unless and until the copyright holder explicitly and +finally terminates your license, and (b) permanently, if the copyright +holder fails to notify you of the violation by some reasonable means +prior to 60 days after the cessation. + + Moreover, your license from a particular copyright holder is +reinstated permanently if the copyright holder notifies you of the +violation by some reasonable means, this is the first time you have +received notice of violation of this License (for any work) from that +copyright holder, and you cure the violation prior to 30 days after +your receipt of the notice. + + Termination of your rights under this section does not terminate the +licenses of parties who have received copies or rights from you under +this License. If your rights have been terminated and not permanently +reinstated, you do not qualify to receive new licenses for the same +material under section 10. + + 9. Acceptance Not Required for Having Copies. + + You are not required to accept this License in order to receive or +run a copy of the Program. Ancillary propagation of a covered work +occurring solely as a consequence of using peer-to-peer transmission +to receive a copy likewise does not require acceptance. However, +nothing other than this License grants you permission to propagate or +modify any covered work. These actions infringe copyright if you do +not accept this License. Therefore, by modifying or propagating a +covered work, you indicate your acceptance of this License to do so. + + 10. Automatic Licensing of Downstream Recipients. + + Each time you convey a covered work, the recipient automatically +receives a license from the original licensors, to run, modify and +propagate that work, subject to this License. You are not responsible +for enforcing compliance by third parties with this License. + + An "entity transaction" is a transaction transferring control of an +organization, or substantially all assets of one, or subdividing an +organization, or merging organizations. If propagation of a covered +work results from an entity transaction, each party to that +transaction who receives a copy of the work also receives whatever +licenses to the work the party's predecessor in interest had or could +give under the previous paragraph, plus a right to possession of the +Corresponding Source of the work from the predecessor in interest, if +the predecessor has it or can get it with reasonable efforts. + + You may not impose any further restrictions on the exercise of the +rights granted or affirmed under this License. For example, you may +not impose a license fee, royalty, or other charge for exercise of +rights granted under this License, and you may not initiate litigation +(including a cross-claim or counterclaim in a lawsuit) alleging that +any patent claim is infringed by making, using, selling, offering for +sale, or importing the Program or any portion of it. + + 11. Patents. + + A "contributor" is a copyright holder who authorizes use under this +License of the Program or a work on which the Program is based. The +work thus licensed is called the contributor's "contributor version". + + A contributor's "essential patent claims" are all patent claims +owned or controlled by the contributor, whether already acquired or +hereafter acquired, that would be infringed by some manner, permitted +by this License, of making, using, or selling its contributor version, +but do not include claims that would be infringed only as a +consequence of further modification of the contributor version. For +purposes of this definition, "control" includes the right to grant +patent sublicenses in a manner consistent with the requirements of +this License. + + Each contributor grants you a non-exclusive, worldwide, royalty-free +patent license under the contributor's essential patent claims, to +make, use, sell, offer for sale, import and otherwise run, modify and +propagate the contents of its contributor version. + + In the following three paragraphs, a "patent license" is any express +agreement or commitment, however denominated, not to enforce a patent +(such as an express permission to practice a patent or covenant not to +sue for patent infringement). To "grant" such a patent license to a +party means to make such an agreement or commitment not to enforce a +patent against the party. + + If you convey a covered work, knowingly relying on a patent license, +and the Corresponding Source of the work is not available for anyone +to copy, free of charge and under the terms of this License, through a +publicly available network server or other readily accessible means, +then you must either (1) cause the Corresponding Source to be so +available, or (2) arrange to deprive yourself of the benefit of the +patent license for this particular work, or (3) arrange, in a manner +consistent with the requirements of this License, to extend the patent +license to downstream recipients. "Knowingly relying" means you have +actual knowledge that, but for the patent license, your conveying the +covered work in a country, or your recipient's use of the covered work +in a country, would infringe one or more identifiable patents in that +country that you have reason to believe are valid. + + If, pursuant to or in connection with a single transaction or +arrangement, you convey, or propagate by procuring conveyance of, a +covered work, and grant a patent license to some of the parties +receiving the covered work authorizing them to use, propagate, modify +or convey a specific copy of the covered work, then the patent license +you grant is automatically extended to all recipients of the covered +work and works based on it. + + A patent license is "discriminatory" if it does not include within +the scope of its coverage, prohibits the exercise of, or is +conditioned on the non-exercise of one or more of the rights that are +specifically granted under this License. You may not convey a covered +work if you are a party to an arrangement with a third party that is +in the business of distributing software, under which you make payment +to the third party based on the extent of your activity of conveying +the work, and under which the third party grants, to any of the +parties who would receive the covered work from you, a discriminatory +patent license (a) in connection with copies of the covered work +conveyed by you (or copies made from those copies), or (b) primarily +for and in connection with specific products or compilations that +contain the covered work, unless you entered into that arrangement, +or that patent license was granted, prior to 28 March 2007. + + Nothing in this License shall be construed as excluding or limiting +any implied license or other defenses to infringement that may +otherwise be available to you under applicable patent law. + + 12. No Surrender of Others' Freedom. + + If conditions are imposed on you (whether by court order, agreement or +otherwise) that contradict the conditions of this License, they do not +excuse you from the conditions of this License. If you cannot convey a +covered work so as to satisfy simultaneously your obligations under this +License and any other pertinent obligations, then as a consequence you may +not convey it at all. For example, if you agree to terms that obligate you +to collect a royalty for further conveying from those to whom you convey +the Program, the only way you could satisfy both those terms and this +License would be to refrain entirely from conveying the Program. + + 13. Use with the GNU Affero General Public License. + + Notwithstanding any other provision of this License, you have +permission to link or combine any covered work with a work licensed +under version 3 of the GNU Affero General Public License into a single +combined work, and to convey the resulting work. The terms of this +License will continue to apply to the part which is the covered work, +but the special requirements of the GNU Affero General Public License, +section 13, concerning interaction through a network will apply to the +combination as such. + + 14. Revised Versions of this License. + + The Free Software Foundation may publish revised and/or new versions of +the GNU General Public License from time to time. Such new versions will +be similar in spirit to the present version, but may differ in detail to +address new problems or concerns. + + Each version is given a distinguishing version number. If the +Program specifies that a certain numbered version of the GNU General +Public License "or any later version" applies to it, you have the +option of following the terms and conditions either of that numbered +version or of any later version published by the Free Software +Foundation. If the Program does not specify a version number of the +GNU General Public License, you may choose any version ever published +by the Free Software Foundation. + + If the Program specifies that a proxy can decide which future +versions of the GNU General Public License can be used, that proxy's +public statement of acceptance of a version permanently authorizes you +to choose that version for the Program. + + Later license versions may give you additional or different +permissions. However, no additional obligations are imposed on any +author or copyright holder as a result of your choosing to follow a +later version. + + 15. Disclaimer of Warranty. + + THERE IS NO WARRANTY FOR THE PROGRAM, TO THE EXTENT PERMITTED BY +APPLICABLE LAW. EXCEPT WHEN OTHERWISE STATED IN WRITING THE COPYRIGHT +HOLDERS AND/OR OTHER PARTIES PROVIDE THE PROGRAM "AS IS" WITHOUT WARRANTY +OF ANY KIND, EITHER EXPRESSED OR IMPLIED, INCLUDING, BUT NOT LIMITED TO, +THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR +PURPOSE. THE ENTIRE RISK AS TO THE QUALITY AND PERFORMANCE OF THE PROGRAM +IS WITH YOU. SHOULD THE PROGRAM PROVE DEFECTIVE, YOU ASSUME THE COST OF +ALL NECESSARY SERVICING, REPAIR OR CORRECTION. + + 16. Limitation of Liability. + + IN NO EVENT UNLESS REQUIRED BY APPLICABLE LAW OR AGREED TO IN WRITING +WILL ANY COPYRIGHT HOLDER, OR ANY OTHER PARTY WHO MODIFIES AND/OR CONVEYS +THE PROGRAM AS PERMITTED ABOVE, BE LIABLE TO YOU FOR DAMAGES, INCLUDING ANY +GENERAL, SPECIAL, INCIDENTAL OR CONSEQUENTIAL DAMAGES ARISING OUT OF THE +USE OR INABILITY TO USE THE PROGRAM (INCLUDING BUT NOT LIMITED TO LOSS OF +DATA OR DATA BEING RENDERED INACCURATE OR LOSSES SUSTAINED BY YOU OR THIRD +PARTIES OR A FAILURE OF THE PROGRAM TO OPERATE WITH ANY OTHER PROGRAMS), +EVEN IF SUCH HOLDER OR OTHER PARTY HAS BEEN ADVISED OF THE POSSIBILITY OF +SUCH DAMAGES. + + 17. Interpretation of Sections 15 and 16. + + If the disclaimer of warranty and limitation of liability provided +above cannot be given local legal effect according to their terms, +reviewing courts shall apply local law that most closely approximates +an absolute waiver of all civil liability in connection with the +Program, unless a warranty or assumption of liability accompanies a +copy of the Program in return for a fee. + + END OF TERMS AND CONDITIONS + + How to Apply These Terms to Your New Programs + + If you develop a new program, and you want it to be of the greatest +possible use to the public, the best way to achieve this is to make it +free software which everyone can redistribute and change under these terms. + + To do so, attach the following notices to the program. It is safest +to attach them to the start of each source file to most effectively +state the exclusion of warranty; and each file should have at least +the "copyright" line and a pointer to where the full notice is found. + + + Copyright (C) + + This program is free software: you can redistribute it and/or modify + it under the terms of the GNU General Public License as published by + the Free Software Foundation, either version 3 of the License, or + (at your option) any later version. + + This program is distributed in the hope that it will be useful, + but WITHOUT ANY WARRANTY; without even the implied warranty of + MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + GNU General Public License for more details. + + You should have received a copy of the GNU General Public License + along with this program. If not, see . + +Also add information on how to contact you by electronic and paper mail. + + If the program does terminal interaction, make it output a short +notice like this when it starts in an interactive mode: + + Copyright (C) + This program comes with ABSOLUTELY NO WARRANTY; for details type `show w'. + This is free software, and you are welcome to redistribute it + under certain conditions; type `show c' for details. + +The hypothetical commands `show w' and `show c' should show the appropriate +parts of the General Public License. Of course, your program's commands +might be different; for a GUI interface, you would use an "about box". + + You should also get your employer (if you work as a programmer) or school, +if any, to sign a "copyright disclaimer" for the program, if necessary. +For more information on this, and how to apply and follow the GNU GPL, see +. + + The GNU General Public License does not permit incorporating your program +into proprietary programs. If your program is a subroutine library, you +may consider it more useful to permit linking proprietary applications with +the library. If this is what you want to do, use the GNU Lesser General +Public License instead of this License. But first, please read +. diff --git a/README.md b/README.md new file mode 100644 index 0000000..79a5000 --- /dev/null +++ b/README.md @@ -0,0 +1,76 @@ +# Deep Learning Flower Classification + +This repository contains implementations of various deep learning models for classifying flower images. + +## Models Included + +The following models are implemented in this repository: + +* **DenseNet:** [DenseNet/](DenseNet/) +* **EfficientNet:** [EfficientNet/](EfficientNet/) +* **GoogLeNet:** [GoogLeNet/](GoogLeNet/) +* **Transformer:** [Transformer/](Transformer/) + +## Data + +The project uses flower image data located in: + +* `data/flower_data/`: Processed data for training and validation. +* `data/flower_photos/`: Original flower photos. + +*Note: Data and pre-trained model weights are excluded from the repository via the [`.gitignore`](.gitignore) file.* + +## Project Structure + +Each model resides in its own directory (e.g., `DenseNet/`, `EfficientNet/`). Within each model's directory, you will typically find: + +* `model.py`: Defines the model architecture. +* `train.py`: Script for training the model. +* `evaluate.py`: Script for evaluating the trained model. +* `predict.py`: Script for making predictions on new images. +* `my_dataset.py`: Defines the custom dataset loading logic. +* `utils.py`: Contains utility functions. +* `class_indices.json`: Maps class indices to class names. + +## Usage + +1. **Clone the repository:** + + ```bash + git clone https://github.com/xixu-me/deep-learning-flower-classification.git + cd deep-learning-flower-classification + ``` + +2. **Prepare Data:** Ensure the required datasets are present in the `data/` directory as expected by the scripts. +3. **Install Dependencies:** Install necessary Python libraries (e.g., PyTorch, torchvision, numpy, matplotlib). *Consider adding a `requirements.txt` file.* +4. **Navigate to a Model Directory:** + + ```bash + cd DenseNet/ # or EfficientNet/, GoogLeNet/, Transformer/ + ``` + +5. **Train:** + + ```bash + python train.py # Add necessary arguments + ``` + +6. **Evaluate:** + + ```bash + python evaluate.py # Add necessary arguments + ``` + +7. **Predict:** + + ```bash + python predict.py --image_path # Add necessary arguments + ``` + +*Refer to the specific scripts within each model directory for detailed usage instructions and available arguments.* + +## License + +Copyright © [Xi Xu](https://xi-xu.me). All rights reserved. + +Licensed under the [GPL-3.0](LICENSE) license. diff --git a/Transformer/class_indices.json b/Transformer/class_indices.json new file mode 100644 index 0000000..84f8a53 --- /dev/null +++ b/Transformer/class_indices.json @@ -0,0 +1,7 @@ +{ + "0": "daisy", + "1": "dandelion", + "2": "roses", + "3": "sunflowers", + "4": "tulips" +} diff --git a/Transformer/evaluate_models.py b/Transformer/evaluate_models.py new file mode 100644 index 0000000..a0acb90 --- /dev/null +++ b/Transformer/evaluate_models.py @@ -0,0 +1,329 @@ +import argparse +import csv +import json +import os + +import matplotlib.pyplot as plt +import numpy as np +import torch +from my_dataset import MyDataSet +from sklearn.metrics import ( + accuracy_score, + confusion_matrix, + f1_score, + precision_score, + recall_score, +) +from torch.utils.data import DataLoader +from torchvision import transforms +from tqdm import tqdm +from utils import read_split_data +from vit_model import ( + vit_base_patch16_224_in21k, + vit_base_patch32_224_in21k, + vit_large_patch16_224_in21k, + vit_large_patch32_224_in21k, +) + + +def load_model(model_name, num_classes, device, weights_path): + """Load a model based on its name.""" + model_mapping = { + "vit_base_patch16_224_in21k": vit_base_patch16_224_in21k, + "vit_base_patch32_224_in21k": vit_base_patch32_224_in21k, + "vit_large_patch16_224_in21k": vit_large_patch16_224_in21k, + "vit_large_patch32_224_in21k": vit_large_patch32_224_in21k, + } + + # Create model + model = model_mapping[model_name](num_classes=num_classes, has_logits=False) + + # Load weights + assert os.path.exists( + weights_path + ), f"Weights file: '{weights_path}' does not exist." + model.load_state_dict(torch.load(weights_path, map_location=device)) + + model.to(device) + model.eval() + + return model + + +def evaluate_model(model, data_loader, device, num_classes): + """Evaluate a model and compute metrics.""" + true_labels = [] + predictions = [] + + # For per-class metrics + class_correct = [0] * num_classes + class_total = [0] * num_classes + + with torch.no_grad(): + for images, labels in tqdm(data_loader): + images = images.to(device) + labels = labels.to(device) + + outputs = model(images) + _, preds = torch.max(outputs, 1) + + true_labels.extend(labels.cpu().numpy()) + predictions.extend(preds.cpu().numpy()) + + correct = (preds == labels).squeeze() + for i in range(len(labels)): + label = labels[i].item() + class_correct[label] += correct[i].item() + class_total[label] += 1 + + # Convert to numpy arrays + true_labels = np.array(true_labels) + predictions = np.array(predictions) + + # Calculate metrics + metrics = { + "accuracy": accuracy_score(true_labels, predictions), + "precision_macro": precision_score(true_labels, predictions, average="macro"), + "precision_weighted": precision_score( + true_labels, predictions, average="weighted" + ), + "recall_macro": recall_score(true_labels, predictions, average="macro"), + "recall_weighted": recall_score(true_labels, predictions, average="weighted"), + "f1_macro": f1_score(true_labels, predictions, average="macro"), + "f1_weighted": f1_score(true_labels, predictions, average="weighted"), + "confusion_matrix": confusion_matrix(true_labels, predictions), + } + + # Calculate per-class accuracy + for i in range(num_classes): + if class_total[i] > 0: + metrics[f"class_{i}_accuracy"] = class_correct[i] / class_total[i] + else: + metrics[f"class_{i}_accuracy"] = 0.0 + + return metrics + + +def visualize_metrics(metrics_dict, class_indict, output_dir): + """Generate visualizations comparing model performance.""" + os.makedirs(output_dir, exist_ok=True) + model_names = list(metrics_dict.keys()) + short_names = [name.split("_")[0:3] for name in model_names] + short_names = [f"{parts[0]} {parts[1]}\n{parts[2]}" for parts in short_names] + + # 1. Overall metrics comparison + plt.figure(figsize=(14, 8)) + metrics_to_plot = ["accuracy", "precision_macro", "recall_macro", "f1_macro"] + metric_names = ["Accuracy", "Precision", "Recall", "F1-Score"] + + x = np.arange(len(model_names)) + width = 0.2 + + for i, metric in enumerate(metrics_to_plot): + values = [metrics_dict[model][metric] for model in model_names] + plt.bar(x + i * width - 0.3, values, width, label=metric_names[i]) + + plt.xlabel("Model Architecture") + plt.ylabel("Score") + plt.title("Model Performance Comparison") + plt.xticks(x, short_names) + plt.legend() + plt.tight_layout() + plt.savefig(os.path.join(output_dir, "overall_metrics.png")) + plt.close() + + # 2. Per-class accuracy comparison + plt.figure(figsize=(14, 8)) + num_classes = len(class_indict) + + for i, model_name in enumerate(model_names): + class_acc = [ + metrics_dict[model_name][f"class_{c}_accuracy"] for c in range(num_classes) + ] + plt.plot( + range(num_classes), + class_acc, + marker="o", + linestyle="-", + linewidth=2, + markersize=8, + label=short_names[i], + ) + + plt.xlabel("Class") + plt.ylabel("Accuracy") + plt.title("Per-Class Accuracy Comparison") + plt.xticks(range(num_classes), [class_indict[str(i)] for i in range(num_classes)]) + plt.legend() + plt.grid(True, linestyle="--", alpha=0.7) + plt.tight_layout() + plt.savefig(os.path.join(output_dir, "per_class_accuracy.png")) + plt.close() + + # 3. Radar chart for comprehensive comparison + metrics_for_radar = ["accuracy", "precision_macro", "recall_macro", "f1_macro"] + labels = ["Accuracy", "Precision", "Recall", "F1-Score"] + + # Create a figure with multiple radar charts + fig, axes = plt.subplots(1, 1, figsize=(10, 10), subplot_kw=dict(polar=True)) + angles = np.linspace(0, 2 * np.pi, len(metrics_for_radar), endpoint=False).tolist() + angles += angles[:1] # Close the loop + + # Define some colors for different models + colors = ["#1f77b4", "#ff7f0e", "#2ca02c", "#d62728"] + + for i, model_name in enumerate(model_names): + values = [metrics_dict[model_name][metric] for metric in metrics_for_radar] + values += values[:1] # Close the loop + + axes.plot(angles, values, color=colors[i], linewidth=2, label=short_names[i]) + axes.fill(angles, values, color=colors[i], alpha=0.25) + + # Set labels and title + axes.set_xticks(angles[:-1]) + axes.set_xticklabels(labels) + axes.set_title("Model Performance Radar Chart") + axes.legend(loc="upper right", bbox_to_anchor=(0.1, 0.1)) + + plt.tight_layout() + plt.savefig(os.path.join(output_dir, "radar_comparison.png")) + plt.close() + + +def create_metrics_table(metrics_dict, output_dir): + """Create and save metrics data in CSV format.""" + # Define which metrics to include in the table + metrics_to_include = ["accuracy", "precision_macro", "recall_macro", "f1_macro"] + metric_names = ["Accuracy", "Precision", "Recall", "F1-Score"] + + # Get model names + model_names = list(metrics_dict.keys()) + short_names = [" ".join(name.split("_")[0:3]) for name in model_names] + + # Prepare CSV file path + csv_path = os.path.join(output_dir, "metrics_table.csv") + + # Write data to CSV + with open(csv_path, "w", newline="") as csvfile: + writer = csv.writer(csvfile) + + # Write header row with metric names + header = ["Model"] + metric_names + writer.writerow(header) + + # Write data for each model + for i, model_name in enumerate(model_names): + row = [short_names[i]] + row.extend( + [ + f"{metrics_dict[model_name][metric]:.4f}" + for metric in metrics_to_include + ] + ) + writer.writerow(row) + + return csv_path + + +def main(): + parser = argparse.ArgumentParser(description="Evaluate ViT models") + parser.add_argument( + "--data-path", type=str, default="../data/flower_photos", help="Dataset path" + ) + parser.add_argument( + "--weights-dir", + type=str, + default="./weights", + help="Directory containing model weights", + ) + parser.add_argument( + "--output-dir", + type=str, + default="./evaluation_results", + help="Directory to save results", + ) + parser.add_argument( + "--batch-size", type=int, default=16, help="Batch size for evaluation" + ) + parser.add_argument("--num-classes", type=int, default=5, help="Number of classes") + parser.add_argument("--device", default="cuda:0", help="Device to use") + + args = parser.parse_args() + + device = torch.device(args.device if torch.cuda.is_available() else "cpu") + os.makedirs(args.output_dir, exist_ok=True) + + # Load data + _, _, val_images_path, val_images_label = read_split_data(args.data_path) + + # Data preprocessing + data_transform = transforms.Compose( + [ + transforms.Resize(256), + transforms.CenterCrop(224), + transforms.ToTensor(), + transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]), + ] + ) + + val_dataset = MyDataSet( + images_path=val_images_path, + images_class=val_images_label, + transform=data_transform, + ) + + nw = min([os.cpu_count(), args.batch_size if args.batch_size > 1 else 0, 8]) + val_loader = DataLoader( + val_dataset, + batch_size=args.batch_size, + shuffle=False, + num_workers=nw, + collate_fn=val_dataset.collate_fn, + ) + + # Load class indices + with open("class_indices.json", "r") as f: + class_indict = json.load(f) + + # Models to evaluate + models_to_evaluate = { + "vit_base_patch16_224_in21k": os.path.join( + args.weights_dir, "vit_base_patch16_224_in21k", "best_model.pth" + ), + "vit_base_patch32_224_in21k": os.path.join( + args.weights_dir, "vit_base_patch32_224_in21k", "best_model.pth" + ), + "vit_large_patch16_224_in21k": os.path.join( + args.weights_dir, "vit_large_patch16_224_in21k", "best_model.pth" + ), + "vit_large_patch32_224_in21k": os.path.join( + args.weights_dir, "vit_large_patch32_224_in21k", "best_model.pth" + ), + } + + all_metrics = {} + + # Evaluate each model + for model_name, weights_path in models_to_evaluate.items(): + print(f"\nEvaluating {model_name}...") + model = load_model(model_name, args.num_classes, device, weights_path) + metrics = evaluate_model(model, val_loader, device, args.num_classes) + all_metrics[model_name] = metrics + + print(f"Accuracy: {metrics['accuracy']:.4f}") + print(f"Precision: {metrics['precision_macro']:.4f}") + print(f"Recall: {metrics['recall_macro']:.4f}") + print(f"F1 Score: {metrics['f1_macro']:.4f}") + + # Visualize results + visualize_metrics(all_metrics, class_indict, args.output_dir) + + # Create and save metrics table as CSV instead of image + table_path = create_metrics_table(all_metrics, args.output_dir) + print(f"Metrics table saved to {table_path}") + + print(f"Evaluation completed. Results saved to {args.output_dir}") + + +if __name__ == "__main__": + main() diff --git a/Transformer/evaluation_results/metrics_table.csv b/Transformer/evaluation_results/metrics_table.csv new file mode 100644 index 0000000..0a071d6 --- /dev/null +++ b/Transformer/evaluation_results/metrics_table.csv @@ -0,0 +1,5 @@ +Model,Accuracy,Precision,Recall,F1-Score +vit base patch16,0.9822,0.9811,0.9820,0.9815 +vit base patch32,0.9726,0.9715,0.9731,0.9720 +vit large patch16,0.9808,0.9792,0.9807,0.9799 +vit large patch32,0.9754,0.9742,0.9762,0.9750 diff --git a/Transformer/evaluation_results/overall_metrics.png b/Transformer/evaluation_results/overall_metrics.png new file mode 100644 index 0000000..13c8946 Binary files /dev/null and b/Transformer/evaluation_results/overall_metrics.png differ diff --git a/Transformer/evaluation_results/per_class_accuracy.png b/Transformer/evaluation_results/per_class_accuracy.png new file mode 100644 index 0000000..c93418a Binary files /dev/null and b/Transformer/evaluation_results/per_class_accuracy.png differ diff --git a/Transformer/evaluation_results/radar_comparison.png b/Transformer/evaluation_results/radar_comparison.png new file mode 100644 index 0000000..0f85cb9 Binary files /dev/null and b/Transformer/evaluation_results/radar_comparison.png differ diff --git a/Transformer/flops.py b/Transformer/flops.py new file mode 100644 index 0000000..761c54e --- /dev/null +++ b/Transformer/flops.py @@ -0,0 +1,23 @@ +import torch +from fvcore.nn import FlopCountAnalysis +from vit_model import Attention + + +def main(): + + a1 = Attention(dim=512, num_heads=1) + a1.proj = torch.nn.Identity() + + a2 = Attention(dim=512, num_heads=8) + + t = (torch.rand(32, 1024, 512),) + + flops1 = FlopCountAnalysis(a1, t) + print("Self-Attention FLOPs:", flops1.total()) + + flops2 = FlopCountAnalysis(a2, t) + print("Multi-Head Attention FLOPs:", flops2.total()) + + +if __name__ == "__main__": + main() diff --git a/Transformer/my_dataset.py b/Transformer/my_dataset.py new file mode 100644 index 0000000..6603067 --- /dev/null +++ b/Transformer/my_dataset.py @@ -0,0 +1,36 @@ +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]) + + if img.mode != "RGB": + img = img.convert("RGB") + + label = self.images_class[item] + + if self.transform is not None: + img = self.transform(img) + + return img, label + + @staticmethod + def collate_fn(batch): + + images, labels = tuple(zip(*batch)) + + images = torch.stack(images, dim=0) + labels = torch.as_tensor(labels) + return images, labels diff --git a/Transformer/predict.py b/Transformer/predict.py new file mode 100644 index 0000000..bbb468a --- /dev/null +++ b/Transformer/predict.py @@ -0,0 +1,63 @@ +import json +import os + +import matplotlib.pyplot as plt +import torch +from PIL import Image +from torchvision import transforms +from vit_model import vit_base_patch16_224_in21k as create_model + + +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.5, 0.5, 0.5], [0.5, 0.5, 0.5]), + ] + ) + + img_path = "../tulip.jpg" + assert os.path.exists(img_path), "file: '{}' dose not exist.".format(img_path) + img = Image.open(img_path) + plt.imshow(img) + + img = data_transform(img) + + img = torch.unsqueeze(img, dim=0) + + json_path = "./class_indices.json" + assert os.path.exists(json_path), "file: '{}' dose not exist.".format(json_path) + + with open(json_path, "r") as f: + class_indict = json.load(f) + + model = create_model(num_classes=5, has_logits=False).to(device) + + model_weight_path = "./weights/model-9.pth" + model.load_state_dict(torch.load(model_weight_path, map_location=device)) + model.eval() + with torch.no_grad(): + + 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() diff --git a/Transformer/train.py b/Transformer/train.py new file mode 100644 index 0000000..eb69ca4 --- /dev/null +++ b/Transformer/train.py @@ -0,0 +1,287 @@ +import argparse +import math +import os + +import matplotlib.pyplot as plt +import numpy as np +import torch +import torch.optim as optim +import torch.optim.lr_scheduler as lr_scheduler +import vit_model +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 train_model( + args, + model_name, + model_creator, + pretrained_weights, + device, + train_loader, + val_loader, + tb_writer, + history, +): + print(f"\n{'='*20} Training {model_name} {'='*20}") + + # Create model + model = model_creator(num_classes=args.num_classes, has_logits=False).to(device) + + # Load pretrained weights + if os.path.exists(pretrained_weights): + print(f"Loading pretrained weights: {pretrained_weights}") + weights_dict = torch.load(pretrained_weights, map_location=device) + # Remove classifier weights that don't match (num_classes difference) + for k in list(weights_dict.keys()): + if "head" in k: + del weights_dict[k] + print(f"Loaded pretrained weights from {pretrained_weights}") + model.load_state_dict(weights_dict, strict=False) + else: + print(f"Warning: Pretrained weights {pretrained_weights} not found!") + + # Freeze layers if needed + if args.freeze_layers: + for name, para in model.named_parameters(): + if "head" not in name and "pre_logits" not in name: + para.requires_grad_(False) + else: + print(f"Training {name}") + + # Set up optimizer + pg = [p for p in model.parameters() if p.requires_grad] + optimizer = optim.SGD(pg, lr=args.lr, momentum=0.9, weight_decay=5e-5) + + # Set up scheduler + lf = ( + lambda x: ((1 + math.cos(x * math.pi / args.epochs)) / 2) * (1 - args.lrf) + + args.lrf + ) + scheduler = lr_scheduler.LambdaLR(optimizer, lr_lambda=lf) + + # Training loop + best_acc = 0.0 + model_history = {"train_loss": [], "train_acc": [], "val_loss": [], "val_acc": []} + + for epoch in range(args.epochs): + # Train + train_loss, train_acc = train_one_epoch( + model=model, + optimizer=optimizer, + data_loader=train_loader, + device=device, + epoch=epoch, + ) + scheduler.step() + + # Evaluate + val_loss, val_acc = evaluate( + model=model, data_loader=val_loader, device=device, epoch=epoch + ) + + # Record metrics + model_history["train_loss"].append(train_loss) + model_history["train_acc"].append(train_acc) + model_history["val_loss"].append(val_loss) + model_history["val_acc"].append(val_acc) + + # TensorBoard logging + tags = ["train_loss", "train_acc", "val_loss", "val_acc", "learning_rate"] + tb_writer.add_scalar(f"{model_name}/{tags[0]}", train_loss, epoch) + tb_writer.add_scalar(f"{model_name}/{tags[1]}", train_acc, epoch) + tb_writer.add_scalar(f"{model_name}/{tags[2]}", val_loss, epoch) + tb_writer.add_scalar(f"{model_name}/{tags[3]}", val_acc, epoch) + tb_writer.add_scalar( + f"{model_name}/{tags[4]}", optimizer.param_groups[0]["lr"], epoch + ) + + # Save model + model_dir = f"./weights/{model_name}" + os.makedirs(model_dir, exist_ok=True) + torch.save(model.state_dict(), f"{model_dir}/model-{epoch}.pth") + + # Save best model + if val_acc > best_acc: + best_acc = val_acc + torch.save(model.state_dict(), f"{model_dir}/best_model.pth") + + # Store the history for this model + history[model_name] = model_history + return history + + +def plot_training_results(history): + """Plot the training and validation results for all models""" + models = list(history.keys()) + epochs = range(1, len(history[models[0]]["train_loss"]) + 1) + + plt.figure(figsize=(20, 15)) + + # Plot training & validation loss + plt.subplot(2, 1, 1) + for model_name in models: + plt.plot( + epochs, history[model_name]["train_loss"], "-o", label=f"{model_name} Train" + ) + plt.plot( + epochs, history[model_name]["val_loss"], "-s", label=f"{model_name} Val" + ) + + plt.title("Training and Validation Loss", fontsize=15) + plt.xlabel("Epochs", fontsize=12) + plt.ylabel("Loss", fontsize=12) + plt.legend(fontsize=10) + plt.grid(True) + + # Plot training & validation accuracy + plt.subplot(2, 1, 2) + for model_name in models: + plt.plot( + epochs, history[model_name]["train_acc"], "-o", label=f"{model_name} Train" + ) + plt.plot( + epochs, history[model_name]["val_acc"], "-s", label=f"{model_name} Val" + ) + + plt.title("Training and Validation Accuracy", fontsize=15) + plt.xlabel("Epochs", fontsize=12) + plt.ylabel("Accuracy", fontsize=12) + plt.legend(fontsize=10) + plt.grid(True) + + plt.tight_layout() + plt.savefig("./training_results.png", dpi=300) + plt.show() + + +def main(args): + device = torch.device(args.device if torch.cuda.is_available() else "cpu") + + if os.path.exists("./weights") is False: + os.makedirs("./weights") + + tb_writer = SummaryWriter() + + 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.5, 0.5, 0.5], [0.5, 0.5, 0.5]), + ] + ), + "val": transforms.Compose( + [ + transforms.Resize(256), + transforms.CenterCrop(224), + transforms.ToTensor(), + transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]), + ] + ), + } + + 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]) + 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, + ) + + # Define models to train and their corresponding pretrained weight files + models_to_train = { + "vit_base_patch16_224_in21k": { + "creator": vit_model.vit_base_patch16_224_in21k, + "weights": "jx_vit_base_patch16_224_in21k-e5005f0a.pth", + }, + "vit_base_patch32_224_in21k": { + "creator": vit_model.vit_base_patch32_224_in21k, + "weights": "jx_vit_base_patch32_224_in21k-8db57226.pth", + }, + "vit_large_patch16_224_in21k": { + "creator": vit_model.vit_large_patch16_224_in21k, + "weights": "jx_vit_large_patch16_224_in21k-606da67d.pth", + }, + "vit_large_patch32_224_in21k": { + "creator": vit_model.vit_large_patch32_224_in21k, + "weights": "jx_vit_large_patch32_224_in21k-9046d2e7.pth", + }, + } + + # Dictionary to store training history for all models + training_history = {} + + # Train each model + for model_name, model_info in models_to_train.items(): + training_history = train_model( + args=args, + model_name=model_name, + model_creator=model_info["creator"], + pretrained_weights=model_info["weights"], + device=device, + train_loader=train_loader, + val_loader=val_loader, + tb_writer=tb_writer, + history=training_history, + ) + + # Plot and save the training results + plot_training_results(training_history) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--num_classes", type=int, default=5) + parser.add_argument( + "--epochs", type=int, default=10 + ) # Reduced epochs for testing all models + parser.add_argument("--batch-size", type=int, default=8) + parser.add_argument("--lr", type=float, default=0.001) + parser.add_argument("--lrf", type=float, default=0.01) + + parser.add_argument("--data-path", type=str, default="../data/flower_photos") + parser.add_argument("--model-name", default="", help="create model name") + + parser.add_argument("--weights", type=str, default="", help="initial weights path") + + parser.add_argument("--freeze-layers", type=bool, default=True) + parser.add_argument( + "--device", default="cuda:0", help="device id (i.e. 0 or 0,1 or cpu)" + ) + + opt = parser.parse_args() + + main(opt) diff --git a/Transformer/training_results.png b/Transformer/training_results.png new file mode 100644 index 0000000..70c807e Binary files /dev/null and b/Transformer/training_results.png differ diff --git a/Transformer/utils.py b/Transformer/utils.py new file mode 100644 index 0000000..c6f65e5 --- /dev/null +++ b/Transformer/utils.py @@ -0,0 +1,180 @@ +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) + + 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") + + 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") + + plt.xlabel("image class") + + 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): + + img = images[i].numpy().transpose(1, 2, 0) + + 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([]) + plt.yticks([]) + 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() + accu_loss = torch.zeros(1).to(device) + accu_num = torch.zeros(1).to(device) + optimizer.zero_grad() + + sample_num = 0 + data_loader = tqdm(data_loader, file=sys.stdout) + for step, data in enumerate(data_loader): + images, labels = data + sample_num += images.shape[0] + + pred = model(images.to(device)) + pred_classes = torch.max(pred, dim=1)[1] + accu_num += torch.eq(pred_classes, labels.to(device)).sum() + + loss = loss_function(pred, labels.to(device)) + loss.backward() + accu_loss += loss.detach() + + data_loader.desc = "[train epoch {}] loss: {:.3f}, acc: {:.3f}".format( + epoch, accu_loss.item() / (step + 1), accu_num.item() / sample_num + ) + + if not torch.isfinite(loss): + print("WARNING: non-finite loss, ending training ", loss) + sys.exit(1) + + optimizer.step() + optimizer.zero_grad() + + return accu_loss.item() / (step + 1), accu_num.item() / sample_num + + +@torch.no_grad() +def evaluate(model, data_loader, device, epoch): + loss_function = torch.nn.CrossEntropyLoss() + + model.eval() + + accu_num = torch.zeros(1).to(device) + accu_loss = torch.zeros(1).to(device) + + sample_num = 0 + data_loader = tqdm(data_loader, file=sys.stdout) + for step, data in enumerate(data_loader): + images, labels = data + sample_num += images.shape[0] + + pred = model(images.to(device)) + pred_classes = torch.max(pred, dim=1)[1] + accu_num += torch.eq(pred_classes, labels.to(device)).sum() + + loss = loss_function(pred, labels.to(device)) + accu_loss += loss + + data_loader.desc = "[valid epoch {}] loss: {:.3f}, acc: {:.3f}".format( + epoch, accu_loss.item() / (step + 1), accu_num.item() / sample_num + ) + + return accu_loss.item() / (step + 1), accu_num.item() / sample_num diff --git a/Transformer/vit_model.py b/Transformer/vit_model.py new file mode 100644 index 0000000..9424d44 --- /dev/null +++ b/Transformer/vit_model.py @@ -0,0 +1,428 @@ +from collections import OrderedDict +from functools import partial + +import torch +import torch.nn as nn + + +def drop_path(x, drop_prob: float = 0.0, training: bool = False): + + if drop_prob == 0.0 or not training: + return x + keep_prob = 1 - drop_prob + shape = (x.shape[0],) + (1,) * (x.ndim - 1) + random_tensor = keep_prob + torch.rand(shape, dtype=x.dtype, device=x.device) + random_tensor.floor_() + output = x.div(keep_prob) * random_tensor + return output + + +class DropPath(nn.Module): + + def __init__(self, drop_prob=None): + super(DropPath, self).__init__() + self.drop_prob = drop_prob + + def forward(self, x): + return drop_path(x, self.drop_prob, self.training) + + +class PatchEmbed(nn.Module): + + def __init__( + self, img_size=224, patch_size=16, in_c=3, embed_dim=768, norm_layer=None + ): + super().__init__() + img_size = (img_size, img_size) + patch_size = (patch_size, patch_size) + self.img_size = img_size + self.patch_size = patch_size + self.grid_size = (img_size[0] // patch_size[0], img_size[1] // patch_size[1]) + self.num_patches = self.grid_size[0] * self.grid_size[1] + + self.proj = nn.Conv2d( + in_c, embed_dim, kernel_size=patch_size, stride=patch_size + ) + self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity() + + def forward(self, x): + B, C, H, W = x.shape + assert ( + H == self.img_size[0] and W == self.img_size[1] + ), f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]})." + + x = self.proj(x).flatten(2).transpose(1, 2) + x = self.norm(x) + return x + + +class Attention(nn.Module): + def __init__( + self, + dim, + num_heads=8, + qkv_bias=False, + qk_scale=None, + attn_drop_ratio=0.0, + proj_drop_ratio=0.0, + ): + super(Attention, self).__init__() + self.num_heads = num_heads + head_dim = dim // num_heads + self.scale = qk_scale or head_dim**-0.5 + self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) + self.attn_drop = nn.Dropout(attn_drop_ratio) + self.proj = nn.Linear(dim, dim) + self.proj_drop = nn.Dropout(proj_drop_ratio) + + def forward(self, x): + + B, N, C = x.shape + + qkv = ( + self.qkv(x) + .reshape(B, N, 3, self.num_heads, C // self.num_heads) + .permute(2, 0, 3, 1, 4) + ) + + q, k, v = qkv[0], qkv[1], qkv[2] + + attn = (q @ k.transpose(-2, -1)) * self.scale + attn = attn.softmax(dim=-1) + attn = self.attn_drop(attn) + + x = (attn @ v).transpose(1, 2).reshape(B, N, C) + x = self.proj(x) + x = self.proj_drop(x) + return x + + +class Mlp(nn.Module): + + def __init__( + self, + in_features, + hidden_features=None, + out_features=None, + act_layer=nn.GELU, + drop=0.0, + ): + super().__init__() + out_features = out_features or in_features + hidden_features = hidden_features or in_features + self.fc1 = nn.Linear(in_features, hidden_features) + self.act = act_layer() + self.fc2 = nn.Linear(hidden_features, out_features) + self.drop = nn.Dropout(drop) + + def forward(self, x): + x = self.fc1(x) + x = self.act(x) + x = self.drop(x) + x = self.fc2(x) + x = self.drop(x) + return x + + +class Block(nn.Module): + def __init__( + self, + dim, + num_heads, + mlp_ratio=4.0, + qkv_bias=False, + qk_scale=None, + drop_ratio=0.0, + attn_drop_ratio=0.0, + drop_path_ratio=0.0, + act_layer=nn.GELU, + norm_layer=nn.LayerNorm, + ): + super(Block, self).__init__() + self.norm1 = norm_layer(dim) + self.attn = Attention( + dim, + num_heads=num_heads, + qkv_bias=qkv_bias, + qk_scale=qk_scale, + attn_drop_ratio=attn_drop_ratio, + proj_drop_ratio=drop_ratio, + ) + + self.drop_path = ( + DropPath(drop_path_ratio) if drop_path_ratio > 0.0 else nn.Identity() + ) + self.norm2 = norm_layer(dim) + mlp_hidden_dim = int(dim * mlp_ratio) + self.mlp = Mlp( + in_features=dim, + hidden_features=mlp_hidden_dim, + act_layer=act_layer, + drop=drop_ratio, + ) + + def forward(self, x): + x = x + self.drop_path(self.attn(self.norm1(x))) + x = x + self.drop_path(self.mlp(self.norm2(x))) + return x + + +class VisionTransformer(nn.Module): + def __init__( + self, + img_size=224, + patch_size=16, + in_c=3, + num_classes=1000, + embed_dim=768, + depth=12, + num_heads=12, + mlp_ratio=4.0, + qkv_bias=True, + qk_scale=None, + representation_size=None, + distilled=False, + drop_ratio=0.0, + attn_drop_ratio=0.0, + drop_path_ratio=0.0, + embed_layer=PatchEmbed, + norm_layer=None, + act_layer=None, + ): + + super(VisionTransformer, self).__init__() + self.num_classes = num_classes + self.num_features = self.embed_dim = embed_dim + self.num_tokens = 2 if distilled else 1 + norm_layer = norm_layer or partial(nn.LayerNorm, eps=1e-6) + act_layer = act_layer or nn.GELU + + self.patch_embed = embed_layer( + img_size=img_size, patch_size=patch_size, in_c=in_c, embed_dim=embed_dim + ) + num_patches = self.patch_embed.num_patches + + self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) + self.dist_token = ( + nn.Parameter(torch.zeros(1, 1, embed_dim)) if distilled else None + ) + self.pos_embed = nn.Parameter( + torch.zeros(1, num_patches + self.num_tokens, embed_dim) + ) + self.pos_drop = nn.Dropout(p=drop_ratio) + + dpr = [x.item() for x in torch.linspace(0, drop_path_ratio, depth)] + self.blocks = nn.Sequential( + *[ + Block( + dim=embed_dim, + num_heads=num_heads, + mlp_ratio=mlp_ratio, + qkv_bias=qkv_bias, + qk_scale=qk_scale, + drop_ratio=drop_ratio, + attn_drop_ratio=attn_drop_ratio, + drop_path_ratio=dpr[i], + norm_layer=norm_layer, + act_layer=act_layer, + ) + for i in range(depth) + ] + ) + self.norm = norm_layer(embed_dim) + + if representation_size and not distilled: + self.has_logits = True + self.num_features = representation_size + self.pre_logits = nn.Sequential( + OrderedDict( + [ + ("fc", nn.Linear(embed_dim, representation_size)), + ("act", nn.Tanh()), + ] + ) + ) + else: + self.has_logits = False + self.pre_logits = nn.Identity() + + self.head = ( + nn.Linear(self.num_features, num_classes) + if num_classes > 0 + else nn.Identity() + ) + self.head_dist = None + if distilled: + self.head_dist = ( + nn.Linear(self.embed_dim, self.num_classes) + if num_classes > 0 + else nn.Identity() + ) + + # Weight init + nn.init.trunc_normal_(self.pos_embed, std=0.02) + if self.dist_token is not None: + nn.init.trunc_normal_(self.dist_token, std=0.02) + + nn.init.trunc_normal_(self.cls_token, std=0.02) + self.apply(_init_vit_weights) + + def forward_features(self, x): + + x = self.patch_embed(x) + + cls_token = self.cls_token.expand(x.shape[0], -1, -1) + if self.dist_token is None: + x = torch.cat((cls_token, x), dim=1) + else: + x = torch.cat( + (cls_token, self.dist_token.expand(x.shape[0], -1, -1), x), dim=1 + ) + + x = self.pos_drop(x + self.pos_embed) + x = self.blocks(x) + x = self.norm(x) + if self.dist_token is None: + return self.pre_logits(x[:, 0]) + else: + return x[:, 0], x[:, 1] + + def forward(self, x): + x = self.forward_features(x) + if self.head_dist is not None: + x, x_dist = self.head(x[0]), self.head_dist(x[1]) + if self.training and not torch.jit.is_scripting(): + + return x, x_dist + else: + return (x + x_dist) / 2 + else: + x = self.head(x) + return x + + +def _init_vit_weights(m): + + if isinstance(m, nn.Linear): + nn.init.trunc_normal_(m.weight, std=0.01) + if m.bias is not None: + nn.init.zeros_(m.bias) + elif isinstance(m, nn.Conv2d): + nn.init.kaiming_normal_(m.weight, mode="fan_out") + if m.bias is not None: + nn.init.zeros_(m.bias) + elif isinstance(m, nn.LayerNorm): + nn.init.zeros_(m.bias) + nn.init.ones_(m.weight) + + +def vit_base_patch16_224(num_classes: int = 1000): + + model = VisionTransformer( + img_size=224, + patch_size=16, + embed_dim=768, + depth=12, + num_heads=12, + representation_size=None, + num_classes=num_classes, + ) + return model + + +def vit_base_patch16_224_in21k(num_classes: int = 21843, has_logits: bool = True): + + model = VisionTransformer( + img_size=224, + patch_size=16, + embed_dim=768, + depth=12, + num_heads=12, + representation_size=768 if has_logits else None, + num_classes=num_classes, + ) + return model + + +def vit_base_patch32_224(num_classes: int = 1000): + + model = VisionTransformer( + img_size=224, + patch_size=32, + embed_dim=768, + depth=12, + num_heads=12, + representation_size=None, + num_classes=num_classes, + ) + return model + + +def vit_base_patch32_224_in21k(num_classes: int = 21843, has_logits: bool = True): + + model = VisionTransformer( + img_size=224, + patch_size=32, + embed_dim=768, + depth=12, + num_heads=12, + representation_size=768 if has_logits else None, + num_classes=num_classes, + ) + return model + + +def vit_large_patch16_224(num_classes: int = 1000): + + model = VisionTransformer( + img_size=224, + patch_size=16, + embed_dim=1024, + depth=24, + num_heads=16, + representation_size=None, + num_classes=num_classes, + ) + return model + + +def vit_large_patch16_224_in21k(num_classes: int = 21843, has_logits: bool = True): + + model = VisionTransformer( + img_size=224, + patch_size=16, + embed_dim=1024, + depth=24, + num_heads=16, + representation_size=1024 if has_logits else None, + num_classes=num_classes, + ) + return model + + +def vit_large_patch32_224_in21k(num_classes: int = 21843, has_logits: bool = True): + + model = VisionTransformer( + img_size=224, + patch_size=32, + embed_dim=1024, + depth=24, + num_heads=16, + representation_size=1024 if has_logits else None, + num_classes=num_classes, + ) + return model + + +def vit_huge_patch14_224_in21k(num_classes: int = 21843, has_logits: bool = True): + + model = VisionTransformer( + img_size=224, + patch_size=14, + embed_dim=1280, + depth=32, + num_heads=16, + representation_size=1280 if has_logits else None, + num_classes=num_classes, + ) + return model