In [3]:
import os
import numpy as np
import pandas as pd
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
from torchvision.models import resnet50, ResNet50_Weights
from PIL import Image
from tqdm.auto import tqdm
from sklearn.utils.class_weight import compute_class_weight
from sklearn.metrics import (
accuracy_score,
balanced_accuracy_score,
f1_score
)
import matplotlib.pyplot as plt
In [4]:
print("PyTorch:", torch.__version__)
if torch.backends.mps.is_available():
device = torch.device("mps")
else:
device = torch.device("cpu")
print("Device:", device)
PyTorch: 2.13.0 Device: mps
In [5]:
PROJECT_DIR = "/Users/applesues01/Documents/Medical_Agent"
DATA_DIR = os.path.join(
PROJECT_DIR,
"data",
"HAM10000"
)
IMAGE_DIR1 = os.path.join(
DATA_DIR,
"HAM10000_images_part_1"
)
IMAGE_DIR2 = os.path.join(
DATA_DIR,
"HAM10000_images_part_2"
)
SPLIT_DIR = os.path.join(
DATA_DIR,
"splits"
)
CHECKPOINT_DIR = os.path.join(
PROJECT_DIR,
"checkpoints"
)
OUTPUT_DIR = os.path.join(
PROJECT_DIR,
"outputs"
)
os.makedirs(
CHECKPOINT_DIR,
exist_ok=True
)
os.makedirs(
OUTPUT_DIR,
exist_ok=True
)
print(DATA_DIR)
/Users/applesues01/Documents/Medical_Agent/data/HAM10000
In [6]:
metadata_path = os.path.join(
DATA_DIR,
"HAM10000_metadata.csv"
)
train_path = os.path.join(
SPLIT_DIR,
"train.csv"
)
val_path = os.path.join(
SPLIT_DIR,
"val.csv"
)
test_path = os.path.join(
SPLIT_DIR,
"test.csv"
)
df = pd.read_csv(metadata_path)
train_df = pd.read_csv(train_path)
val_df = pd.read_csv(val_path)
test_df = pd.read_csv(test_path)
print("全部图片:", len(df))
print("训练:", len(train_df))
print("验证:", len(val_df))
print("测试:", len(test_df))
全部图片: 10015 训练: 7002 验证: 1532 测试: 1481
In [7]:
label_map = {
"akiec":0,
"bcc":1,
"bkl":2,
"df":3,
"mel":4,
"nv":5,
"vasc":6
}
class_names = [
"akiec",
"bcc",
"bkl",
"df",
"mel",
"nv",
"vasc"
]
In [8]:
train_transform = transforms.Compose([
transforms.Resize(
(224,224)
),
transforms.RandomHorizontalFlip(),
transforms.RandomRotation(10),
transforms.ToTensor(),
transforms.Normalize(
mean=[
0.485,
0.456,
0.406
],
std=[
0.229,
0.224,
0.225
]
)
])
val_transform = transforms.Compose([
transforms.Resize(
(224,224)
),
transforms.ToTensor(),
transforms.Normalize(
mean=[
0.485,
0.456,
0.406
],
std=[
0.229,
0.224,
0.225
]
)
])
In [9]:
class HAMDataset(Dataset):
def __init__(
self,
dataframe,
image_dir1,
image_dir2,
transform=None
):
self.df = dataframe.reset_index(
drop=True
)
self.image_dir1 = image_dir1
self.image_dir2 = image_dir2
self.transform = transform
def __len__(self):
return len(self.df)
def __getitem__(
self,
idx
):
row = self.df.iloc[idx]
image_id = row["image_id"]
filename = image_id + ".jpg"
path1 = os.path.join(
self.image_dir1,
filename
)
path2 = os.path.join(
self.image_dir2,
filename
)
if os.path.exists(path1):
image_path = path1
else:
image_path = path2
image = Image.open(
image_path
).convert("RGB")
if self.transform:
image = self.transform(image)
label = label_map[
row["dx"]
]
return image, label
In [10]:
train_dataset = HAMDataset(
train_df,
IMAGE_DIR1,
IMAGE_DIR2,
train_transform
)
val_dataset = HAMDataset(
val_df,
IMAGE_DIR1,
IMAGE_DIR2,
val_transform
)
test_dataset = HAMDataset(
test_df,
IMAGE_DIR1,
IMAGE_DIR2,
val_transform
)
print(
len(train_dataset),
len(val_dataset),
len(test_dataset)
)
7002 1532 1481
In [11]:
train_loader = DataLoader(
train_dataset,
batch_size=16,
shuffle=True,
num_workers=0
)
val_loader = DataLoader(
val_dataset,
batch_size=16,
shuffle=False,
num_workers=0
)
test_loader = DataLoader(
test_dataset,
batch_size=16,
shuffle=False,
num_workers=0
)
In [12]:
weights = ResNet50_Weights.DEFAULT
model = resnet50(
weights=weights
)
# 冻结主干
for p in model.parameters():
p.requires_grad = False
# 修改分类层
model.fc = nn.Linear(
model.fc.in_features,
7
)
model = model.to(device)
print(model.fc)
Linear(in_features=2048, out_features=7, bias=True)
In [13]:
class_weights = compute_class_weight(
class_weight="balanced",
classes=np.array(class_names),
y=train_df["dx"]
)
class_weights = torch.tensor(
class_weights,
dtype=torch.float32
).to(device)
print(class_weights)
tensor([ 4.3491, 2.7330, 1.2924, 13.1617, 1.2857, 0.2138, 10.1039],
device='mps:0')
In [14]:
criterion = nn.CrossEntropyLoss(
weight=class_weights
)
optimizer = torch.optim.Adam(
model.fc.parameters(),
lr=1e-3
)
In [15]:
def train_one_epoch():
model.train()
total_loss = 0
preds = []
truths = []
for images, labels in tqdm(
train_loader
):
images = images.to(device)
labels = labels.to(device)
optimizer.zero_grad()
outputs = model(images)
loss = criterion(
outputs,
labels
)
loss.backward()
optimizer.step()
total_loss += (
loss.item()
*
images.size(0)
)
preds.extend(
outputs.argmax(1)
.detach()
.cpu()
.numpy()
)
truths.extend(
labels.cpu().numpy()
)
return {
"loss":
total_loss /
len(train_loader.dataset),
"acc":
accuracy_score(
truths,
preds
),
"f1":
f1_score(
truths,
preds,
average="macro"
)
}
In [16]:
def evaluate():
model.eval()
preds=[]
truths=[]
loss_sum=0
with torch.no_grad():
for images,labels in val_loader:
images=images.to(device)
labels=labels.to(device)
outputs=model(images)
loss=criterion(
outputs,
labels
)
loss_sum += (
loss.item()
*
images.size(0)
)
preds.extend(
outputs.argmax(1)
.cpu()
.numpy()
)
truths.extend(
labels.cpu()
.numpy()
)
return {
"loss":
loss_sum /
len(val_loader.dataset),
"acc":
accuracy_score(
truths,
preds
),
"balanced":
balanced_accuracy_score(
truths,
preds
),
"f1":
f1_score(
truths,
preds,
average="macro"
)
}
In [17]:
epochs = 10
best_f1 = 0
for epoch in range(epochs):
train_metrics = train_one_epoch()
val_metrics = evaluate()
print(
f"""
Epoch {epoch+1}
Train:
Loss {train_metrics['loss']:.4f}
Acc {train_metrics['acc']:.4f}
F1 {train_metrics['f1']:.4f}
Val:
Loss {val_metrics['loss']:.4f}
Acc {val_metrics['acc']:.4f}
Balanced {val_metrics['balanced']:.4f}
F1 {val_metrics['f1']:.4f}
"""
)
if val_metrics["f1"] > best_f1:
best_f1 = val_metrics["f1"]
torch.save(
model.state_dict(),
os.path.join(
CHECKPOINT_DIR,
"resnet50_image_only_best.pth"
)
)
print("保存最佳模型")
0%| | 0/438 [00:00<?, ?it/s]
Epoch 1 Train: Loss 1.5082 Acc 0.5921 F1 0.3524 Val: Loss 1.0923 Acc 0.6208 Balanced 0.4961 F1 0.4039 保存最佳模型
0%| | 0/438 [00:00<?, ?it/s]
Epoch 2 Train: Loss 1.1781 Acc 0.6390 F1 0.4475 Val: Loss 0.8972 Acc 0.6678 Balanced 0.5354 F1 0.4571 保存最佳模型
0%| | 0/438 [00:00<?, ?it/s]
Epoch 3 Train: Loss 1.0770 Acc 0.6742 F1 0.4999 Val: Loss 1.0274 Acc 0.6188 Balanced 0.5289 F1 0.4382
0%| | 0/438 [00:00<?, ?it/s]
Epoch 4 Train: Loss 0.9878 Acc 0.6798 F1 0.5176 Val: Loss 0.9742 Acc 0.6312 Balanced 0.5708 F1 0.4397
0%| | 0/438 [00:00<?, ?it/s]
Epoch 5 Train: Loss 0.9558 Acc 0.6842 F1 0.5277 Val: Loss 0.9793 Acc 0.6247 Balanced 0.5643 F1 0.4393
0%| | 0/438 [00:00<?, ?it/s]
Epoch 6 Train: Loss 0.8981 Acc 0.6965 F1 0.5513 Val: Loss 0.9010 Acc 0.6573 Balanced 0.5703 F1 0.4509
0%| | 0/438 [00:00<?, ?it/s]
Epoch 7 Train: Loss 0.8721 Acc 0.7069 F1 0.5557 Val: Loss 0.8433 Acc 0.6756 Balanced 0.5495 F1 0.4760 保存最佳模型
0%| | 0/438 [00:00<?, ?it/s]
Epoch 8 Train: Loss 0.8774 Acc 0.6982 F1 0.5566 Val: Loss 0.9819 Acc 0.6305 Balanced 0.6169 F1 0.4627
0%| | 0/438 [00:00<?, ?it/s]
Epoch 9 Train: Loss 0.8516 Acc 0.7144 F1 0.5886 Val: Loss 0.9413 Acc 0.6410 Balanced 0.5824 F1 0.4717
0%| | 0/438 [00:00<?, ?it/s]
Epoch 10 Train: Loss 0.8288 Acc 0.7081 F1 0.5839 Val: Loss 0.9063 Acc 0.6547 Balanced 0.5590 F1 0.4897 保存最佳模型
解冻ResNet50模型的layer4,开始模型微调¶
In [18]:
best_frozen_path = os.path.join(
CHECKPOINT_DIR,
"resnet50_image_only_best.pth"
)
best_finetune_path = os.path.join(
CHECKPOINT_DIR,
"resnet50_image_only_finetuned_best.pth"
)
In [19]:
model.load_state_dict(
torch.load(
best_frozen_path,
map_location=device,
weights_only=True
)
)
model = model.to(device)
print("已加载 Frozen ResNet50 最佳模型")
已加载 Frozen ResNet50 最佳模型
In [20]:
for parameter in model.parameters():
parameter.requires_grad = False
for parameter in model.layer4.parameters():
parameter.requires_grad = True
for parameter in model.fc.parameters():
parameter.requires_grad = True
In [21]:
trainable_parameters = [
name
for name, parameter in model.named_parameters()
if parameter.requires_grad
]
print("可训练参数前10项:")
print(trainable_parameters[:10])
print("可训练参数总数:", len(trainable_parameters))
可训练参数前10项: ['layer4.0.conv1.weight', 'layer4.0.bn1.weight', 'layer4.0.bn1.bias', 'layer4.0.conv2.weight', 'layer4.0.bn2.weight', 'layer4.0.bn2.bias', 'layer4.0.conv3.weight', 'layer4.0.bn3.weight', 'layer4.0.bn3.bias', 'layer4.0.downsample.0.weight'] 可训练参数总数: 32
In [22]:
optimizer = torch.optim.Adam(
[
{
"params": model.layer4.parameters(),
"lr": 1e-5
},
{
"params": model.fc.parameters(),
"lr": 1e-4
}
],
weight_decay=1e-4
)
In [23]:
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
optimizer,
mode="max",
factor=0.5,
patience=2
)
In [24]:
def evaluate_model(
model,
data_loader,
criterion,
device
):
model.eval()
total_loss = 0.0
all_labels = []
all_predictions = []
with torch.no_grad():
for images, labels in tqdm(
data_loader,
desc="Evaluating",
leave=False
):
images = images.to(device)
labels = labels.to(device)
outputs = model(images)
loss = criterion(outputs, labels)
total_loss += loss.item() * images.size(0)
predictions = outputs.argmax(dim=1)
all_labels.extend(
labels.cpu().numpy()
)
all_predictions.extend(
predictions.cpu().numpy()
)
average_loss = total_loss / len(data_loader.dataset)
accuracy = accuracy_score(
all_labels,
all_predictions
)
balanced_accuracy = balanced_accuracy_score(
all_labels,
all_predictions
)
macro_f1 = f1_score(
all_labels,
all_predictions,
average="macro",
zero_division=0
)
return {
"loss": average_loss,
"accuracy": accuracy,
"balanced_accuracy": balanced_accuracy,
"macro_f1": macro_f1
}
In [25]:
def train_one_epoch(
model,
data_loader,
criterion,
optimizer,
device
):
model.train()
total_loss = 0.0
all_labels = []
all_predictions = []
for images, labels in tqdm(
data_loader,
desc="Training",
leave=False
):
images = images.to(device)
labels = labels.to(device)
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
total_loss += loss.item() * images.size(0)
predictions = outputs.argmax(dim=1)
all_labels.extend(
labels.detach().cpu().numpy()
)
all_predictions.extend(
predictions.detach().cpu().numpy()
)
average_loss = total_loss / len(data_loader.dataset)
accuracy = accuracy_score(
all_labels,
all_predictions
)
balanced_accuracy = balanced_accuracy_score(
all_labels,
all_predictions
)
macro_f1 = f1_score(
all_labels,
all_predictions,
average="macro",
zero_division=0
)
return {
"loss": average_loss,
"accuracy": accuracy,
"balanced_accuracy": balanced_accuracy,
"macro_f1": macro_f1
}
In [26]:
num_finetune_epochs = 10
early_stopping_patience = 4
best_val_macro_f1 = -1.0
epochs_without_improvement = 0
finetune_history = {
"train_loss": [],
"train_accuracy": [],
"train_balanced_accuracy": [],
"train_macro_f1": [],
"val_loss": [],
"val_accuracy": [],
"val_balanced_accuracy": [],
"val_macro_f1": []
}
In [27]:
for epoch in range(num_finetune_epochs):
train_metrics = train_one_epoch(
model,
train_loader,
criterion,
optimizer,
device
)
val_metrics = evaluate_model(
model,
val_loader,
criterion,
device
)
finetune_history["train_loss"].append(
train_metrics["loss"]
)
finetune_history["train_accuracy"].append(
train_metrics["accuracy"]
)
finetune_history["train_balanced_accuracy"].append(
train_metrics["balanced_accuracy"]
)
finetune_history["train_macro_f1"].append(
train_metrics["macro_f1"]
)
finetune_history["val_loss"].append(
val_metrics["loss"]
)
finetune_history["val_accuracy"].append(
val_metrics["accuracy"]
)
finetune_history["val_balanced_accuracy"].append(
val_metrics["balanced_accuracy"]
)
finetune_history["val_macro_f1"].append(
val_metrics["macro_f1"]
)
layer4_lr = optimizer.param_groups[0]["lr"]
fc_lr = optimizer.param_groups[1]["lr"]
print(f"\nFine-tune Epoch {epoch + 1}/{num_finetune_epochs}")
print(
f"Train Loss: {train_metrics['loss']:.4f} | "
f"Accuracy: {train_metrics['accuracy']:.4f} | "
f"Balanced Accuracy: "
f"{train_metrics['balanced_accuracy']:.4f} | "
f"Macro-F1: {train_metrics['macro_f1']:.4f}"
)
print(
f"Val Loss: {val_metrics['loss']:.4f} | "
f"Accuracy: {val_metrics['accuracy']:.4f} | "
f"Balanced Accuracy: "
f"{val_metrics['balanced_accuracy']:.4f} | "
f"Macro-F1: {val_metrics['macro_f1']:.4f}"
)
print(
f"Learning Rate | "
f"layer4: {layer4_lr:.8f} | "
f"fc: {fc_lr:.8f}"
)
scheduler.step(
val_metrics["macro_f1"]
)
if val_metrics["macro_f1"] > best_val_macro_f1:
best_val_macro_f1 = val_metrics["macro_f1"]
epochs_without_improvement = 0
torch.save(
model.state_dict(),
best_finetune_path
)
print("已保存新的 Fine-tuned 最佳模型")
else:
epochs_without_improvement += 1
print(
"验证集 Macro-F1 未提升,"
f"连续 {epochs_without_improvement} 轮"
)
if epochs_without_improvement >= early_stopping_patience:
print("触发 Early Stopping,结束微调")
break
#保存训练记录
finetune_history_df = pd.DataFrame(
finetune_history
)
finetune_history_path = os.path.join(
OUTPUT_DIR,
"resnet50_finetune_history.csv"
)
finetune_history_df.to_csv(
finetune_history_path,
index=False
)
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 1/10 Train Loss: 0.7481 | Accuracy: 0.7384 | Balanced Accuracy: 0.7445 | Macro-F1: 0.6350 Val Loss: 0.8641 | Accuracy: 0.6775 | Balanced Accuracy: 0.6068 | Macro-F1: 0.5127 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 已保存新的 Fine-tuned 最佳模型
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 2/10 Train Loss: 0.6927 | Accuracy: 0.7495 | Balanced Accuracy: 0.7599 | Macro-F1: 0.6406 Val Loss: 0.8572 | Accuracy: 0.6821 | Balanced Accuracy: 0.6180 | Macro-F1: 0.5143 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 已保存新的 Fine-tuned 最佳模型
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 3/10 Train Loss: 0.6311 | Accuracy: 0.7615 | Balanced Accuracy: 0.7922 | Macro-F1: 0.6736 Val Loss: 0.7926 | Accuracy: 0.7043 | Balanced Accuracy: 0.6249 | Macro-F1: 0.5389 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 已保存新的 Fine-tuned 最佳模型
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 4/10 Train Loss: 0.5917 | Accuracy: 0.7702 | Balanced Accuracy: 0.8160 | Macro-F1: 0.6885 Val Loss: 0.7952 | Accuracy: 0.6932 | Balanced Accuracy: 0.6142 | Macro-F1: 0.5373 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 验证集 Macro-F1 未提升,连续 1 轮
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 5/10 Train Loss: 0.5443 | Accuracy: 0.7808 | Balanced Accuracy: 0.8238 | Macro-F1: 0.7089 Val Loss: 0.7579 | Accuracy: 0.7167 | Balanced Accuracy: 0.6170 | Macro-F1: 0.5599 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 已保存新的 Fine-tuned 最佳模型
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 6/10 Train Loss: 0.5308 | Accuracy: 0.7873 | Balanced Accuracy: 0.8183 | Macro-F1: 0.7098 Val Loss: 0.7753 | Accuracy: 0.7134 | Balanced Accuracy: 0.6475 | Macro-F1: 0.5550 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 验证集 Macro-F1 未提升,连续 1 轮
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 7/10 Train Loss: 0.5055 | Accuracy: 0.7956 | Balanced Accuracy: 0.8370 | Macro-F1: 0.7307 Val Loss: 0.7939 | Accuracy: 0.6971 | Balanced Accuracy: 0.6259 | Macro-F1: 0.5554 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 验证集 Macro-F1 未提升,连续 2 轮
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 8/10 Train Loss: 0.4569 | Accuracy: 0.8132 | Balanced Accuracy: 0.8611 | Macro-F1: 0.7644 Val Loss: 0.7504 | Accuracy: 0.7298 | Balanced Accuracy: 0.6289 | Macro-F1: 0.5803 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 已保存新的 Fine-tuned 最佳模型
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 9/10 Train Loss: 0.4612 | Accuracy: 0.8062 | Balanced Accuracy: 0.8569 | Macro-F1: 0.7550 Val Loss: 0.7291 | Accuracy: 0.7317 | Balanced Accuracy: 0.6184 | Macro-F1: 0.5707 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 验证集 Macro-F1 未提升,连续 1 轮
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 10/10 Train Loss: 0.4316 | Accuracy: 0.8191 | Balanced Accuracy: 0.8623 | Macro-F1: 0.7677 Val Loss: 0.7259 | Accuracy: 0.7356 | Balanced Accuracy: 0.6266 | Macro-F1: 0.5908 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 已保存新的 Fine-tuned 最佳模型
In [28]:
plt.figure(figsize=(8, 5))
plt.plot(
finetune_history_df["train_macro_f1"],
marker="o",
label="Train Macro-F1"
)
plt.plot(
finetune_history_df["val_macro_f1"],
marker="o",
label="Validation Macro-F1"
)
plt.xlabel("Epoch")
plt.ylabel("Macro-F1")
plt.title("ResNet50 Fine-tuning Macro-F1")
plt.legend()
plt.grid(True)
plt.show()
In [29]:
plt.figure(figsize=(8, 5))
plt.plot(
finetune_history_df["train_loss"],
marker="o",
label="Train Loss"
)
plt.plot(
finetune_history_df["val_loss"],
marker="o",
label="Validation Loss"
)
plt.xlabel("Epoch")
plt.ylabel("Loss")
plt.title("ResNet50 Fine-tuning Loss")
plt.legend()
plt.grid(True)
plt.show()
In [30]:
additional_epochs = 5
early_stopping_patience = 4
# 当前最好成绩来自前10轮
best_val_macro_f1 = 0.5908
epochs_without_improvement = 0
for extra_epoch in range(additional_epochs):
epoch_number = 11 + extra_epoch
train_metrics = train_one_epoch(
model,
train_loader,
criterion,
optimizer,
device
)
val_metrics = evaluate_model(
model,
val_loader,
criterion,
device
)
layer4_lr = optimizer.param_groups[0]["lr"]
fc_lr = optimizer.param_groups[1]["lr"]
print(f"\nFine-tune Epoch {epoch_number}/15")
print(
f"Train Loss: {train_metrics['loss']:.4f} | "
f"Accuracy: {train_metrics['accuracy']:.4f} | "
f"Balanced Accuracy: "
f"{train_metrics['balanced_accuracy']:.4f} | "
f"Macro-F1: {train_metrics['macro_f1']:.4f}"
)
print(
f"Val Loss: {val_metrics['loss']:.4f} | "
f"Accuracy: {val_metrics['accuracy']:.4f} | "
f"Balanced Accuracy: "
f"{val_metrics['balanced_accuracy']:.4f} | "
f"Macro-F1: {val_metrics['macro_f1']:.4f}"
)
print(
f"Learning Rate | "
f"layer4: {layer4_lr:.8f} | "
f"fc: {fc_lr:.8f}"
)
# 先根据当前结果更新学习率调度器
scheduler.step(val_metrics["macro_f1"])
if val_metrics["macro_f1"] > best_val_macro_f1:
best_val_macro_f1 = val_metrics["macro_f1"]
epochs_without_improvement = 0
torch.save(
model.state_dict(),
best_finetune_path
)
print("已保存新的 Fine-tuned 最佳模型")
else:
epochs_without_improvement += 1
print(
"验证集 Macro-F1 未提升,"
f"连续 {epochs_without_improvement} 轮"
)
if epochs_without_improvement >= early_stopping_patience:
print("触发 Early Stopping,结束微调")
break
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 11/15 Train Loss: 0.3894 | Accuracy: 0.8335 | Balanced Accuracy: 0.8807 | Macro-F1: 0.7984 Val Loss: 0.7231 | Accuracy: 0.7474 | Balanced Accuracy: 0.6374 | Macro-F1: 0.5935 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 已保存新的 Fine-tuned 最佳模型
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 12/15 Train Loss: 0.3917 | Accuracy: 0.8369 | Balanced Accuracy: 0.8807 | Macro-F1: 0.7980 Val Loss: 0.7549 | Accuracy: 0.7258 | Balanced Accuracy: 0.6521 | Macro-F1: 0.5922 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 验证集 Macro-F1 未提升,连续 1 轮
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 13/15 Train Loss: 0.3697 | Accuracy: 0.8342 | Balanced Accuracy: 0.8854 | Macro-F1: 0.7948 Val Loss: 0.7429 | Accuracy: 0.7258 | Balanced Accuracy: 0.6160 | Macro-F1: 0.5755 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 验证集 Macro-F1 未提升,连续 2 轮
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 14/15 Train Loss: 0.3451 | Accuracy: 0.8428 | Balanced Accuracy: 0.8942 | Macro-F1: 0.8111 Val Loss: 0.7656 | Accuracy: 0.7278 | Balanced Accuracy: 0.6389 | Macro-F1: 0.5728 Learning Rate | layer4: 0.00001000 | fc: 0.00010000 验证集 Macro-F1 未提升,连续 3 轮
Training: 0%| | 0/438 [00:00<?, ?it/s]
Evaluating: 0%| | 0/96 [00:00<?, ?it/s]
Fine-tune Epoch 15/15 Train Loss: 0.3165 | Accuracy: 0.8489 | Balanced Accuracy: 0.9080 | Macro-F1: 0.8242 Val Loss: 0.7164 | Accuracy: 0.7520 | Balanced Accuracy: 0.6453 | Macro-F1: 0.6095 Learning Rate | layer4: 0.00000500 | fc: 0.00005000 已保存新的 Fine-tuned 最佳模型
15epochs训练完毕,下面正式开始调用测试集¶
In [32]:
model.load_state_dict(
torch.load(
best_finetune_path,
map_location=device
)
)
model = model.to(device)
model.eval()
print("Best fine-tuned model loaded")
Best fine-tuned model loaded
In [33]:
test_metrics = evaluate_model(
model,
test_loader,
criterion,
device
)
print(
"Test Loss:",
test_metrics["loss"]
)
print(
"Test Accuracy:",
test_metrics["accuracy"]
)
print(
"Test Balanced Accuracy:",
test_metrics["balanced_accuracy"]
)
print(
"Test Macro-F1:",
test_metrics["macro_f1"]
)
Evaluating: 0%| | 0/93 [00:00<?, ?it/s]
Test Loss: 0.6434045841742576 Test Accuracy: 0.7771775827143822 Test Balanced Accuracy: 0.6372819453123182 Test Macro-F1: 0.6083234281767498
In [34]:
from sklearn.metrics import classification_report
all_labels = []
all_predictions = []
model.eval()
with torch.no_grad():
for images, labels in test_loader:
images = images.to(device)
outputs = model(images)
predictions = outputs.argmax(1)
all_labels.extend(
labels.numpy()
)
all_predictions.extend(
predictions.cpu().numpy()
)
report = classification_report(
all_labels,
all_predictions,
target_names=class_names,
digits=4,
zero_division=0
)
print(report)
precision recall f1-score support
akiec 0.4286 0.5870 0.4954 46
bcc 0.5976 0.6901 0.6405 71
bkl 0.5795 0.6726 0.6226 168
df 0.4286 0.4500 0.4390 20
mel 0.4913 0.5152 0.5030 165
nv 0.9223 0.8619 0.8911 992
vasc 0.6500 0.6842 0.6667 19
accuracy 0.7772 1481
macro avg 0.5854 0.6373 0.6083 1481
weighted avg 0.7944 0.7772 0.7841 1481
-------------------------------------------------------------¶
In [36]:
print(df["localization"].unique())
print(df["sex"].unique())
print(df["age"].isnull().sum())
['scalp' 'ear' 'face' 'back' 'trunk' 'chest' 'upper extremity' 'abdomen' 'unknown' 'lower extremity' 'genital' 'neck' 'hand' 'foot' 'acral'] ['male' 'female' 'unknown'] 57
In [37]:
from sklearn.preprocessing import OneHotEncoder
In [38]:
sex_categories = [
["male"],
["female"],
["unknown"]
]
localization_categories = [
[
'scalp',
'ear',
'face',
'back',
'trunk',
'chest',
'upper extremity',
'abdomen',
'unknown',
'lower extremity',
'genital',
'neck',
'hand',
'foot',
'acral'
]
]
In [39]:
import numpy as np
def process_metadata(row):
features = []
# age
age = row["age"]
if pd.isna(age):
age = train_df["age"].mean()
features.append(
age / 100
)
# sex
sex_map = {
"male":[1,0,0],
"female":[0,1,0],
"unknown":[0,0,1]
}
features.extend(
sex_map[row["sex"]]
)
# localization
locations = [
'scalp',
'ear',
'face',
'back',
'trunk',
'chest',
'upper extremity',
'abdomen',
'unknown',
'lower extremity',
'genital',
'neck',
'hand',
'foot',
'acral'
]
loc_vector = [
0
] * len(locations)
loc_vector[
locations.index(
row["localization"]
)
] = 1
features.extend(
loc_vector
)
return np.array(
features,
dtype=np.float32
)
In [41]:
class HAMMetadataDataset(Dataset):
def __init__(
self,
dataframe,
image_dir1,
image_dir2,
transform=None
):
self.df = dataframe.reset_index(
drop=True
)
self.image_dir1 = image_dir1
self.image_dir2 = image_dir2
self.transform = transform
def __len__(self):
return len(self.df)
def __getitem__(
self,
idx
):
row = self.df.iloc[idx]
image_id = row["image_id"]
filename = image_id + ".jpg"
path1 = os.path.join(
self.image_dir1,
filename
)
path2 = os.path.join(
self.image_dir2,
filename
)
if os.path.exists(path1):
image_path = path1
else:
image_path = path2
image = Image.open(
image_path
).convert("RGB")
if self.transform:
image = self.transform(image)
metadata = process_metadata(
row
)
metadata = torch.tensor(
metadata
)
label = label_map[
row["dx"]
]
return (
image,
metadata,
label
)
In [42]:
train_metadata_dataset = HAMMetadataDataset(
train_df,
IMAGE_DIR1,
IMAGE_DIR2,
train_transform
)
val_metadata_dataset = HAMMetadataDataset(
val_df,
IMAGE_DIR1,
IMAGE_DIR2,
val_transform
)
test_metadata_dataset = HAMMetadataDataset(
test_df,
IMAGE_DIR1,
IMAGE_DIR2,
val_transform
)
In [44]:
train_metadata_loader = DataLoader(
train_metadata_dataset,
batch_size=16,
shuffle=True,
num_workers=0
)
val_metadata_loader = DataLoader(
val_metadata_dataset,
batch_size=16,
shuffle=False,
num_workers=0
)
test_metadata_loader = DataLoader(
test_metadata_dataset,
batch_size=16,
shuffle=False,
num_workers=0
)
In [46]:
images, metadata, labels = next(
iter(train_metadata_loader)
)
print(images.shape)
print(metadata.shape)
print(labels.shape)
torch.Size([16, 3, 224, 224]) torch.Size([16, 19]) torch.Size([16])
In [47]:
class ResNet50FeatureExtractor(nn.Module):
def __init__(self, backbone):
super().__init__()
self.features = nn.Sequential(
*list(backbone.children())[:-1]
)
def forward(self, x):
x = self.features(x)
x = torch.flatten(
x,
1
)
return x
In [48]:
image_model = resnet50(
weights=None
)
image_model.fc = nn.Linear(
image_model.fc.in_features,
7
)
image_model.load_state_dict(
torch.load(
best_finetune_path,
map_location=device
)
)
image_model = image_model.to(device)
image_model.eval()
print("Image backbone loaded")
Image backbone loaded
In [ ]:
feature_extractor = ResNet50FeatureExtractor(
image_model
)
feature_extractor = feature_extractor.to(device)
for param in feature_extractor.parameters():
param.requires_grad = False
#注意:这里冻结。因为 Baseline 2 的目的:测试 metadata 是否提供额外信息。不是重新训练视觉模型。
In [50]:
class ImageMetadataModel(nn.Module):
def __init__(self):
super().__init__()
self.image_encoder = feature_extractor
self.classifier = nn.Sequential(
nn.Linear(
2048 + 19,
512
),
nn.ReLU(),
nn.Dropout(0.3),
nn.Linear(
512,
128
),
nn.ReLU(),
nn.Linear(
128,
7
)
)
def forward(
self,
image,
metadata
):
image_feature = self.image_encoder(
image
)
fused_feature = torch.cat(
[
image_feature,
metadata
],
dim=1
)
output = self.classifier(
fused_feature
)
return output
In [51]:
metadata_model = ImageMetadataModel()
metadata_model = metadata_model.to(device)
print(metadata_model)
ImageMetadataModel(
(image_encoder): ResNet50FeatureExtractor(
(features): Sequential(
(0): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
(1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(2): ReLU(inplace=True)
(3): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
(4): Sequential(
(0): Bottleneck(
(conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
(downsample): Sequential(
(0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
(1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
)
)
(1): Bottleneck(
(conv1): Conv2d(256, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
)
(2): Bottleneck(
(conv1): Conv2d(256, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
)
)
(5): Sequential(
(0): Bottleneck(
(conv1): Conv2d(256, 128, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(128, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
(downsample): Sequential(
(0): Conv2d(256, 512, kernel_size=(1, 1), stride=(2, 2), bias=False)
(1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
)
)
(1): Bottleneck(
(conv1): Conv2d(512, 128, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(128, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
)
(2): Bottleneck(
(conv1): Conv2d(512, 128, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(128, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
)
(3): Bottleneck(
(conv1): Conv2d(512, 128, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(128, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
)
)
(6): Sequential(
(0): Bottleneck(
(conv1): Conv2d(512, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
(downsample): Sequential(
(0): Conv2d(512, 1024, kernel_size=(1, 1), stride=(2, 2), bias=False)
(1): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
)
)
(1): Bottleneck(
(conv1): Conv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
)
(2): Bottleneck(
(conv1): Conv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
)
(3): Bottleneck(
(conv1): Conv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
)
(4): Bottleneck(
(conv1): Conv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
)
(5): Bottleneck(
(conv1): Conv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
)
)
(7): Sequential(
(0): Bottleneck(
(conv1): Conv2d(1024, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
(downsample): Sequential(
(0): Conv2d(1024, 2048, kernel_size=(1, 1), stride=(2, 2), bias=False)
(1): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
)
)
(1): Bottleneck(
(conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
)
(2): Bottleneck(
(conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
(bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
(bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, bias=True, track_running_stats=True)
(relu): ReLU(inplace=True)
)
)
(8): AdaptiveAvgPool2d(output_size=(1, 1))
)
)
(classifier): Sequential(
(0): Linear(in_features=2067, out_features=512, bias=True)
(1): ReLU()
(2): Dropout(p=0.3, inplace=False)
(3): Linear(in_features=512, out_features=128, bias=True)
(4): ReLU()
(5): Linear(in_features=128, out_features=7, bias=True)
)
)
In [52]:
images, metadata, labels = next(
iter(train_metadata_loader)
)
images = images.to(device)
metadata = metadata.to(device)
outputs = metadata_model(
images,
metadata
)
print(outputs.shape)
torch.Size([16, 7])
In [53]:
criterion_metadata = nn.CrossEntropyLoss(
weight=class_weights
)
In [54]:
optimizer_metadata = torch.optim.Adam(
metadata_model.classifier.parameters(),
lr=1e-4
)
In [55]:
def train_metadata_epoch():
metadata_model.train()
total_loss = 0
preds=[]
truths=[]
for images, metadata, labels in tqdm(
train_metadata_loader
):
images = images.to(device)
metadata = metadata.to(device)
labels = labels.to(device)
optimizer_metadata.zero_grad()
outputs = metadata_model(
images,
metadata
)
loss = criterion_metadata(
outputs,
labels
)
loss.backward()
optimizer_metadata.step()
total_loss += (
loss.item()
*
images.size(0)
)
preds.extend(
outputs.argmax(1)
.detach()
.cpu()
.numpy()
)
truths.extend(
labels.cpu()
.numpy()
)
return {
"loss":
total_loss /
len(train_metadata_loader.dataset),
"accuracy":
accuracy_score(
truths,
preds
),
"macro_f1":
f1_score(
truths,
preds,
average="macro"
)
}
In [56]:
def evaluate_metadata_model():
metadata_model.eval()
total_loss = 0
preds = []
truths = []
with torch.no_grad():
for images, metadata, labels in tqdm(
val_metadata_loader
):
images = images.to(device)
metadata = metadata.to(device)
labels = labels.to(device)
outputs = metadata_model(
images,
metadata
)
loss = criterion_metadata(
outputs,
labels
)
total_loss += (
loss.item()
*
images.size(0)
)
predictions = outputs.argmax(
dim=1
)
preds.extend(
predictions.cpu().numpy()
)
truths.extend(
labels.cpu().numpy()
)
return {
"loss":
total_loss /
len(val_metadata_loader.dataset),
"accuracy":
accuracy_score(
truths,
preds
),
"balanced_accuracy":
balanced_accuracy_score(
truths,
preds
),
"macro_f1":
f1_score(
truths,
preds,
average="macro",
zero_division=0
)
}
In [57]:
# 加入元数据开始训练20个回合
num_epochs = 20
best_metadata_f1 = 0
metadata_history = {
"train_loss": [],
"val_loss": [],
"train_f1": [],
"val_f1": [],
"val_accuracy": [],
"val_balanced_accuracy": []
}
for epoch in range(num_epochs):
train_metrics = train_metadata_epoch()
val_metrics = evaluate_metadata_model()
metadata_history["train_loss"].append(
train_metrics["loss"]
)
metadata_history["val_loss"].append(
val_metrics["loss"]
)
metadata_history["train_f1"].append(
train_metrics["macro_f1"]
)
metadata_history["val_f1"].append(
val_metrics["macro_f1"]
)
metadata_history["val_accuracy"].append(
val_metrics["accuracy"]
)
metadata_history["val_balanced_accuracy"].append(
val_metrics["balanced_accuracy"]
)
print(
f"""
Epoch {epoch+1}/{num_epochs}
Train:
Loss:
{train_metrics['loss']:.4f}
Accuracy:
{train_metrics['accuracy']:.4f}
Macro-F1:
{train_metrics['macro_f1']:.4f}
Validation:
Loss:
{val_metrics['loss']:.4f}
Accuracy:
{val_metrics['accuracy']:.4f}
Balanced Accuracy:
{val_metrics['balanced_accuracy']:.4f}
Macro-F1:
{val_metrics['macro_f1']:.4f}
"""
)
if val_metrics["macro_f1"] > best_metadata_f1:
best_metadata_f1 = (
val_metrics["macro_f1"]
)
torch.save(
metadata_model.state_dict(),
os.path.join(
CHECKPOINT_DIR,
"image_all_metadata_best.pth"
)
)
print(
"保存最佳 Image+Metadata 模型"
)
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 1/20 Train: Loss: 1.4755 Accuracy: 0.6404 Macro-F1: 0.4105 Validation: Loss: 1.1038 Accuracy: 0.6253 Balanced Accuracy: 0.5612 Macro-F1: 0.4662 保存最佳 Image+Metadata 模型
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 2/20 Train: Loss: 0.8758 Accuracy: 0.7128 Macro-F1: 0.5992 Validation: Loss: 0.8069 Accuracy: 0.7076 Balanced Accuracy: 0.6288 Macro-F1: 0.5200 保存最佳 Image+Metadata 模型
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 3/20 Train: Loss: 0.6682 Accuracy: 0.7579 Macro-F1: 0.6642 Validation: Loss: 0.8659 Accuracy: 0.6821 Balanced Accuracy: 0.6428 Macro-F1: 0.5273 保存最佳 Image+Metadata 模型
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 4/20 Train: Loss: 0.5594 Accuracy: 0.7828 Macro-F1: 0.7096 Validation: Loss: 0.7521 Accuracy: 0.7245 Balanced Accuracy: 0.5905 Macro-F1: 0.5551 保存最佳 Image+Metadata 模型
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 5/20 Train: Loss: 0.5483 Accuracy: 0.7749 Macro-F1: 0.7025 Validation: Loss: 0.8503 Accuracy: 0.7063 Balanced Accuracy: 0.6255 Macro-F1: 0.5297
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 6/20 Train: Loss: 0.4810 Accuracy: 0.7866 Macro-F1: 0.7263 Validation: Loss: 0.7575 Accuracy: 0.7063 Balanced Accuracy: 0.6266 Macro-F1: 0.5670 保存最佳 Image+Metadata 模型
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 7/20 Train: Loss: 0.4559 Accuracy: 0.7961 Macro-F1: 0.7418 Validation: Loss: 0.8354 Accuracy: 0.6854 Balanced Accuracy: 0.6310 Macro-F1: 0.5624
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 8/20 Train: Loss: 0.4330 Accuracy: 0.8081 Macro-F1: 0.7520 Validation: Loss: 0.7605 Accuracy: 0.7298 Balanced Accuracy: 0.6287 Macro-F1: 0.5805 保存最佳 Image+Metadata 模型
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 9/20 Train: Loss: 0.4050 Accuracy: 0.8163 Macro-F1: 0.7724 Validation: Loss: 0.7413 Accuracy: 0.7474 Balanced Accuracy: 0.6293 Macro-F1: 0.5766
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 10/20 Train: Loss: 0.3859 Accuracy: 0.8193 Macro-F1: 0.7761 Validation: Loss: 0.8207 Accuracy: 0.6919 Balanced Accuracy: 0.6241 Macro-F1: 0.5649
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 11/20 Train: Loss: 0.3832 Accuracy: 0.8266 Macro-F1: 0.7749 Validation: Loss: 0.7330 Accuracy: 0.7480 Balanced Accuracy: 0.6399 Macro-F1: 0.5955 保存最佳 Image+Metadata 模型
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 12/20 Train: Loss: 0.3681 Accuracy: 0.8236 Macro-F1: 0.7903 Validation: Loss: 0.9029 Accuracy: 0.6691 Balanced Accuracy: 0.6352 Macro-F1: 0.5756
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 13/20 Train: Loss: 0.3554 Accuracy: 0.8332 Macro-F1: 0.7919 Validation: Loss: 0.8380 Accuracy: 0.7134 Balanced Accuracy: 0.6169 Macro-F1: 0.5702
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 14/20 Train: Loss: 0.3526 Accuracy: 0.8318 Macro-F1: 0.8021 Validation: Loss: 0.9344 Accuracy: 0.6867 Balanced Accuracy: 0.6025 Macro-F1: 0.5411
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 15/20 Train: Loss: 0.3653 Accuracy: 0.8369 Macro-F1: 0.7941 Validation: Loss: 0.7728 Accuracy: 0.7467 Balanced Accuracy: 0.6305 Macro-F1: 0.5903
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 16/20 Train: Loss: 0.3520 Accuracy: 0.8399 Macro-F1: 0.8017 Validation: Loss: 0.8160 Accuracy: 0.7134 Balanced Accuracy: 0.6215 Macro-F1: 0.5713
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 17/20 Train: Loss: 0.3493 Accuracy: 0.8422 Macro-F1: 0.8119 Validation: Loss: 0.7642 Accuracy: 0.7383 Balanced Accuracy: 0.6248 Macro-F1: 0.5913
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 18/20 Train: Loss: 0.3236 Accuracy: 0.8480 Macro-F1: 0.8073 Validation: Loss: 0.7502 Accuracy: 0.7631 Balanced Accuracy: 0.6146 Macro-F1: 0.6028 保存最佳 Image+Metadata 模型
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 19/20 Train: Loss: 0.3163 Accuracy: 0.8425 Macro-F1: 0.8050 Validation: Loss: 0.8062 Accuracy: 0.7448 Balanced Accuracy: 0.6030 Macro-F1: 0.5813
0%| | 0/438 [00:00<?, ?it/s]
0%| | 0/96 [00:00<?, ?it/s]
Epoch 20/20 Train: Loss: 0.3249 Accuracy: 0.8478 Macro-F1: 0.8156 Validation: Loss: 0.7941 Accuracy: 0.7428 Balanced Accuracy: 0.6319 Macro-F1: 0.5830
In [58]:
import os
print(
os.path.exists(
os.path.join(
CHECKPOINT_DIR,
"image_all_metadata_best.pth"
)
)
)
True
In [59]:
metadata_model.load_state_dict(
torch.load(
os.path.join(
CHECKPOINT_DIR,
"image_all_metadata_best.pth"
),
map_location=device
)
)
metadata_model.eval()
print("Best metadata model loaded")
Best metadata model loaded
In [60]:
def evaluate_metadata_test():
metadata_model.eval()
preds=[]
truths=[]
total_loss=0
with torch.no_grad():
for images, metadata, labels in tqdm(
test_metadata_loader
):
images=images.to(device)
metadata=metadata.to(device)
labels=labels.to(device)
outputs=metadata_model(
images,
metadata
)
loss=criterion_metadata(
outputs,
labels
)
total_loss += (
loss.item()
*
images.size(0)
)
preds.extend(
outputs.argmax(1)
.cpu()
.numpy()
)
truths.extend(
labels.cpu()
.numpy()
)
return {
"loss":
total_loss /
len(test_metadata_loader.dataset),
"accuracy":
accuracy_score(
truths,
preds
),
"balanced_accuracy":
balanced_accuracy_score(
truths,
preds
),
"macro_f1":
f1_score(
truths,
preds,
average="macro"
),
"truths":
truths,
"preds":
preds
}
In [61]:
metadata_test_result = evaluate_metadata_test()
print(metadata_test_result)
0%| | 0/93 [00:00<?, ?it/s]
{'loss': 0.678672323963032, 'accuracy': 0.7798784604996624, 'balanced_accuracy': 0.6135215876495631, 'macro_f1': 0.5845429210957046, 'truths': [np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(3), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0)], 'preds': [np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(5), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(1), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(5), np.int64(2), np.int64(1), np.int64(2), np.int64(4), np.int64(2), np.int64(5), np.int64(5), np.int64(2), np.int64(2), np.int64(5), np.int64(2), np.int64(5), np.int64(2), np.int64(2), np.int64(2), np.int64(5), np.int64(5), np.int64(2), np.int64(2), np.int64(5), np.int64(2), np.int64(1), np.int64(4), np.int64(2), np.int64(5), np.int64(4), np.int64(2), np.int64(2), np.int64(2), np.int64(1), np.int64(4), np.int64(2), np.int64(4), np.int64(4), np.int64(4), np.int64(0), np.int64(2), np.int64(2), np.int64(4), np.int64(2), np.int64(0), np.int64(2), np.int64(2), np.int64(5), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(0), np.int64(0), np.int64(1), np.int64(2), np.int64(0), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(4), np.int64(2), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(0), np.int64(2), np.int64(2), np.int64(2), np.int64(4), np.int64(2), np.int64(2), np.int64(5), np.int64(4), np.int64(5), np.int64(0), np.int64(0), np.int64(4), np.int64(1), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(0), np.int64(4), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(5), np.int64(2), np.int64(3), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(3), np.int64(3), np.int64(2), np.int64(2), np.int64(5), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(5), np.int64(1), np.int64(3), np.int64(5), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(2), np.int64(3), np.int64(1), np.int64(1), np.int64(3), np.int64(0), np.int64(1), np.int64(3), np.int64(5), np.int64(0), np.int64(0), np.int64(3), np.int64(0), np.int64(0), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(2), np.int64(3), np.int64(3), np.int64(4), np.int64(4), np.int64(4), np.int64(0), np.int64(3), np.int64(1), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(5), np.int64(2), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(3), np.int64(5), np.int64(2), np.int64(2), np.int64(4), np.int64(4), np.int64(5), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(5), np.int64(4), np.int64(5), np.int64(0), np.int64(4), np.int64(2), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(5), np.int64(3), np.int64(4), np.int64(2), np.int64(4), np.int64(4), np.int64(4), np.int64(2), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(0), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(2), np.int64(4), np.int64(0), np.int64(4), np.int64(4), np.int64(4), np.int64(0), np.int64(0), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(0), np.int64(4), np.int64(5), np.int64(4), np.int64(5), np.int64(2), np.int64(5), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(5), np.int64(2), np.int64(5), np.int64(2), np.int64(2), np.int64(2), np.int64(4), np.int64(4), np.int64(4), np.int64(5), np.int64(2), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(5), np.int64(4), np.int64(4), np.int64(5), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(0), np.int64(2), np.int64(2), np.int64(4), np.int64(2), np.int64(0), np.int64(5), np.int64(2), np.int64(2), np.int64(4), np.int64(0), np.int64(4), np.int64(4), np.int64(2), np.int64(2), np.int64(0), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(2), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(2), np.int64(2), np.int64(2), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(2), np.int64(2), np.int64(4), np.int64(0), np.int64(2), np.int64(4), np.int64(4), np.int64(4), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(1), np.int64(6), np.int64(4), np.int64(4), np.int64(5), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(6), np.int64(5), np.int64(6), np.int64(6), np.int64(5), np.int64(6), np.int64(1), np.int64(0), np.int64(1), np.int64(1), np.int64(1), np.int64(2), np.int64(2), np.int64(0), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(2), np.int64(1), np.int64(4), np.int64(4), np.int64(5), np.int64(6), np.int64(4), np.int64(5), np.int64(5), np.int64(1), np.int64(1), np.int64(1), np.int64(5), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(0), np.int64(1), np.int64(2), np.int64(1), np.int64(0), np.int64(5), np.int64(5), np.int64(2), np.int64(1), np.int64(1), np.int64(0), np.int64(1), np.int64(1), np.int64(1), np.int64(5), np.int64(5), np.int64(1), np.int64(0), np.int64(2), np.int64(1), np.int64(1), np.int64(1), np.int64(3), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(5), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(0), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(1), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(3), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(6), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(3), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(3), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(1), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(3), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(1), np.int64(1), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(6), np.int64(2), np.int64(1), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(6), np.int64(4), np.int64(5), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(4), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(5), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(6), np.int64(6), np.int64(5), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(6), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(1), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(6), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(2), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(1), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(4), np.int64(4), np.int64(3), np.int64(3), np.int64(4), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(6), np.int64(4), np.int64(5), np.int64(2), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(1), np.int64(5), np.int64(5), np.int64(5), np.int64(1), np.int64(1), np.int64(5), np.int64(4), np.int64(2), np.int64(0), np.int64(2), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(2), np.int64(4), np.int64(2), np.int64(4), np.int64(5), np.int64(1), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(2), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(1), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(2), np.int64(1), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(4), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(2), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(4), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(5), np.int64(0), np.int64(4), np.int64(1), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(1), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(3), np.int64(1), np.int64(2), np.int64(1), np.int64(2), np.int64(0), np.int64(0), np.int64(1), np.int64(2), np.int64(2), np.int64(0), np.int64(2), np.int64(0), np.int64(1), np.int64(2), np.int64(1), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(0), np.int64(2), np.int64(0), np.int64(0), np.int64(0), np.int64(2)]}
In [62]:
from sklearn.metrics import classification_report
print(
classification_report(
metadata_test_result["truths"],
metadata_test_result["preds"],
target_names=class_names,
digits=4
)
)
precision recall f1-score support
akiec 0.4603 0.6304 0.5321 46
bcc 0.5867 0.6197 0.6027 71
bkl 0.6034 0.6429 0.6225 168
df 0.2857 0.3000 0.2927 20
mel 0.4762 0.5455 0.5085 165
nv 0.9281 0.8720 0.8992 992
vasc 0.5909 0.6842 0.6341 19
accuracy 0.7799 1481
macro avg 0.5616 0.6135 0.5845 1481
weighted avg 0.7970 0.7799 0.7871 1481
In [63]:
import pandas as pd
import os
from datetime import datetime
results = []
# ==========================
# Image Only baseline
# ==========================
results.append({
"Method": "Image Only",
"Accuracy": 0.7771775827143822,
"Balanced Accuracy": 0.6372819453123182,
"Macro-F1": 0.6083234281767498,
"Description":
"Only dermoscopic image input"
})
# ==========================
# Image + Metadata
# ==========================
results.append({
"Method": "Image + All Metadata",
"Accuracy": 0.7798784604996624,
"Balanced Accuracy": 0.6135215876495631,
"Macro-F1": 0.5845429210957046,
"Description":
"Image with age, sex and localization metadata"
})
# 转DataFrame
baseline_df = pd.DataFrame(results)
baseline_df
Out[63]:
| Method | Accuracy | Balanced Accuracy | Macro-F1 | Description | |
|---|---|---|---|---|---|
| 0 | Image Only | 0.777178 | 0.637282 | 0.608323 | Only dermoscopic image input |
| 1 | Image + All Metadata | 0.779878 | 0.613522 | 0.584543 | Image with age, sex and localization metadata |
In [64]:
save_path = "./baseline_results.csv"
baseline_df.to_csv(
save_path,
index=False,
encoding="utf-8-sig"
)
print(
"Saved:",
os.path.abspath(save_path)
)
Saved: /Users/applesues01/Documents/Medical_Agent/notebooks/baseline_results.csv
In [65]:
check = pd.read_csv(
"./baseline_results.csv"
)
print(check)
Method Accuracy Balanced Accuracy Macro-F1 \
0 Image Only 0.777178 0.637282 0.608323
1 Image + All Metadata 0.779878 0.613522 0.584543
Description
0 Only dermoscopic image input
1 Image with age, sex and localization metadata
In [66]:
import json
experiment_info = {
"dataset":
"HAM10000",
"classes":
[
"akiec",
"bcc",
"bkl",
"df",
"mel",
"nv",
"vasc"
],
"split":
{
"train":7002,
"val":1532,
"test":1481
},
"models":
{
"image_only":
{
"accuracy":0.7772,
"balanced_accuracy":0.6373,
"macro_f1":0.6083
},
"image_metadata":
{
"accuracy":0.7799,
"balanced_accuracy":0.6135,
"macro_f1":0.5845
}
},
"date":
str(datetime.now())
}
with open(
"experiment_summary.json",
"w",
encoding="utf-8"
) as f:
json.dump(
experiment_info,
f,
indent=4,
ensure_ascii=False
)
print("JSON saved")
JSON saved
In [67]:
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.metrics import confusion_matrix
In [68]:
metadata_cm = confusion_matrix(
metadata_test_result["truths"],
metadata_test_result["preds"]
)
plt.figure(figsize=(8,6))
sns.heatmap(
metadata_cm,
annot=True,
fmt="d",
cmap="Blues",
xticklabels=class_names,
yticklabels=class_names
)
plt.xlabel("Predicted")
plt.ylabel("True")
plt.title(
"Confusion Matrix - Image + Metadata"
)
plt.tight_layout()
plt.savefig(
"confusion_matrix_image_metadata.png",
dpi=300,
bbox_inches="tight"
)
plt.show()
In [72]:
import torch
import torch.nn as nn
from torchvision import models
device = torch.device(
"mps" if torch.backends.mps.is_available()
else "cpu"
)
print(device)
mps
In [73]:
image_only_model = models.resnet50(
weights=None
)
image_only_model.fc = nn.Linear(
image_only_model.fc.in_features,
7
)
image_only_model = image_only_model.to(device)
print("model created")
model created
In [77]:
checkpoint_path = "../checkpoints/resnet50_image_only_best.pth"
state_dict = torch.load(
checkpoint_path,
map_location=device
)
image_only_model.load_state_dict(
state_dict
)
image_only_model.eval()
print("Image Only best model loaded")
Image Only best model loaded
In [78]:
image_only_truths = []
image_only_preds = []
with torch.no_grad():
for images, labels in test_loader:
images = images.to(device)
labels = labels.to(device)
outputs = image_only_model(images)
preds = outputs.argmax(dim=1)
image_only_truths.extend(
labels.cpu().numpy()
)
image_only_preds.extend(
preds.cpu().numpy()
)
print(len(image_only_truths))
1481
In [79]:
image_only_cm = confusion_matrix(
image_only_truths,
image_only_preds
)
plt.figure(figsize=(8,6))
sns.heatmap(
image_only_cm,
annot=True,
fmt="d",
cmap="Blues",
xticklabels=class_names,
yticklabels=class_names
)
plt.xlabel("Predicted")
plt.ylabel("True")
plt.title(
"Confusion Matrix - Image Only"
)
plt.tight_layout()
plt.savefig(
"confusion_matrix_image_only.png",
dpi=300,
bbox_inches="tight"
)
plt.show()
In [81]:
import os
save_dir = "../results"
os.makedirs(save_dir, exist_ok=True)
print("保存目录:", os.path.abspath(save_dir))
保存目录: /Users/applesues01/Documents/Medical_Agent/results
In [82]:
import numpy as np
np.save(
os.path.join(
save_dir,
"image_only_truths.npy"
),
np.array(image_only_truths)
)
np.save(
os.path.join(
save_dir,
"image_only_preds.npy"
),
np.array(image_only_preds)
)
print("Image Only prediction saved")
Image Only prediction saved
In [83]:
np.save(
os.path.join(
save_dir,
"metadata_truths.npy"
),
np.array(
metadata_test_result["truths"]
)
)
np.save(
os.path.join(
save_dir,
"metadata_preds.npy"
),
np.array(
metadata_test_result["preds"]
)
)
print("Metadata prediction saved")
Metadata prediction saved
In [84]:
from sklearn.metrics import confusion_matrix
image_only_cm = confusion_matrix(
image_only_truths,
image_only_preds
)
metadata_cm = confusion_matrix(
metadata_test_result["truths"],
metadata_test_result["preds"]
)
np.save(
os.path.join(
save_dir,
"confusion_matrix_image_only.npy"
),
image_only_cm
)
np.save(
os.path.join(
save_dir,
"confusion_matrix_metadata.npy"
),
metadata_cm
)
print("Confusion matrix saved")
Confusion matrix saved
In [85]:
for file in os.listdir(save_dir):
print(file)
confusion_matrix_metadata.npy image_only_preds.npy metadata_preds.npy confusion_matrix_image_only.npy metadata_truths.npy image_only_truths.npy
In [87]:
test_truth = np.load(
"../results/image_only_truths.npy"
)
test_pred = np.load(
"../results/image_only_preds.npy"
)
print(test_truth.shape)
print(test_pred.shape)
(1481,) (1481,)