In [1]:
import os
import copy
import json
import numpy as np
import pandas as pd

import torch
import torch.nn as nn

from torch.utils.data import DataLoader, Dataset
from torchvision import transforms
from torchvision.models import resnet50

from PIL import Image

from sklearn.metrics import (
    accuracy_score,
    balanced_accuracy_score,
    f1_score
)
from sklearn.utils.class_weight import compute_class_weight
In [2]:
device = torch.device(
    "mps" if torch.backends.mps.is_available()
    else "cuda" if torch.cuda.is_available()
    else "cpu"
)

print("Device:", device)

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")
Device: mps
In [3]:
train_df = pd.read_csv(os.path.join(SPLIT_DIR, "train.csv"))
val_df = pd.read_csv(os.path.join(SPLIT_DIR, "val.csv"))
test_df = pd.read_csv(os.path.join(SPLIT_DIR, "test.csv"))

print(len(train_df), len(val_df), len(test_df))
7002 1532 1481
In [4]:
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"]

LOCATIONS = [
    "scalp", "ear", "face", "back", "trunk", "chest",
    "upper extremity", "abdomen", "unknown", "lower extremity",
    "genital", "neck", "hand", "foot", "acral"
]

SEX_MAP = {
    "male": [1.0, 0.0, 0.0],
    "female": [0.0, 1.0, 0.0],
    "unknown": [0.0, 0.0, 1.0],
}

def process_metadata(row, train_age_mean):
    features = []

    age = row["age"]
    if pd.isna(age):
        age = train_age_mean
    features.append(float(age) / 100.0)

    sex_key = row["sex"] if row["sex"] in SEX_MAP else "unknown"
    features.extend(SEX_MAP[sex_key])

    loc_key = row["localization"] if row["localization"] in LOCATIONS else "unknown"
    loc_vector = [0.0] * len(LOCATIONS)
    loc_vector[LOCATIONS.index(loc_key)] = 1.0
    features.extend(loc_vector)

    return np.array(features, dtype=np.float32)

train_age_mean = train_df["age"].mean()
print("metadata dim:", len(process_metadata(train_df.iloc[0], train_age_mean)))
metadata dim: 19
In [5]:
IMAGE_SIZE = 224
BATCH_SIZE = 16

eval_transform = transforms.Compose([
    transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),
    transforms.ToTensor(),
    transforms.Normalize(
        mean=[0.485, 0.456, 0.406],
        std=[0.229, 0.224, 0.225]
    )
])
In [6]:
def resolve_image_path(image_id):
    filename = f"{image_id}.jpg"
    path1 = os.path.join(IMAGE_DIR1, filename)
    path2 = os.path.join(IMAGE_DIR2, filename)
    return path1 if os.path.exists(path1) else path2
In [7]:
class HAMImageDataset(Dataset):
    def __init__(self, dataframe, transform=None):
        self.df = dataframe.reset_index(drop=True)
        self.transform = transform

    def __len__(self):
        return len(self.df)

    def __getitem__(self, idx):
        row = self.df.iloc[idx]
        image = Image.open(resolve_image_path(row["image_id"])).convert("RGB")

        if self.transform:
            image = self.transform(image)

        label = LABEL_MAP[row["dx"]]
        return image, label
In [8]:
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 [9]:
image_only_model = resnet50(weights=None)
image_only_model.fc = nn.Linear(image_only_model.fc.in_features, 7)

checkpoint_path = os.path.join(
    CHECKPOINT_DIR,
    "resnet50_image_only_finetuned_best.pth"
)

image_only_model.load_state_dict(
    torch.load(checkpoint_path, map_location=device)
)

image_only_model = image_only_model.to(device)
image_only_model.eval()

print("Image Only checkpoint loaded")
Image Only checkpoint loaded
In [10]:
feature_extractor = ResNet50FeatureExtractor(image_only_model).to(device)
feature_extractor.eval()

for param in feature_extractor.parameters():
    param.requires_grad = False

print("Feature extractor ready")
Feature extractor ready
In [12]:
train_image_loader = DataLoader(
    HAMImageDataset(train_df, eval_transform),
    batch_size=BATCH_SIZE,
    shuffle=False,
    num_workers=0
)

val_image_loader = DataLoader(
    HAMImageDataset(val_df, eval_transform),
    batch_size=BATCH_SIZE,
    shuffle=False,
    num_workers=0
)

test_image_loader = DataLoader(
    HAMImageDataset(test_df, eval_transform),
    batch_size=BATCH_SIZE,
    shuffle=False,
    num_workers=0
)
In [13]:
def extract_image_features(data_loader, feature_extractor, device):
    feature_extractor.eval()

    all_features = []
    all_labels = []

    with torch.no_grad():
        for images, labels in data_loader:
            images = images.to(device)
            features = feature_extractor(images).cpu()

            all_features.append(features)
            all_labels.append(labels)

    return torch.cat(all_features, dim=0), torch.cat(all_labels, dim=0)
In [14]:
train_image_features, train_labels = extract_image_features(
    train_image_loader,
    feature_extractor,
    device
)

val_image_features, val_labels = extract_image_features(
    val_image_loader,
    feature_extractor,
    device
)

test_image_features, test_labels = extract_image_features(
    test_image_loader,
    feature_extractor,
    device
)

print(train_image_features.shape, train_labels.shape)
print(val_image_features.shape, val_labels.shape)
print(test_image_features.shape, test_labels.shape)
torch.Size([7002, 2048]) torch.Size([7002])
torch.Size([1532, 2048]) torch.Size([1532])
torch.Size([1481, 2048]) torch.Size([1481])
In [15]:
def process_selected_metadata(row, selected_features, train_age_mean):
    features = []

    if "age" in selected_features:
        age = row["age"]
        if pd.isna(age):
            age = train_age_mean
        features.append(float(age) / 100.0)

    if "sex" in selected_features:
        sex_key = row["sex"] if row["sex"] in SEX_MAP else "unknown"
        features.extend(SEX_MAP[sex_key])

    if "location" in selected_features:
        loc_key = row["localization"] if row["localization"] in LOCATIONS else "unknown"
        loc_vector = [0.0] * len(LOCATIONS)
        loc_vector[LOCATIONS.index(loc_key)] = 1.0
        features.extend(loc_vector)

    return np.array(features, dtype=np.float32)
In [16]:
def build_metadata_matrix(dataframe, selected_features, train_age_mean):
    rows = [
        process_selected_metadata(
            dataframe.iloc[idx],
            selected_features,
            train_age_mean
        )
        for idx in range(len(dataframe))
    ]
    return np.stack(rows, axis=0)
In [17]:
metadata_experiments = [
    {"name": "image_age", "features": ["age"]},
    {"name": "image_sex", "features": ["sex"]},
    {"name": "image_location", "features": ["location"]},
    {"name": "image_age_sex", "features": ["age", "sex"]},
    {"name": "image_age_location", "features": ["age", "location"]},
    {"name": "image_sex_location", "features": ["sex", "location"]},
    {"name": "image_all_metadata", "features": ["age", "sex", "location"]},
]
In [18]:
for exp in metadata_experiments:
    train_metadata = build_metadata_matrix(
        train_df,
        exp["features"],
        train_age_mean
    )
    print(exp["name"], train_metadata.shape)
image_age (7002, 1)
image_sex (7002, 3)
image_location (7002, 15)
image_age_sex (7002, 4)
image_age_location (7002, 16)
image_sex_location (7002, 18)
image_all_metadata (7002, 19)
In [19]:
class CachedFusionDataset(Dataset):
    def __init__(self, image_features, metadata_features, labels):
        self.image_features = image_features.float()
        self.metadata_features = torch.tensor(metadata_features, dtype=torch.float32)
        self.labels = labels.long()

    def __len__(self):
        return len(self.labels)

    def __getitem__(self, idx):
        return (
            self.image_features[idx],
            self.metadata_features[idx],
            self.labels[idx]
        )
In [20]:
class MetadataFusionClassifier(nn.Module):
    def __init__(self, metadata_dim, num_classes=7):
        super().__init__()
        self.classifier = nn.Sequential(
            nn.Linear(2048 + metadata_dim, 512),
            nn.ReLU(),
            nn.Dropout(0.3),
            nn.Linear(512, 128),
            nn.ReLU(),
            nn.Linear(128, num_classes)
        )

    def forward(self, image_features, metadata):
        fused = torch.cat([image_features, metadata], dim=1)
        return self.classifier(fused)
In [21]:
def evaluate_cached_model(model, data_loader, criterion, device):
    model.eval()

    preds = []
    truths = []
    total_loss = 0.0

    with torch.no_grad():
        for image_features, metadata, labels in data_loader:
            image_features = image_features.to(device)
            metadata = metadata.to(device)
            labels = labels.to(device)

            outputs = model(image_features, metadata)
            loss = criterion(outputs, labels)

            total_loss += loss.item() * image_features.size(0)

            preds.extend(outputs.argmax(1).cpu().numpy())
            truths.extend(labels.cpu().numpy())

    return {
        "loss": total_loss / len(data_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 [22]:
def train_cached_one_experiment(model, data_loader, criterion, optimizer, device):
    model.train()

    preds = []
    truths = []
    total_loss = 0.0

    for image_features, metadata, labels in data_loader:
        image_features = image_features.to(device)
        metadata = metadata.to(device)
        labels = labels.to(device)

        optimizer.zero_grad()

        outputs = model(image_features, metadata)
        loss = criterion(outputs, labels)

        loss.backward()
        optimizer.step()

        total_loss += loss.item() * image_features.size(0)

        preds.extend(outputs.argmax(1).detach().cpu().numpy())
        truths.extend(labels.cpu().numpy())

    return {
        "loss": total_loss / len(data_loader.dataset),
        "accuracy": accuracy_score(truths, preds),
        "macro_f1": f1_score(truths, preds, average="macro", zero_division=0),
    }
In [23]:
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 [24]:
def build_cached_loaders_for_experiment(selected_features):
    train_metadata = build_metadata_matrix(
        train_df,
        selected_features,
        train_age_mean
    )
    val_metadata = build_metadata_matrix(
        val_df,
        selected_features,
        train_age_mean
    )
    test_metadata = build_metadata_matrix(
        test_df,
        selected_features,
        train_age_mean
    )

    train_dataset = CachedFusionDataset(
        train_image_features,
        train_metadata,
        train_labels
    )
    val_dataset = CachedFusionDataset(
        val_image_features,
        val_metadata,
        val_labels
    )
    test_dataset = CachedFusionDataset(
        test_image_features,
        test_metadata,
        test_labels
    )

    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
    )

    return train_loader, val_loader, test_loader, train_metadata.shape[1]
In [25]:
def run_one_metadata_experiment(
    experiment_name,
    selected_features,
    num_epochs=20
):
    train_loader, val_loader, test_loader, metadata_dim = build_cached_loaders_for_experiment(
        selected_features
    )

    model = MetadataFusionClassifier(
        metadata_dim=metadata_dim,
        num_classes=7
    ).to(device)

    criterion = nn.CrossEntropyLoss(
        weight=class_weights
    )

    optimizer = torch.optim.Adam(
        model.parameters(),
        lr=1e-4
    )

    best_val_f1 = -1.0
    best_state = None
    history = []

    for epoch in range(num_epochs):
        train_metrics = train_cached_one_experiment(
            model,
            train_loader,
            criterion,
            optimizer,
            device
        )

        val_metrics = evaluate_cached_model(
            model,
            val_loader,
            criterion,
            device
        )

        history.append({
            "epoch": epoch + 1,
            "train_loss": train_metrics["loss"],
            "train_accuracy": train_metrics["accuracy"],
            "train_macro_f1": train_metrics["macro_f1"],
            "val_loss": val_metrics["loss"],
            "val_accuracy": val_metrics["accuracy"],
            "val_balanced_accuracy": val_metrics["balanced_accuracy"],
            "val_macro_f1": val_metrics["macro_f1"],
        })

        print(
            f"[{experiment_name}] "
            f"epoch {epoch+1}/{num_epochs} | "
            f"train_f1={train_metrics['macro_f1']:.4f} | "
            f"val_f1={val_metrics['macro_f1']:.4f} | "
            f"val_bal_acc={val_metrics['balanced_accuracy']:.4f}"
        )

        if val_metrics["macro_f1"] > best_val_f1:
            best_val_f1 = val_metrics["macro_f1"]
            best_state = copy.deepcopy(model.state_dict())

    model.load_state_dict(best_state)

    test_metrics = evaluate_cached_model(
        model,
        test_loader,
        criterion,
        device
    )

    return {
        "name": experiment_name,
        "features": selected_features,
        "metadata_dim": metadata_dim,
        "best_val_macro_f1": best_val_f1,
        "test_accuracy": test_metrics["accuracy"],
        "test_balanced_accuracy": test_metrics["balanced_accuracy"],
        "test_macro_f1": test_metrics["macro_f1"],
        "history": history
    }
In [26]:
single_result = run_one_metadata_experiment(
    experiment_name="image_age",
    selected_features=["age"],
    num_epochs=5
)

single_result
[image_age] epoch 1/5 | train_f1=0.4882 | val_f1=0.5060 | val_bal_acc=0.5606
[image_age] epoch 2/5 | train_f1=0.7032 | val_f1=0.4872 | val_bal_acc=0.5589
[image_age] epoch 3/5 | train_f1=0.7837 | val_f1=0.5508 | val_bal_acc=0.6028
[image_age] epoch 4/5 | train_f1=0.8242 | val_f1=0.5655 | val_bal_acc=0.5920
[image_age] epoch 5/5 | train_f1=0.8565 | val_f1=0.5735 | val_bal_acc=0.5980
Out[26]:
{'name': 'image_age',
 'features': ['age'],
 'metadata_dim': 1,
 'best_val_macro_f1': 0.5734997843740176,
 'test_accuracy': 0.7521944632005402,
 'test_balanced_accuracy': 0.6197927455467547,
 'test_macro_f1': 0.5849735880974463,
 'history': [{'epoch': 1,
   'train_loss': 1.3839546795471571,
   'train_accuracy': 0.6990859754355898,
   'train_macro_f1': 0.48822919560862854,
   'val_loss': 0.9466811304926561,
   'val_accuracy': 0.652088772845953,
   'val_balanced_accuracy': 0.5606359786821179,
   'val_macro_f1': 0.5059884445289596},
  {'epoch': 2,
   'train_loss': 0.6245764851127479,
   'train_accuracy': 0.768494715795487,
   'train_macro_f1': 0.7032388669681201,
   'val_loss': 0.8368626000364517,
   'val_accuracy': 0.6742819843342036,
   'val_balanced_accuracy': 0.5588679356074729,
   'val_macro_f1': 0.4872402187896393},
  {'epoch': 3,
   'train_loss': 0.4115912770925335,
   'train_accuracy': 0.8199085975435589,
   'train_macro_f1': 0.7836684385321069,
   'val_loss': 0.7735871212756353,
   'val_accuracy': 0.7030026109660574,
   'val_balanced_accuracy': 0.6027994565925806,
   'val_macro_f1': 0.5508253382275645},
  {'epoch': 4,
   'train_loss': 0.3064043220452873,
   'train_accuracy': 0.8429020279920023,
   'train_macro_f1': 0.8241916157482325,
   'val_loss': 0.7659847238745453,
   'val_accuracy': 0.7212793733681462,
   'val_balanced_accuracy': 0.5920181909197215,
   'val_macro_f1': 0.5654593046886761},
  {'epoch': 5,
   'train_loss': 0.24235398945247266,
   'train_accuracy': 0.8686089688660382,
   'train_macro_f1': 0.8565311292610884,
   'val_loss': 0.8055419812937007,
   'val_accuracy': 0.7173629242819843,
   'val_balanced_accuracy': 0.5980349626666673,
   'val_macro_f1': 0.5734997843740176}]}
In [27]:
all_metadata_results = []

for exp in metadata_experiments:
    print("\n" + "=" * 60)
    print("Running:", exp["name"], exp["features"])
    print("=" * 60)

    result = run_one_metadata_experiment(
        experiment_name=exp["name"],
        selected_features=exp["features"],
        num_epochs=20
    )

    all_metadata_results.append(result)
============================================================
Running: image_age ['age']
============================================================
[image_age] epoch 1/20 | train_f1=0.5012 | val_f1=0.5149 | val_bal_acc=0.6030
[image_age] epoch 2/20 | train_f1=0.7098 | val_f1=0.5307 | val_bal_acc=0.5850
[image_age] epoch 3/20 | train_f1=0.7853 | val_f1=0.5331 | val_bal_acc=0.5962
[image_age] epoch 4/20 | train_f1=0.8168 | val_f1=0.5570 | val_bal_acc=0.6056
[image_age] epoch 5/20 | train_f1=0.8527 | val_f1=0.5543 | val_bal_acc=0.5901
[image_age] epoch 6/20 | train_f1=0.8765 | val_f1=0.5880 | val_bal_acc=0.6021
[image_age] epoch 7/20 | train_f1=0.8978 | val_f1=0.5816 | val_bal_acc=0.6035
[image_age] epoch 8/20 | train_f1=0.9104 | val_f1=0.5784 | val_bal_acc=0.6071
[image_age] epoch 9/20 | train_f1=0.9218 | val_f1=0.5972 | val_bal_acc=0.5850
[image_age] epoch 10/20 | train_f1=0.9347 | val_f1=0.5826 | val_bal_acc=0.5867
[image_age] epoch 11/20 | train_f1=0.9475 | val_f1=0.5894 | val_bal_acc=0.5908
[image_age] epoch 12/20 | train_f1=0.9524 | val_f1=0.5958 | val_bal_acc=0.5768
[image_age] epoch 13/20 | train_f1=0.9569 | val_f1=0.5903 | val_bal_acc=0.5996
[image_age] epoch 14/20 | train_f1=0.9636 | val_f1=0.5757 | val_bal_acc=0.5836
[image_age] epoch 15/20 | train_f1=0.9680 | val_f1=0.5833 | val_bal_acc=0.5616
[image_age] epoch 16/20 | train_f1=0.9723 | val_f1=0.5786 | val_bal_acc=0.5625
[image_age] epoch 17/20 | train_f1=0.9728 | val_f1=0.5963 | val_bal_acc=0.5711
[image_age] epoch 18/20 | train_f1=0.9793 | val_f1=0.5891 | val_bal_acc=0.5786
[image_age] epoch 19/20 | train_f1=0.9822 | val_f1=0.5904 | val_bal_acc=0.5687
[image_age] epoch 20/20 | train_f1=0.9862 | val_f1=0.6015 | val_bal_acc=0.5752

============================================================
Running: image_sex ['sex']
============================================================
[image_sex] epoch 1/20 | train_f1=0.4632 | val_f1=0.4550 | val_bal_acc=0.5105
[image_sex] epoch 2/20 | train_f1=0.6957 | val_f1=0.5361 | val_bal_acc=0.5902
[image_sex] epoch 3/20 | train_f1=0.7799 | val_f1=0.5808 | val_bal_acc=0.6039
[image_sex] epoch 4/20 | train_f1=0.8306 | val_f1=0.5589 | val_bal_acc=0.5893
[image_sex] epoch 5/20 | train_f1=0.8519 | val_f1=0.5749 | val_bal_acc=0.5920
[image_sex] epoch 6/20 | train_f1=0.8738 | val_f1=0.5544 | val_bal_acc=0.5865
[image_sex] epoch 7/20 | train_f1=0.8980 | val_f1=0.5896 | val_bal_acc=0.5855
[image_sex] epoch 8/20 | train_f1=0.9071 | val_f1=0.5693 | val_bal_acc=0.5934
[image_sex] epoch 9/20 | train_f1=0.9208 | val_f1=0.5789 | val_bal_acc=0.5726
[image_sex] epoch 10/20 | train_f1=0.9367 | val_f1=0.5866 | val_bal_acc=0.5887
[image_sex] epoch 11/20 | train_f1=0.9435 | val_f1=0.5625 | val_bal_acc=0.5687
[image_sex] epoch 12/20 | train_f1=0.9475 | val_f1=0.5888 | val_bal_acc=0.5791
[image_sex] epoch 13/20 | train_f1=0.9620 | val_f1=0.5830 | val_bal_acc=0.5765
[image_sex] epoch 14/20 | train_f1=0.9666 | val_f1=0.5782 | val_bal_acc=0.5649
[image_sex] epoch 15/20 | train_f1=0.9663 | val_f1=0.5914 | val_bal_acc=0.5779
[image_sex] epoch 16/20 | train_f1=0.9732 | val_f1=0.5878 | val_bal_acc=0.5696
[image_sex] epoch 17/20 | train_f1=0.9781 | val_f1=0.5825 | val_bal_acc=0.5800
[image_sex] epoch 18/20 | train_f1=0.9760 | val_f1=0.5875 | val_bal_acc=0.5754
[image_sex] epoch 19/20 | train_f1=0.9772 | val_f1=0.5746 | val_bal_acc=0.5460
[image_sex] epoch 20/20 | train_f1=0.9840 | val_f1=0.5956 | val_bal_acc=0.5649

============================================================
Running: image_location ['location']
============================================================
[image_location] epoch 1/20 | train_f1=0.4688 | val_f1=0.4996 | val_bal_acc=0.5669
[image_location] epoch 2/20 | train_f1=0.7125 | val_f1=0.5245 | val_bal_acc=0.5685
[image_location] epoch 3/20 | train_f1=0.7910 | val_f1=0.5335 | val_bal_acc=0.6017
[image_location] epoch 4/20 | train_f1=0.8364 | val_f1=0.5662 | val_bal_acc=0.6093
[image_location] epoch 5/20 | train_f1=0.8599 | val_f1=0.5769 | val_bal_acc=0.6020
[image_location] epoch 6/20 | train_f1=0.8846 | val_f1=0.5499 | val_bal_acc=0.5829
[image_location] epoch 7/20 | train_f1=0.8923 | val_f1=0.5505 | val_bal_acc=0.5811
[image_location] epoch 8/20 | train_f1=0.9070 | val_f1=0.5615 | val_bal_acc=0.5782
[image_location] epoch 9/20 | train_f1=0.9212 | val_f1=0.5816 | val_bal_acc=0.5858
[image_location] epoch 10/20 | train_f1=0.9301 | val_f1=0.6034 | val_bal_acc=0.5798
[image_location] epoch 11/20 | train_f1=0.9418 | val_f1=0.5847 | val_bal_acc=0.5901
[image_location] epoch 12/20 | train_f1=0.9498 | val_f1=0.5795 | val_bal_acc=0.5664
[image_location] epoch 13/20 | train_f1=0.9594 | val_f1=0.5965 | val_bal_acc=0.5902
[image_location] epoch 14/20 | train_f1=0.9655 | val_f1=0.5884 | val_bal_acc=0.5941
[image_location] epoch 15/20 | train_f1=0.9692 | val_f1=0.5811 | val_bal_acc=0.5801
[image_location] epoch 16/20 | train_f1=0.9739 | val_f1=0.5920 | val_bal_acc=0.5833
[image_location] epoch 17/20 | train_f1=0.9779 | val_f1=0.5895 | val_bal_acc=0.5599
[image_location] epoch 18/20 | train_f1=0.9798 | val_f1=0.5950 | val_bal_acc=0.5891
[image_location] epoch 19/20 | train_f1=0.9843 | val_f1=0.5859 | val_bal_acc=0.5646
[image_location] epoch 20/20 | train_f1=0.9858 | val_f1=0.5970 | val_bal_acc=0.5769

============================================================
Running: image_age_sex ['age', 'sex']
============================================================
[image_age_sex] epoch 1/20 | train_f1=0.4752 | val_f1=0.5057 | val_bal_acc=0.5597
[image_age_sex] epoch 2/20 | train_f1=0.7065 | val_f1=0.5307 | val_bal_acc=0.5940
[image_age_sex] epoch 3/20 | train_f1=0.7807 | val_f1=0.5083 | val_bal_acc=0.5936
[image_age_sex] epoch 4/20 | train_f1=0.8348 | val_f1=0.5687 | val_bal_acc=0.6064
[image_age_sex] epoch 5/20 | train_f1=0.8505 | val_f1=0.5811 | val_bal_acc=0.5880
[image_age_sex] epoch 6/20 | train_f1=0.8822 | val_f1=0.5826 | val_bal_acc=0.6037
[image_age_sex] epoch 7/20 | train_f1=0.8937 | val_f1=0.5630 | val_bal_acc=0.5960
[image_age_sex] epoch 8/20 | train_f1=0.9064 | val_f1=0.5834 | val_bal_acc=0.5591
[image_age_sex] epoch 9/20 | train_f1=0.9319 | val_f1=0.5944 | val_bal_acc=0.5772
[image_age_sex] epoch 10/20 | train_f1=0.9352 | val_f1=0.5895 | val_bal_acc=0.5870
[image_age_sex] epoch 11/20 | train_f1=0.9415 | val_f1=0.5991 | val_bal_acc=0.5921
[image_age_sex] epoch 12/20 | train_f1=0.9506 | val_f1=0.5688 | val_bal_acc=0.5918
[image_age_sex] epoch 13/20 | train_f1=0.9603 | val_f1=0.5931 | val_bal_acc=0.5819
[image_age_sex] epoch 14/20 | train_f1=0.9586 | val_f1=0.5739 | val_bal_acc=0.5733
[image_age_sex] epoch 15/20 | train_f1=0.9675 | val_f1=0.5854 | val_bal_acc=0.5914
[image_age_sex] epoch 16/20 | train_f1=0.9744 | val_f1=0.5945 | val_bal_acc=0.5987
[image_age_sex] epoch 17/20 | train_f1=0.9746 | val_f1=0.5883 | val_bal_acc=0.5715
[image_age_sex] epoch 18/20 | train_f1=0.9803 | val_f1=0.5759 | val_bal_acc=0.5634
[image_age_sex] epoch 19/20 | train_f1=0.9863 | val_f1=0.5937 | val_bal_acc=0.5808
[image_age_sex] epoch 20/20 | train_f1=0.9838 | val_f1=0.5844 | val_bal_acc=0.5744

============================================================
Running: image_age_location ['age', 'location']
============================================================
[image_age_location] epoch 1/20 | train_f1=0.5011 | val_f1=0.5017 | val_bal_acc=0.5698
[image_age_location] epoch 2/20 | train_f1=0.6965 | val_f1=0.5357 | val_bal_acc=0.5855
[image_age_location] epoch 3/20 | train_f1=0.7875 | val_f1=0.5528 | val_bal_acc=0.5981
[image_age_location] epoch 4/20 | train_f1=0.8351 | val_f1=0.5584 | val_bal_acc=0.5980
[image_age_location] epoch 5/20 | train_f1=0.8621 | val_f1=0.5727 | val_bal_acc=0.5914
[image_age_location] epoch 6/20 | train_f1=0.8741 | val_f1=0.5708 | val_bal_acc=0.5931
[image_age_location] epoch 7/20 | train_f1=0.8991 | val_f1=0.5810 | val_bal_acc=0.5600
[image_age_location] epoch 8/20 | train_f1=0.9189 | val_f1=0.5685 | val_bal_acc=0.5722
[image_age_location] epoch 9/20 | train_f1=0.9281 | val_f1=0.5804 | val_bal_acc=0.5928
[image_age_location] epoch 10/20 | train_f1=0.9360 | val_f1=0.5894 | val_bal_acc=0.5909
[image_age_location] epoch 11/20 | train_f1=0.9496 | val_f1=0.5799 | val_bal_acc=0.5827
[image_age_location] epoch 12/20 | train_f1=0.9534 | val_f1=0.5899 | val_bal_acc=0.5819
[image_age_location] epoch 13/20 | train_f1=0.9588 | val_f1=0.5854 | val_bal_acc=0.5744
[image_age_location] epoch 14/20 | train_f1=0.9628 | val_f1=0.5895 | val_bal_acc=0.5724
[image_age_location] epoch 15/20 | train_f1=0.9694 | val_f1=0.5879 | val_bal_acc=0.5619
[image_age_location] epoch 16/20 | train_f1=0.9745 | val_f1=0.5868 | val_bal_acc=0.5575
[image_age_location] epoch 17/20 | train_f1=0.9806 | val_f1=0.5880 | val_bal_acc=0.5554
[image_age_location] epoch 18/20 | train_f1=0.9813 | val_f1=0.5847 | val_bal_acc=0.5463
[image_age_location] epoch 19/20 | train_f1=0.9802 | val_f1=0.5866 | val_bal_acc=0.5580
[image_age_location] epoch 20/20 | train_f1=0.9813 | val_f1=0.5761 | val_bal_acc=0.5728

============================================================
Running: image_sex_location ['sex', 'location']
============================================================
[image_sex_location] epoch 1/20 | train_f1=0.5017 | val_f1=0.4928 | val_bal_acc=0.5956
[image_sex_location] epoch 2/20 | train_f1=0.7018 | val_f1=0.5180 | val_bal_acc=0.5935
[image_sex_location] epoch 3/20 | train_f1=0.7877 | val_f1=0.5325 | val_bal_acc=0.6116
[image_sex_location] epoch 4/20 | train_f1=0.8316 | val_f1=0.5575 | val_bal_acc=0.6076
[image_sex_location] epoch 5/20 | train_f1=0.8544 | val_f1=0.5678 | val_bal_acc=0.6074
[image_sex_location] epoch 6/20 | train_f1=0.8765 | val_f1=0.5677 | val_bal_acc=0.5662
[image_sex_location] epoch 7/20 | train_f1=0.8964 | val_f1=0.5733 | val_bal_acc=0.5982
[image_sex_location] epoch 8/20 | train_f1=0.9153 | val_f1=0.5631 | val_bal_acc=0.5776
[image_sex_location] epoch 9/20 | train_f1=0.9239 | val_f1=0.5944 | val_bal_acc=0.5935
[image_sex_location] epoch 10/20 | train_f1=0.9324 | val_f1=0.5760 | val_bal_acc=0.5876
[image_sex_location] epoch 11/20 | train_f1=0.9395 | val_f1=0.5810 | val_bal_acc=0.5643
[image_sex_location] epoch 12/20 | train_f1=0.9534 | val_f1=0.5963 | val_bal_acc=0.5926
[image_sex_location] epoch 13/20 | train_f1=0.9587 | val_f1=0.5823 | val_bal_acc=0.5758
[image_sex_location] epoch 14/20 | train_f1=0.9641 | val_f1=0.6067 | val_bal_acc=0.5851
[image_sex_location] epoch 15/20 | train_f1=0.9642 | val_f1=0.5745 | val_bal_acc=0.5500
[image_sex_location] epoch 16/20 | train_f1=0.9734 | val_f1=0.5822 | val_bal_acc=0.5826
[image_sex_location] epoch 17/20 | train_f1=0.9737 | val_f1=0.5987 | val_bal_acc=0.5812
[image_sex_location] epoch 18/20 | train_f1=0.9806 | val_f1=0.5735 | val_bal_acc=0.5692
[image_sex_location] epoch 19/20 | train_f1=0.9852 | val_f1=0.5873 | val_bal_acc=0.5760
[image_sex_location] epoch 20/20 | train_f1=0.9716 | val_f1=0.5764 | val_bal_acc=0.5698

============================================================
Running: image_all_metadata ['age', 'sex', 'location']
============================================================
[image_all_metadata] epoch 1/20 | train_f1=0.4778 | val_f1=0.5028 | val_bal_acc=0.5301
[image_all_metadata] epoch 2/20 | train_f1=0.7088 | val_f1=0.5214 | val_bal_acc=0.6097
[image_all_metadata] epoch 3/20 | train_f1=0.7820 | val_f1=0.5618 | val_bal_acc=0.6057
[image_all_metadata] epoch 4/20 | train_f1=0.8324 | val_f1=0.5492 | val_bal_acc=0.6055
[image_all_metadata] epoch 5/20 | train_f1=0.8543 | val_f1=0.5639 | val_bal_acc=0.5608
[image_all_metadata] epoch 6/20 | train_f1=0.8742 | val_f1=0.5952 | val_bal_acc=0.5957
[image_all_metadata] epoch 7/20 | train_f1=0.8961 | val_f1=0.5670 | val_bal_acc=0.5806
[image_all_metadata] epoch 8/20 | train_f1=0.9113 | val_f1=0.5891 | val_bal_acc=0.5825
[image_all_metadata] epoch 9/20 | train_f1=0.9171 | val_f1=0.5754 | val_bal_acc=0.5714
[image_all_metadata] epoch 10/20 | train_f1=0.9360 | val_f1=0.5677 | val_bal_acc=0.5721
[image_all_metadata] epoch 11/20 | train_f1=0.9387 | val_f1=0.5813 | val_bal_acc=0.5804
[image_all_metadata] epoch 12/20 | train_f1=0.9542 | val_f1=0.6008 | val_bal_acc=0.5955
[image_all_metadata] epoch 13/20 | train_f1=0.9598 | val_f1=0.5708 | val_bal_acc=0.5676
[image_all_metadata] epoch 14/20 | train_f1=0.9580 | val_f1=0.6113 | val_bal_acc=0.5836
[image_all_metadata] epoch 15/20 | train_f1=0.9706 | val_f1=0.5751 | val_bal_acc=0.5417
[image_all_metadata] epoch 16/20 | train_f1=0.9740 | val_f1=0.5789 | val_bal_acc=0.5646
[image_all_metadata] epoch 17/20 | train_f1=0.9784 | val_f1=0.5853 | val_bal_acc=0.5745
[image_all_metadata] epoch 18/20 | train_f1=0.9713 | val_f1=0.5878 | val_bal_acc=0.5716
[image_all_metadata] epoch 19/20 | train_f1=0.9849 | val_f1=0.5895 | val_bal_acc=0.5776
[image_all_metadata] epoch 20/20 | train_f1=0.9818 | val_f1=0.5980 | val_bal_acc=0.5844
In [28]:
metadata_result_rows = []

for result in all_metadata_results:
    metadata_result_rows.append({
        "Method": result["name"],
        "Features": ",".join(result["features"]),
        "Metadata Dim": result["metadata_dim"],
        "Best Val Macro-F1": result["best_val_macro_f1"],
        "Test Accuracy": result["test_accuracy"],
        "Test Balanced Accuracy": result["test_balanced_accuracy"],
        "Test Macro-F1": result["test_macro_f1"],
    })

metadata_results_df = pd.DataFrame(metadata_result_rows)

metadata_results_df = metadata_results_df.sort_values(
    by="Test Macro-F1",
    ascending=False
).reset_index(drop=True)

metadata_results_df
Out[28]:
Method Features Metadata Dim Best Val Macro-F1 Test Accuracy Test Balanced Accuracy Test Macro-F1
0 image_age age 1 0.601548 0.791357 0.567975 0.587573
1 image_age_sex age,sex 4 0.599118 0.766374 0.574931 0.578003
2 image_sex_location sex,location 18 0.606667 0.775827 0.560076 0.570476
3 image_age_location age,location 16 0.589866 0.775827 0.571165 0.567851
4 image_all_metadata age,sex,location 19 0.611333 0.767725 0.554128 0.566352
5 image_sex sex 3 0.595635 0.781904 0.538633 0.563919
6 image_location location 15 0.603405 0.770425 0.557796 0.562553
In [30]:
save_path = os.path.join(
    PROJECT_DIR,
    "supports",
    "0803_metadata_combination_results.csv"
)

metadata_results_df.to_csv(
    save_path,
    index=False,
    encoding="utf-8-sig"
)

print("Saved to:", save_path)
Saved to: /Users/applesues01/Documents/Medical_Agent/supports/0803_metadata_combination_results.csv
In [31]:
import matplotlib.pyplot as plt

plt.figure(figsize=(10, 5))

plt.bar(
    metadata_results_df["Method"],
    metadata_results_df["Test Macro-F1"]
)

plt.xticks(rotation=45, ha="right")
plt.ylabel("Test Macro-F1")
plt.title("Performance of Different Metadata Combinations")
plt.tight_layout()
plt.show()
No description has been provided for this image
In [32]:
class MetadataEncoderFusionClassifier(nn.Module):
    def __init__(self, metadata_input_dim, metadata_embed_dim=64, num_classes=7):
        super().__init__()

        self.metadata_encoder = nn.Sequential(
            nn.Linear(metadata_input_dim, metadata_embed_dim),
            nn.BatchNorm1d(metadata_embed_dim),
            nn.ReLU(),
            nn.Dropout(0.2)
        )

        self.classifier = nn.Sequential(
            nn.Linear(2048 + metadata_embed_dim, 512),
            nn.ReLU(),
            nn.Dropout(0.3),
            nn.Linear(512, 128),
            nn.ReLU(),
            nn.Linear(128, num_classes)
        )

    def forward(self, image_features, metadata):
        metadata_features = self.metadata_encoder(metadata)
        fused = torch.cat([image_features, metadata_features], dim=1)
        return self.classifier(fused)
In [33]:
def run_metadata_encoder_experiment(
    experiment_name,
    selected_features,
    metadata_embed_dim=64,
    num_epochs=20
):
    train_loader, val_loader, test_loader, metadata_dim = build_cached_loaders_for_experiment(
        selected_features
    )

    model = MetadataEncoderFusionClassifier(
        metadata_input_dim=metadata_dim,
        metadata_embed_dim=metadata_embed_dim,
        num_classes=7
    ).to(device)

    criterion = nn.CrossEntropyLoss(weight=class_weights)

    optimizer = torch.optim.Adam(
        model.parameters(),
        lr=1e-4
    )

    best_val_f1 = -1.0
    best_state = None
    history = []

    for epoch in range(num_epochs):
        train_metrics = train_cached_one_experiment(
            model,
            train_loader,
            criterion,
            optimizer,
            device
        )

        val_metrics = evaluate_cached_model(
            model,
            val_loader,
            criterion,
            device
        )

        history.append({
            "epoch": epoch + 1,
            "train_loss": train_metrics["loss"],
            "train_accuracy": train_metrics["accuracy"],
            "train_macro_f1": train_metrics["macro_f1"],
            "val_loss": val_metrics["loss"],
            "val_accuracy": val_metrics["accuracy"],
            "val_balanced_accuracy": val_metrics["balanced_accuracy"],
            "val_macro_f1": val_metrics["macro_f1"],
        })

        print(
            f"[{experiment_name}] "
            f"epoch {epoch+1}/{num_epochs} | "
            f"train_f1={train_metrics['macro_f1']:.4f} | "
            f"val_f1={val_metrics['macro_f1']:.4f} | "
            f"val_bal_acc={val_metrics['balanced_accuracy']:.4f}"
        )

        if val_metrics["macro_f1"] > best_val_f1:
            best_val_f1 = val_metrics["macro_f1"]
            best_state = copy.deepcopy(model.state_dict())

    model.load_state_dict(best_state)

    test_metrics = evaluate_cached_model(
        model,
        test_loader,
        criterion,
        device
    )

    return {
        "name": experiment_name,
        "features": selected_features,
        "metadata_dim": metadata_dim,
        "metadata_embed_dim": metadata_embed_dim,
        "best_val_macro_f1": best_val_f1,
        "test_accuracy": test_metrics["accuracy"],
        "test_balanced_accuracy": test_metrics["balanced_accuracy"],
        "test_macro_f1": test_metrics["macro_f1"],
        "history": history
    }
In [34]:
encoded64_result = run_metadata_encoder_experiment(
    experiment_name="image_all_metadata_encoded64",
    selected_features=["age", "sex", "location"],
    metadata_embed_dim=64,
    num_epochs=20
)

encoded64_result
[image_all_metadata_encoded64] epoch 1/20 | train_f1=0.4675 | val_f1=0.5298 | val_bal_acc=0.5590
[image_all_metadata_encoded64] epoch 2/20 | train_f1=0.7219 | val_f1=0.5643 | val_bal_acc=0.6264
[image_all_metadata_encoded64] epoch 3/20 | train_f1=0.7902 | val_f1=0.5439 | val_bal_acc=0.6471
[image_all_metadata_encoded64] epoch 4/20 | train_f1=0.8352 | val_f1=0.5624 | val_bal_acc=0.5849
[image_all_metadata_encoded64] epoch 5/20 | train_f1=0.8634 | val_f1=0.5909 | val_bal_acc=0.5903
[image_all_metadata_encoded64] epoch 6/20 | train_f1=0.8843 | val_f1=0.5746 | val_bal_acc=0.5723
[image_all_metadata_encoded64] epoch 7/20 | train_f1=0.8999 | val_f1=0.5637 | val_bal_acc=0.6034
[image_all_metadata_encoded64] epoch 8/20 | train_f1=0.9187 | val_f1=0.5983 | val_bal_acc=0.6013
[image_all_metadata_encoded64] epoch 9/20 | train_f1=0.9239 | val_f1=0.6157 | val_bal_acc=0.5993
[image_all_metadata_encoded64] epoch 10/20 | train_f1=0.9291 | val_f1=0.6034 | val_bal_acc=0.5772
[image_all_metadata_encoded64] epoch 11/20 | train_f1=0.9457 | val_f1=0.5832 | val_bal_acc=0.5870
[image_all_metadata_encoded64] epoch 12/20 | train_f1=0.9532 | val_f1=0.5830 | val_bal_acc=0.6113
[image_all_metadata_encoded64] epoch 13/20 | train_f1=0.9600 | val_f1=0.5969 | val_bal_acc=0.5803
[image_all_metadata_encoded64] epoch 14/20 | train_f1=0.9634 | val_f1=0.5939 | val_bal_acc=0.5977
[image_all_metadata_encoded64] epoch 15/20 | train_f1=0.9665 | val_f1=0.5839 | val_bal_acc=0.5700
[image_all_metadata_encoded64] epoch 16/20 | train_f1=0.9773 | val_f1=0.5928 | val_bal_acc=0.5767
[image_all_metadata_encoded64] epoch 17/20 | train_f1=0.9740 | val_f1=0.6069 | val_bal_acc=0.5827
[image_all_metadata_encoded64] epoch 18/20 | train_f1=0.9788 | val_f1=0.5872 | val_bal_acc=0.5855
[image_all_metadata_encoded64] epoch 19/20 | train_f1=0.9831 | val_f1=0.6058 | val_bal_acc=0.5734
[image_all_metadata_encoded64] epoch 20/20 | train_f1=0.9872 | val_f1=0.5902 | val_bal_acc=0.5562
Out[34]:
{'name': 'image_all_metadata_encoded64',
 'features': ['age', 'sex', 'location'],
 'metadata_dim': 19,
 'metadata_embed_dim': 64,
 'best_val_macro_f1': 0.6157358157993051,
 'test_accuracy': 0.7832545577312626,
 'test_balanced_accuracy': 0.5784537893091156,
 'test_macro_f1': 0.581137785797784,
 'history': [{'epoch': 1,
   'train_loss': 1.3357206415837235,
   'train_accuracy': 0.6423878891745216,
   'train_macro_f1': 0.46752364134565044,
   'val_loss': 0.8145539368412824,
   'val_accuracy': 0.6860313315926893,
   'val_balanced_accuracy': 0.5589958286747888,
   'val_macro_f1': 0.529844716712117},
  {'epoch': 2,
   'train_loss': 0.5933823679403317,
   'train_accuracy': 0.7842045129962868,
   'train_macro_f1': 0.7219091238192819,
   'val_loss': 0.7064913759499244,
   'val_accuracy': 0.7271540469973891,
   'val_balanced_accuracy': 0.6264104888232229,
   'val_macro_f1': 0.564314919697155},
  {'epoch': 3,
   'train_loss': 0.3923551866035058,
   'train_accuracy': 0.8290488431876607,
   'train_macro_f1': 0.7902266395050317,
   'val_loss': 0.877769479085509,
   'val_accuracy': 0.6612271540469974,
   'val_balanced_accuracy': 0.6470999559191623,
   'val_macro_f1': 0.5439472681164729},
  {'epoch': 4,
   'train_loss': 0.2988949242624955,
   'train_accuracy': 0.8577549271636675,
   'train_macro_f1': 0.8352326332503784,
   'val_loss': 0.71763523306221,
   'val_accuracy': 0.7343342036553525,
   'val_balanced_accuracy': 0.5849282399077159,
   'val_macro_f1': 0.5624342194523481},
  {'epoch': 5,
   'train_loss': 0.23664439923931482,
   'train_accuracy': 0.8770351328191945,
   'train_macro_f1': 0.8634419584494657,
   'val_loss': 0.7153918144460913,
   'val_accuracy': 0.7408616187989556,
   'val_balanced_accuracy': 0.5903126497961575,
   'val_macro_f1': 0.5909086320354072},
  {'epoch': 6,
   'train_loss': 0.1905775090875846,
   'train_accuracy': 0.890745501285347,
   'train_macro_f1': 0.8842949950038163,
   'val_loss': 0.7429845648533214,
   'val_accuracy': 0.7369451697127938,
   'val_balanced_accuracy': 0.5723462071315655,
   'val_macro_f1': 0.5745934978263565},
  {'epoch': 7,
   'train_loss': 0.17204279489698188,
   'train_accuracy': 0.9038846043987432,
   'train_macro_f1': 0.8998669403348616,
   'val_loss': 0.7783091074687383,
   'val_accuracy': 0.7467362924281984,
   'val_balanced_accuracy': 0.6034478314597946,
   'val_macro_f1': 0.5637253063134304},
  {'epoch': 8,
   'train_loss': 0.1406310166215529,
   'train_accuracy': 0.9171665238503285,
   'train_macro_f1': 0.9186953047632763,
   'val_loss': 0.7738940058854791,
   'val_accuracy': 0.762402088772846,
   'val_balanced_accuracy': 0.6013058474598657,
   'val_macro_f1': 0.598280527961947},
  {'epoch': 9,
   'train_loss': 0.11683716404447808,
   'train_accuracy': 0.9253070551271065,
   'train_macro_f1': 0.9238928841639754,
   'val_loss': 0.7905030644671412,
   'val_accuracy': 0.7715404699738904,
   'val_balanced_accuracy': 0.599288688981812,
   'val_macro_f1': 0.6157358157993051},
  {'epoch': 10,
   'train_loss': 0.11439687799118614,
   'train_accuracy': 0.9247357897743502,
   'train_macro_f1': 0.9290616170493092,
   'val_loss': 0.8208377551841829,
   'val_accuracy': 0.7637075718015666,
   'val_balanced_accuracy': 0.5772384870927307,
   'val_macro_f1': 0.6034019699694605},
  {'epoch': 11,
   'train_loss': 0.09125680477557647,
   'train_accuracy': 0.9418737503570408,
   'train_macro_f1': 0.9457417800916122,
   'val_loss': 0.8944401063423788,
   'val_accuracy': 0.7441253263707572,
   'val_balanced_accuracy': 0.5870414246215354,
   'val_macro_f1': 0.5831839491883171},
  {'epoch': 12,
   'train_loss': 0.07788306252189753,
   'train_accuracy': 0.9481576692373608,
   'train_macro_f1': 0.9531971860457158,
   'val_loss': 0.9210110665578638,
   'val_accuracy': 0.7434725848563969,
   'val_balanced_accuracy': 0.6113389312745247,
   'val_macro_f1': 0.5829601932673681},
  {'epoch': 13,
   'train_loss': 0.07236704741751201,
   'train_accuracy': 0.9525849757212225,
   'train_macro_f1': 0.9599688476546783,
   'val_loss': 0.8845066959670607,
   'val_accuracy': 0.7715404699738904,
   'val_balanced_accuracy': 0.5803334998666624,
   'val_macro_f1': 0.5968973247230877},
  {'epoch': 14,
   'train_loss': 0.06066363946564221,
   'train_accuracy': 0.9594401599542988,
   'train_macro_f1': 0.9634069603286285,
   'val_loss': 0.9544521311599538,
   'val_accuracy': 0.77088772845953,
   'val_balanced_accuracy': 0.5977145433779574,
   'val_macro_f1': 0.5938940889974026},
  {'epoch': 15,
   'train_loss': 0.05386901470430576,
   'train_accuracy': 0.9632962010854041,
   'train_macro_f1': 0.9664668192913553,
   'val_loss': 0.9910254857877352,
   'val_accuracy': 0.7617493472584856,
   'val_balanced_accuracy': 0.569999787355182,
   'val_macro_f1': 0.5838641381233565},
  {'epoch': 16,
   'train_loss': 0.0410140141606931,
   'train_accuracy': 0.9718651813767495,
   'train_macro_f1': 0.9773320137390435,
   'val_loss': 1.0144228580074661,
   'val_accuracy': 0.7813315926892951,
   'val_balanced_accuracy': 0.5766725038217965,
   'val_macro_f1': 0.5928227104779502},
  {'epoch': 17,
   'train_loss': 0.04143895722353196,
   'train_accuracy': 0.9712939160239932,
   'train_macro_f1': 0.9739741902846506,
   'val_loss': 0.990178325275499,
   'val_accuracy': 0.7924281984334204,
   'val_balanced_accuracy': 0.5826623368826253,
   'val_macro_f1': 0.6068757050481833},
  {'epoch': 18,
   'train_loss': 0.036266747962122675,
   'train_accuracy': 0.9748643244787204,
   'train_macro_f1': 0.9788230045260337,
   'val_loss': 1.2343657589943073,
   'val_accuracy': 0.7323759791122716,
   'val_balanced_accuracy': 0.5855029836687579,
   'val_macro_f1': 0.5872478871044825},
  {'epoch': 19,
   'train_loss': 0.030731554383343372,
   'train_accuracy': 0.9801485289917167,
   'train_macro_f1': 0.9830872186944252,
   'val_loss': 1.0749318152344658,
   'val_accuracy': 0.7891644908616188,
   'val_balanced_accuracy': 0.573380753088806,
   'val_macro_f1': 0.6058054569794089},
  {'epoch': 20,
   'train_loss': 0.024434133543867285,
   'train_accuracy': 0.9830048557554985,
   'train_macro_f1': 0.9871951877521231,
   'val_loss': 1.1455151122083915,
   'val_accuracy': 0.77088772845953,
   'val_balanced_accuracy': 0.5562054487125988,
   'val_macro_f1': 0.5901614616424423}]}
In [35]:
encoded64_row = pd.DataFrame([{
    "Method": encoded64_result["name"],
    "Features": ",".join(encoded64_result["features"]),
    "Metadata Dim": encoded64_result["metadata_dim"],
    "Metadata Embed Dim": encoded64_result["metadata_embed_dim"],
    "Best Val Macro-F1": encoded64_result["best_val_macro_f1"],
    "Test Accuracy": encoded64_result["test_accuracy"],
    "Test Balanced Accuracy": encoded64_result["test_balanced_accuracy"],
    "Test Macro-F1": encoded64_result["test_macro_f1"],
}])

encoded64_row
Out[35]:
Method Features Metadata Dim Metadata Embed Dim Best Val Macro-F1 Test Accuracy Test Balanced Accuracy Test Macro-F1
0 image_all_metadata_encoded64 age,sex,location 19 64 0.615736 0.783255 0.578454 0.581138
In [36]:
encoded64_save_path = os.path.join(
    PROJECT_DIR,
    "supports",
    "metadata_encoded64_result.csv"
)

encoded64_row.to_csv(
    encoded64_save_path,
    index=False,
    encoding="utf-8-sig"
)

print("Saved:", encoded64_save_path)
Saved: /Users/applesues01/Documents/Medical_Agent/supports/metadata_encoded64_result.csv
In [37]:
encoded64_history_df = pd.DataFrame(encoded64_result["history"])

encoded64_history_path = os.path.join(
    PROJECT_DIR,
    "supports",
    "metadata_encoded64_history.csv"
)

encoded64_history_df.to_csv(
    encoded64_history_path,
    index=False,
    encoding="utf-8-sig"
)

print("Saved:", encoded64_history_path)
Saved: /Users/applesues01/Documents/Medical_Agent/supports/metadata_encoded64_history.csv
In [38]:
current_summary_path = os.path.join(
    PROJECT_DIR,
    "supports",
    "metadata_combination_results.csv"
)

if os.path.exists(current_summary_path):
    current_df = pd.read_csv(current_summary_path)
    merged_df = pd.concat([current_df, encoded64_row], ignore_index=True)
else:
    merged_df = encoded64_row.copy()

merged_df.to_csv(
    current_summary_path,
    index=False,
    encoding="utf-8-sig"
)

merged_df
Out[38]:
Method Features Metadata Dim Best Val Macro-F1 Test Accuracy Test Balanced Accuracy Test Macro-F1 Metadata Embed Dim
0 image_age age 1 0.601548 0.791357 0.567975 0.587573 NaN
1 image_age_sex age,sex 4 0.599118 0.766374 0.574931 0.578003 NaN
2 image_sex_location sex,location 18 0.606667 0.775827 0.560076 0.570476 NaN
3 image_age_location age,location 16 0.589866 0.775827 0.571165 0.567851 NaN
4 image_all_metadata age,sex,location 19 0.611333 0.767725 0.554128 0.566352 NaN
5 image_sex sex 3 0.595635 0.781904 0.538633 0.563919 NaN
6 image_location location 15 0.603405 0.770425 0.557796 0.562553 NaN
7 image_all_metadata_encoded64 age,sex,location 19 0.615736 0.783255 0.578454 0.581138 64.0
In [39]:
embed_dim_results = []

for embed_dim in [32, 64, 128]:
    print("\n" + "=" * 60)
    print(f"Running metadata encoder with embed_dim = {embed_dim}")
    print("=" * 60)

    result = run_metadata_encoder_experiment(
        experiment_name=f"image_all_metadata_encoded{embed_dim}",
        selected_features=["age", "sex", "location"],
        metadata_embed_dim=embed_dim,
        num_epochs=20
    )

    embed_dim_results.append(result)
============================================================
Running metadata encoder with embed_dim = 32
============================================================
[image_all_metadata_encoded32] epoch 1/20 | train_f1=0.4845 | val_f1=0.5016 | val_bal_acc=0.5758
[image_all_metadata_encoded32] epoch 2/20 | train_f1=0.7044 | val_f1=0.5603 | val_bal_acc=0.6086
[image_all_metadata_encoded32] epoch 3/20 | train_f1=0.7851 | val_f1=0.5721 | val_bal_acc=0.6070
[image_all_metadata_encoded32] epoch 4/20 | train_f1=0.8320 | val_f1=0.5655 | val_bal_acc=0.5894
[image_all_metadata_encoded32] epoch 5/20 | train_f1=0.8613 | val_f1=0.5260 | val_bal_acc=0.5937
[image_all_metadata_encoded32] epoch 6/20 | train_f1=0.8746 | val_f1=0.5868 | val_bal_acc=0.5888
[image_all_metadata_encoded32] epoch 7/20 | train_f1=0.8992 | val_f1=0.5778 | val_bal_acc=0.5800
[image_all_metadata_encoded32] epoch 8/20 | train_f1=0.9171 | val_f1=0.6000 | val_bal_acc=0.5766
[image_all_metadata_encoded32] epoch 9/20 | train_f1=0.9259 | val_f1=0.5888 | val_bal_acc=0.5782
[image_all_metadata_encoded32] epoch 10/20 | train_f1=0.9373 | val_f1=0.5939 | val_bal_acc=0.5888
[image_all_metadata_encoded32] epoch 11/20 | train_f1=0.9408 | val_f1=0.5952 | val_bal_acc=0.5957
[image_all_metadata_encoded32] epoch 12/20 | train_f1=0.9454 | val_f1=0.6038 | val_bal_acc=0.5976
[image_all_metadata_encoded32] epoch 13/20 | train_f1=0.9617 | val_f1=0.5884 | val_bal_acc=0.5877
[image_all_metadata_encoded32] epoch 14/20 | train_f1=0.9571 | val_f1=0.5914 | val_bal_acc=0.5886
[image_all_metadata_encoded32] epoch 15/20 | train_f1=0.9660 | val_f1=0.6039 | val_bal_acc=0.5905
[image_all_metadata_encoded32] epoch 16/20 | train_f1=0.9727 | val_f1=0.5936 | val_bal_acc=0.5877
[image_all_metadata_encoded32] epoch 17/20 | train_f1=0.9779 | val_f1=0.6004 | val_bal_acc=0.5890
[image_all_metadata_encoded32] epoch 18/20 | train_f1=0.9805 | val_f1=0.5812 | val_bal_acc=0.5586
[image_all_metadata_encoded32] epoch 19/20 | train_f1=0.9788 | val_f1=0.5868 | val_bal_acc=0.5792
[image_all_metadata_encoded32] epoch 20/20 | train_f1=0.9819 | val_f1=0.5883 | val_bal_acc=0.5653

============================================================
Running metadata encoder with embed_dim = 64
============================================================
[image_all_metadata_encoded64] epoch 1/20 | train_f1=0.4987 | val_f1=0.4986 | val_bal_acc=0.5797
[image_all_metadata_encoded64] epoch 2/20 | train_f1=0.7053 | val_f1=0.5356 | val_bal_acc=0.6124
[image_all_metadata_encoded64] epoch 3/20 | train_f1=0.7836 | val_f1=0.5596 | val_bal_acc=0.5960
[image_all_metadata_encoded64] epoch 4/20 | train_f1=0.8345 | val_f1=0.5576 | val_bal_acc=0.6102
[image_all_metadata_encoded64] epoch 5/20 | train_f1=0.8637 | val_f1=0.5460 | val_bal_acc=0.5997
[image_all_metadata_encoded64] epoch 6/20 | train_f1=0.8875 | val_f1=0.5586 | val_bal_acc=0.5950
[image_all_metadata_encoded64] epoch 7/20 | train_f1=0.9029 | val_f1=0.5693 | val_bal_acc=0.5884
[image_all_metadata_encoded64] epoch 8/20 | train_f1=0.9138 | val_f1=0.5851 | val_bal_acc=0.5921
[image_all_metadata_encoded64] epoch 9/20 | train_f1=0.9258 | val_f1=0.5775 | val_bal_acc=0.5766
[image_all_metadata_encoded64] epoch 10/20 | train_f1=0.9352 | val_f1=0.5638 | val_bal_acc=0.5601
[image_all_metadata_encoded64] epoch 11/20 | train_f1=0.9332 | val_f1=0.5759 | val_bal_acc=0.6066
[image_all_metadata_encoded64] epoch 12/20 | train_f1=0.9486 | val_f1=0.6053 | val_bal_acc=0.5831
[image_all_metadata_encoded64] epoch 13/20 | train_f1=0.9529 | val_f1=0.5710 | val_bal_acc=0.5616
[image_all_metadata_encoded64] epoch 14/20 | train_f1=0.9647 | val_f1=0.5995 | val_bal_acc=0.5851
[image_all_metadata_encoded64] epoch 15/20 | train_f1=0.9672 | val_f1=0.6000 | val_bal_acc=0.5834
[image_all_metadata_encoded64] epoch 16/20 | train_f1=0.9695 | val_f1=0.6031 | val_bal_acc=0.5828
[image_all_metadata_encoded64] epoch 17/20 | train_f1=0.9716 | val_f1=0.5931 | val_bal_acc=0.5582
[image_all_metadata_encoded64] epoch 18/20 | train_f1=0.9772 | val_f1=0.5940 | val_bal_acc=0.5727
[image_all_metadata_encoded64] epoch 19/20 | train_f1=0.9778 | val_f1=0.5933 | val_bal_acc=0.5770
[image_all_metadata_encoded64] epoch 20/20 | train_f1=0.9794 | val_f1=0.5752 | val_bal_acc=0.5272

============================================================
Running metadata encoder with embed_dim = 128
============================================================
[image_all_metadata_encoded128] epoch 1/20 | train_f1=0.4746 | val_f1=0.5074 | val_bal_acc=0.5715
[image_all_metadata_encoded128] epoch 2/20 | train_f1=0.7034 | val_f1=0.5574 | val_bal_acc=0.5739
[image_all_metadata_encoded128] epoch 3/20 | train_f1=0.7796 | val_f1=0.5467 | val_bal_acc=0.6330
[image_all_metadata_encoded128] epoch 4/20 | train_f1=0.8277 | val_f1=0.5741 | val_bal_acc=0.6150
[image_all_metadata_encoded128] epoch 5/20 | train_f1=0.8562 | val_f1=0.5563 | val_bal_acc=0.6057
[image_all_metadata_encoded128] epoch 6/20 | train_f1=0.8697 | val_f1=0.5648 | val_bal_acc=0.5855
[image_all_metadata_encoded128] epoch 7/20 | train_f1=0.9021 | val_f1=0.5657 | val_bal_acc=0.5999
[image_all_metadata_encoded128] epoch 8/20 | train_f1=0.9049 | val_f1=0.5743 | val_bal_acc=0.5956
[image_all_metadata_encoded128] epoch 9/20 | train_f1=0.9254 | val_f1=0.5664 | val_bal_acc=0.5787
[image_all_metadata_encoded128] epoch 10/20 | train_f1=0.9282 | val_f1=0.5757 | val_bal_acc=0.5903
[image_all_metadata_encoded128] epoch 11/20 | train_f1=0.9500 | val_f1=0.6039 | val_bal_acc=0.5842
[image_all_metadata_encoded128] epoch 12/20 | train_f1=0.9474 | val_f1=0.5901 | val_bal_acc=0.5851
[image_all_metadata_encoded128] epoch 13/20 | train_f1=0.9528 | val_f1=0.5928 | val_bal_acc=0.5720
[image_all_metadata_encoded128] epoch 14/20 | train_f1=0.9612 | val_f1=0.5862 | val_bal_acc=0.5659
[image_all_metadata_encoded128] epoch 15/20 | train_f1=0.9688 | val_f1=0.5976 | val_bal_acc=0.5784
[image_all_metadata_encoded128] epoch 16/20 | train_f1=0.9688 | val_f1=0.5679 | val_bal_acc=0.5719
[image_all_metadata_encoded128] epoch 17/20 | train_f1=0.9741 | val_f1=0.5909 | val_bal_acc=0.5868
[image_all_metadata_encoded128] epoch 18/20 | train_f1=0.9765 | val_f1=0.5977 | val_bal_acc=0.5963
[image_all_metadata_encoded128] epoch 19/20 | train_f1=0.9782 | val_f1=0.5875 | val_bal_acc=0.5875
[image_all_metadata_encoded128] epoch 20/20 | train_f1=0.9800 | val_f1=0.5814 | val_bal_acc=0.5845
In [40]:
embed_dim_rows = []

for result in embed_dim_results:
    embed_dim_rows.append({
        "Method": result["name"],
        "Features": ",".join(result["features"]),
        "Metadata Dim": result["metadata_dim"],
        "Metadata Embed Dim": result["metadata_embed_dim"],
        "Best Val Macro-F1": result["best_val_macro_f1"],
        "Test Accuracy": result["test_accuracy"],
        "Test Balanced Accuracy": result["test_balanced_accuracy"],
        "Test Macro-F1": result["test_macro_f1"],
    })

embed_dim_df = pd.DataFrame(embed_dim_rows)
embed_dim_df = embed_dim_df.sort_values(
    by="Test Macro-F1",
    ascending=False
).reset_index(drop=True)

embed_dim_df
Out[40]:
Method Features Metadata Dim Metadata Embed Dim Best Val Macro-F1 Test Accuracy Test Balanced Accuracy Test Macro-F1
0 image_all_metadata_encoded128 age,sex,location 19 128 0.603923 0.799460 0.594053 0.600579
1 image_all_metadata_encoded64 age,sex,location 19 64 0.605293 0.804186 0.581747 0.593228
2 image_all_metadata_encoded32 age,sex,location 19 32 0.603866 0.794733 0.581390 0.582997
In [41]:
embed_dim_save_path = os.path.join(
    PROJECT_DIR,
    "supports",
    "metadata_encoder_dim_results.csv"
)

embed_dim_df.to_csv(
    embed_dim_save_path,
    index=False,
    encoding="utf-8-sig"
)

print("Saved:", embed_dim_save_path)
Saved: /Users/applesues01/Documents/Medical_Agent/supports/metadata_encoder_dim_results.csv
In [42]:
for result in embed_dim_results:
    history_df = pd.DataFrame(result["history"])

    history_path = os.path.join(
        PROJECT_DIR,
        "supports",
        f"{result['name']}_history.csv"
    )

    history_df.to_csv(
        history_path,
        index=False,
        encoding="utf-8-sig"
    )

    print("Saved:", history_path)
Saved: /Users/applesues01/Documents/Medical_Agent/supports/image_all_metadata_encoded32_history.csv
Saved: /Users/applesues01/Documents/Medical_Agent/supports/image_all_metadata_encoded64_history.csv
Saved: /Users/applesues01/Documents/Medical_Agent/supports/image_all_metadata_encoded128_history.csv
In [43]:
plt.figure(figsize=(8, 5))

plt.bar(
    embed_dim_df["Method"],
    embed_dim_df["Test Macro-F1"]
)

plt.xticks(rotation=30, ha="right")
plt.ylabel("Test Macro-F1")
plt.title("Effect of Metadata Embedding Dimension")
plt.tight_layout()
plt.show()
No description has been provided for this image
In [44]:
result_256 = run_metadata_encoder_experiment(
    experiment_name="image_all_metadata_encoded256",
    selected_features=["age", "sex", "location"],
    metadata_embed_dim=256,
    num_epochs=20
)

result_256
[image_all_metadata_encoded256] epoch 1/20 | train_f1=0.4688 | val_f1=0.5351 | val_bal_acc=0.6231
[image_all_metadata_encoded256] epoch 2/20 | train_f1=0.6871 | val_f1=0.5486 | val_bal_acc=0.6292
[image_all_metadata_encoded256] epoch 3/20 | train_f1=0.7793 | val_f1=0.5620 | val_bal_acc=0.6166
[image_all_metadata_encoded256] epoch 4/20 | train_f1=0.8137 | val_f1=0.6022 | val_bal_acc=0.6022
[image_all_metadata_encoded256] epoch 5/20 | train_f1=0.8586 | val_f1=0.5520 | val_bal_acc=0.6214
[image_all_metadata_encoded256] epoch 6/20 | train_f1=0.8655 | val_f1=0.5681 | val_bal_acc=0.6134
[image_all_metadata_encoded256] epoch 7/20 | train_f1=0.8864 | val_f1=0.5741 | val_bal_acc=0.5772
[image_all_metadata_encoded256] epoch 8/20 | train_f1=0.9055 | val_f1=0.5708 | val_bal_acc=0.5715
[image_all_metadata_encoded256] epoch 9/20 | train_f1=0.9188 | val_f1=0.5711 | val_bal_acc=0.5915
[image_all_metadata_encoded256] epoch 10/20 | train_f1=0.9272 | val_f1=0.6053 | val_bal_acc=0.5925
[image_all_metadata_encoded256] epoch 11/20 | train_f1=0.9376 | val_f1=0.6004 | val_bal_acc=0.5948
[image_all_metadata_encoded256] epoch 12/20 | train_f1=0.9409 | val_f1=0.5996 | val_bal_acc=0.5951
[image_all_metadata_encoded256] epoch 13/20 | train_f1=0.9531 | val_f1=0.5838 | val_bal_acc=0.5805
[image_all_metadata_encoded256] epoch 14/20 | train_f1=0.9590 | val_f1=0.5999 | val_bal_acc=0.5934
[image_all_metadata_encoded256] epoch 15/20 | train_f1=0.9626 | val_f1=0.5759 | val_bal_acc=0.5678
[image_all_metadata_encoded256] epoch 16/20 | train_f1=0.9648 | val_f1=0.5857 | val_bal_acc=0.5757
[image_all_metadata_encoded256] epoch 17/20 | train_f1=0.9694 | val_f1=0.5917 | val_bal_acc=0.5647
[image_all_metadata_encoded256] epoch 18/20 | train_f1=0.9707 | val_f1=0.5884 | val_bal_acc=0.5648
[image_all_metadata_encoded256] epoch 19/20 | train_f1=0.9786 | val_f1=0.5841 | val_bal_acc=0.5880
[image_all_metadata_encoded256] epoch 20/20 | train_f1=0.9809 | val_f1=0.5989 | val_bal_acc=0.5741
Out[44]:
{'name': 'image_all_metadata_encoded256',
 'features': ['age', 'sex', 'location'],
 'metadata_dim': 19,
 'metadata_embed_dim': 256,
 'best_val_macro_f1': 0.6052600613888408,
 'test_accuracy': 0.8041863605671843,
 'test_balanced_accuracy': 0.5939975839433717,
 'test_macro_f1': 0.5921817771622297,
 'history': [{'epoch': 1,
   'train_loss': 1.352913881499098,
   'train_accuracy': 0.6652385032847757,
   'train_macro_f1': 0.46880625023212286,
   'val_loss': 0.8628827635364184,
   'val_accuracy': 0.6899477806788512,
   'val_balanced_accuracy': 0.6231305303835579,
   'val_macro_f1': 0.5350691450567766},
  {'epoch': 2,
   'train_loss': 0.6303250630458264,
   'train_accuracy': 0.7767780634104542,
   'train_macro_f1': 0.6870592602352151,
   'val_loss': 0.8162107201142349,
   'val_accuracy': 0.6886422976501305,
   'val_balanced_accuracy': 0.6291837911190326,
   'val_macro_f1': 0.5486212283553559},
  {'epoch': 3,
   'train_loss': 0.41193746244420465,
   'train_accuracy': 0.824764353041988,
   'train_macro_f1': 0.7793244966961235,
   'val_loss': 0.7041458753825168,
   'val_accuracy': 0.7473890339425587,
   'val_balanced_accuracy': 0.6166200434087853,
   'val_macro_f1': 0.5619804261222198},
  {'epoch': 4,
   'train_loss': 0.3230803258008393,
   'train_accuracy': 0.8457583547557841,
   'train_macro_f1': 0.8136845473736318,
   'val_loss': 0.6803763021299173,
   'val_accuracy': 0.7526109660574413,
   'val_balanced_accuracy': 0.6021878423872209,
   'val_macro_f1': 0.6022489750485691},
  {'epoch': 5,
   'train_loss': 0.2682647559258707,
   'train_accuracy': 0.8721793773207654,
   'train_macro_f1': 0.8586025700476955,
   'val_loss': 0.8132532330973964,
   'val_accuracy': 0.70822454308094,
   'val_balanced_accuracy': 0.6214067954956846,
   'val_macro_f1': 0.5520438509159691},
  {'epoch': 6,
   'train_loss': 0.21354676806919917,
   'train_accuracy': 0.8854612967723507,
   'train_macro_f1': 0.8654556956641717,
   'val_loss': 0.8368096381581484,
   'val_accuracy': 0.7173629242819843,
   'val_balanced_accuracy': 0.6134213702819871,
   'val_macro_f1': 0.5680681038536705},
  {'epoch': 7,
   'train_loss': 0.1872922113174304,
   'train_accuracy': 0.8940302770636961,
   'train_macro_f1': 0.8864249104705886,
   'val_loss': 0.7722278858581193,
   'val_accuracy': 0.7356396866840731,
   'val_balanced_accuracy': 0.5771700855277967,
   'val_macro_f1': 0.574093240581487},
  {'epoch': 8,
   'train_loss': 0.15816285639205377,
   'train_accuracy': 0.9081690945444159,
   'train_macro_f1': 0.9054880507129834,
   'val_loss': 0.8378136268970082,
   'val_accuracy': 0.7265013054830287,
   'val_balanced_accuracy': 0.5715461073728038,
   'val_macro_f1': 0.5708212973164419},
  {'epoch': 9,
   'train_loss': 0.13744294162338477,
   'train_accuracy': 0.9184518708940302,
   'train_macro_f1': 0.9188192448262782,
   'val_loss': 0.8692586455134281,
   'val_accuracy': 0.7271540469973891,
   'val_balanced_accuracy': 0.5914779564691794,
   'val_macro_f1': 0.5710982384370265},
  {'epoch': 10,
   'train_loss': 0.1137242777471507,
   'train_accuracy': 0.9271636675235647,
   'train_macro_f1': 0.9271527900907918,
   'val_loss': 0.7693033254712665,
   'val_accuracy': 0.7898172323759791,
   'val_balanced_accuracy': 0.5925027022073825,
   'val_macro_f1': 0.6052600613888408},
  {'epoch': 11,
   'train_loss': 0.10518049037786797,
   'train_accuracy': 0.932447872036561,
   'train_macro_f1': 0.9376480156636964,
   'val_loss': 0.7839548142174595,
   'val_accuracy': 0.7774151436031331,
   'val_balanced_accuracy': 0.5947992593128967,
   'val_macro_f1': 0.6004041166587201},
  {'epoch': 12,
   'train_loss': 0.09531168254308874,
   'train_accuracy': 0.9367323621822337,
   'train_macro_f1': 0.9408931946638842,
   'val_loss': 0.8802234442325858,
   'val_accuracy': 0.762402088772846,
   'val_balanced_accuracy': 0.5951315411987812,
   'val_macro_f1': 0.59956768464803},
  {'epoch': 13,
   'train_loss': 0.08338177378266258,
   'train_accuracy': 0.9474435875464153,
   'train_macro_f1': 0.9530654747556888,
   'val_loss': 0.9786938316633745,
   'val_accuracy': 0.7297650130548303,
   'val_balanced_accuracy': 0.5804650159628568,
   'val_macro_f1': 0.5838314261708452},
  {'epoch': 14,
   'train_loss': 0.06680626502344095,
   'train_accuracy': 0.9551556698086261,
   'train_macro_f1': 0.9589962620743984,
   'val_loss': 0.9116787690092997,
   'val_accuracy': 0.7669712793733682,
   'val_balanced_accuracy': 0.5934295808621081,
   'val_macro_f1': 0.5998608371547743},
  {'epoch': 15,
   'train_loss': 0.066234799880046,
   'train_accuracy': 0.9545844044558698,
   'train_macro_f1': 0.9626183858512709,
   'val_loss': 0.9995350481567822,
   'val_accuracy': 0.7460835509138382,
   'val_balanced_accuracy': 0.5678310112122485,
   'val_macro_f1': 0.5758791268218397},
  {'epoch': 16,
   'train_loss': 0.06322339564477783,
   'train_accuracy': 0.9578691802342187,
   'train_macro_f1': 0.9648397590277972,
   'val_loss': 0.9599722753924514,
   'val_accuracy': 0.7643603133159269,
   'val_balanced_accuracy': 0.5757035106511016,
   'val_macro_f1': 0.5857152822584181},
  {'epoch': 17,
   'train_loss': 0.050073418022079944,
   'train_accuracy': 0.963010568409026,
   'train_macro_f1': 0.9693934390680655,
   'val_loss': 0.9452066939739345,
   'val_accuracy': 0.77088772845953,
   'val_balanced_accuracy': 0.5646819289597449,
   'val_macro_f1': 0.5916529022555528},
  {'epoch': 18,
   'train_loss': 0.04483924051124248,
   'train_accuracy': 0.9705798343330477,
   'train_macro_f1': 0.9707068516883075,
   'val_loss': 1.0637989053548271,
   'val_accuracy': 0.7552219321148825,
   'val_balanced_accuracy': 0.5648423065309214,
   'val_macro_f1': 0.5883831334936879},
  {'epoch': 19,
   'train_loss': 0.03535685211941983,
   'train_accuracy': 0.9762924878606113,
   'train_macro_f1': 0.9785525632902224,
   'val_loss': 1.1463026389541644,
   'val_accuracy': 0.7473890339425587,
   'val_balanced_accuracy': 0.5879559118887961,
   'val_macro_f1': 0.5840859922256031},
  {'epoch': 20,
   'train_loss': 0.031908363061436996,
   'train_accuracy': 0.9787203656098258,
   'train_macro_f1': 0.9808583281422241,
   'val_loss': 1.1057374976540724,
   'val_accuracy': 0.7741514360313316,
   'val_balanced_accuracy': 0.5740750336839182,
   'val_macro_f1': 0.5988664765168278}]}
In [45]:
extra_row = pd.DataFrame([{
    "Method": result_256["name"],
    "Features": ",".join(result_256["features"]),
    "Metadata Dim": result_256["metadata_dim"],
    "Metadata Embed Dim": result_256["metadata_embed_dim"],
    "Best Val Macro-F1": result_256["best_val_macro_f1"],
    "Test Accuracy": result_256["test_accuracy"],
    "Test Balanced Accuracy": result_256["test_balanced_accuracy"],
    "Test Macro-F1": result_256["test_macro_f1"],
}])

embed_dim_df = pd.concat([embed_dim_df, extra_row], ignore_index=True)
embed_dim_df = embed_dim_df.sort_values(
    by="Test Macro-F1",
    ascending=False
).reset_index(drop=True)

embed_dim_df
Out[45]:
Method Features Metadata Dim Metadata Embed Dim Best Val Macro-F1 Test Accuracy Test Balanced Accuracy Test Macro-F1
0 image_all_metadata_encoded128 age,sex,location 19 128 0.603923 0.799460 0.594053 0.600579
1 image_all_metadata_encoded64 age,sex,location 19 64 0.605293 0.804186 0.581747 0.593228
2 image_all_metadata_encoded256 age,sex,location 19 256 0.605260 0.804186 0.593998 0.592182
3 image_all_metadata_encoded32 age,sex,location 19 32 0.603866 0.794733 0.581390 0.582997
In [ ]: