Only Available On Mac!¶

Practice Baseline Model¶

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()
No description has been provided for this image
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()
No description has been provided for this image
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

Unfrozen Model Was Set done, Means A ResNet50 Model with Unfrozen weight ,providing with Image only, shows a great effect, without overfitting!¶

This Model is a bottom baseline, without any metadata, Image only!¶

-------------------------------------------------------------¶

Next, we are going to provide the model with metadate, to build a top baseline, or limit.¶

Adding age,sex,localization (3factors)¶

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()
No description has been provided for this image
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()
No description has been provided for this image
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,)