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()
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()
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 [ ]: