本实验主要测试每个因素叠加所使用的最佳权重维度,只有6个实验,分别用raw,32,64,128进行测试,全部数据的之前已经测过了,在128的时候最好¶

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")
SUPPORT_DIR = os.path.join(PROJECT_DIR, "supports")
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],
}

train_age_mean = train_df["age"].mean()

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 [5]:
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"]},
]

embed_dims_to_try = [32, 64, 128]

metadata_experiments
Out[5]:
[{'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']}]
In [6]:
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]
    )
])

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]:
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 [9]:
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 [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 [17]:
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 [13]:
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
)

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 [14]:
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 [15]:
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 [16]:
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=BATCH_SIZE,
        shuffle=True,
        num_workers=0
    )
    val_loader = DataLoader(
        val_dataset,
        batch_size=BATCH_SIZE,
        shuffle=False,
        num_workers=0
    )
    test_loader = DataLoader(
        test_dataset,
        batch_size=BATCH_SIZE,
        shuffle=False,
        num_workers=0
    )

    return train_loader, val_loader, test_loader, train_metadata.shape[1]
In [18]:
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 [19]:
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),
    }


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 [20]:
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 [21]:
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 [22]:
trial_result = run_metadata_encoder_experiment(
    experiment_name="image_age_encoded32",
    selected_features=["age"],
    metadata_embed_dim=32,
    num_epochs=5
)

trial_result
[image_age_encoded32] epoch 1/5 | train_f1=0.4757 | val_f1=0.4573 | val_bal_acc=0.5159
[image_age_encoded32] epoch 2/5 | train_f1=0.7009 | val_f1=0.5360 | val_bal_acc=0.6046
[image_age_encoded32] epoch 3/5 | train_f1=0.7846 | val_f1=0.5353 | val_bal_acc=0.5824
[image_age_encoded32] epoch 4/5 | train_f1=0.8265 | val_f1=0.5524 | val_bal_acc=0.5819
[image_age_encoded32] epoch 5/5 | train_f1=0.8549 | val_f1=0.5591 | val_bal_acc=0.5878
Out[22]:
{'name': 'image_age_encoded32',
 'features': ['age'],
 'metadata_dim': 1,
 'metadata_embed_dim': 32,
 'best_val_macro_f1': 0.5590734262658515,
 'test_accuracy': 0.7643484132343011,
 'test_balanced_accuracy': 0.6185011277441932,
 'test_macro_f1': 0.5769331456378782,
 'history': [{'epoch': 1,
   'train_loss': 1.3757865863743934,
   'train_accuracy': 0.7076549557269352,
   'train_macro_f1': 0.4757486695259101,
   'val_loss': 0.9205009663385137,
   'val_accuracy': 0.6755874673629243,
   'val_balanced_accuracy': 0.5158705515494769,
   'val_macro_f1': 0.4572665622579147},
  {'epoch': 2,
   'train_loss': 0.6629017387771906,
   'train_accuracy': 0.7853470437017995,
   'train_macro_f1': 0.7008942334104724,
   'val_loss': 0.811258202355462,
   'val_accuracy': 0.6951697127937336,
   'val_balanced_accuracy': 0.604645605174234,
   'val_macro_f1': 0.5359549649618064},
  {'epoch': 3,
   'train_loss': 0.4142191265947783,
   'train_accuracy': 0.8303341902313625,
   'train_macro_f1': 0.7846178056249696,
   'val_loss': 0.889071045594178,
   'val_accuracy': 0.6690600522193212,
   'val_balanced_accuracy': 0.5824252454612833,
   'val_macro_f1': 0.5352937918190129},
  {'epoch': 4,
   'train_loss': 0.30858423042658295,
   'train_accuracy': 0.8533276206798057,
   'train_macro_f1': 0.8264624074773869,
   'val_loss': 0.8315554425392698,
   'val_accuracy': 0.6945169712793734,
   'val_balanced_accuracy': 0.5819068126572706,
   'val_macro_f1': 0.5523535341053204},
  {'epoch': 5,
   'train_loss': 0.24944227131153304,
   'train_accuracy': 0.8727506426735219,
   'train_macro_f1': 0.8549049486049007,
   'val_loss': 0.7976379494715298,
   'val_accuracy': 0.7271540469973891,
   'val_balanced_accuracy': 0.5878216285337252,
   'val_macro_f1': 0.5590734262658515}]}
In [23]:
# Cell A: 定义这次完整实验搜索空间
SEARCH_EXPERIMENTS = [
    {"name": "image_age", "features": ["age"], "metadata_dim": 1},
    {"name": "image_sex", "features": ["sex"], "metadata_dim": 3},
    {"name": "image_location", "features": ["location"], "metadata_dim": 15},
    {"name": "image_age_sex", "features": ["age", "sex"], "metadata_dim": 4},
    {"name": "image_age_location", "features": ["age", "location"], "metadata_dim": 16},
    {"name": "image_sex_location", "features": ["sex", "location"], "metadata_dim": 18},
]

EMBED_DIMS = [32, 64, 128]

SEARCH_EXPERIMENTS
Out[23]:
[{'name': 'image_age', 'features': ['age'], 'metadata_dim': 1},
 {'name': 'image_sex', 'features': ['sex'], 'metadata_dim': 3},
 {'name': 'image_location', 'features': ['location'], 'metadata_dim': 15},
 {'name': 'image_age_sex', 'features': ['age', 'sex'], 'metadata_dim': 4},
 {'name': 'image_age_location',
  'features': ['age', 'location'],
  'metadata_dim': 16},
 {'name': 'image_sex_location',
  'features': ['sex', 'location'],
  'metadata_dim': 18}]
In [25]:
# Cell B: 一次性完整跑完
all_search_results = []

for exp in SEARCH_EXPERIMENTS:
    for embed_dim in EMBED_DIMS:
        exp_name = f"{exp['name']}_encoded{embed_dim}"
        print(f"\n===== Running {exp_name} =====")
        
        result = run_metadata_encoder_experiment(
            experiment_name=exp_name,
            selected_features=exp["features"],
            metadata_embed_dim=embed_dim,
            num_epochs=20  # 先5轮快速筛;如果你想更稳可以改10
        )
        
        all_search_results.append({
            "method": exp_name,
            "base_method": exp["name"],
            "features": ",".join(exp["features"]),
            "metadata_dim": exp["metadata_dim"],
            "metadata_embed_dim": 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"],
            "history": result["history"],
        })

len(all_search_results)
===== Running image_age_encoded32 =====
[image_age_encoded32] epoch 1/20 | train_f1=0.4541 | val_f1=0.5271 | val_bal_acc=0.5556
[image_age_encoded32] epoch 2/20 | train_f1=0.7024 | val_f1=0.5259 | val_bal_acc=0.5920
[image_age_encoded32] epoch 3/20 | train_f1=0.7896 | val_f1=0.5261 | val_bal_acc=0.5875
[image_age_encoded32] epoch 4/20 | train_f1=0.8233 | val_f1=0.5572 | val_bal_acc=0.6171
[image_age_encoded32] epoch 5/20 | train_f1=0.8552 | val_f1=0.5665 | val_bal_acc=0.5910
[image_age_encoded32] epoch 6/20 | train_f1=0.8794 | val_f1=0.5601 | val_bal_acc=0.5931
[image_age_encoded32] epoch 7/20 | train_f1=0.9030 | val_f1=0.5734 | val_bal_acc=0.5858
[image_age_encoded32] epoch 8/20 | train_f1=0.9045 | val_f1=0.5734 | val_bal_acc=0.5886
[image_age_encoded32] epoch 9/20 | train_f1=0.9223 | val_f1=0.5687 | val_bal_acc=0.5662
[image_age_encoded32] epoch 10/20 | train_f1=0.9325 | val_f1=0.5707 | val_bal_acc=0.5719
[image_age_encoded32] epoch 11/20 | train_f1=0.9401 | val_f1=0.5746 | val_bal_acc=0.5792
[image_age_encoded32] epoch 12/20 | train_f1=0.9397 | val_f1=0.5832 | val_bal_acc=0.5769
[image_age_encoded32] epoch 13/20 | train_f1=0.9571 | val_f1=0.5826 | val_bal_acc=0.5827
[image_age_encoded32] epoch 14/20 | train_f1=0.9614 | val_f1=0.5899 | val_bal_acc=0.5929
[image_age_encoded32] epoch 15/20 | train_f1=0.9670 | val_f1=0.5635 | val_bal_acc=0.5396
[image_age_encoded32] epoch 16/20 | train_f1=0.9706 | val_f1=0.5925 | val_bal_acc=0.5799
[image_age_encoded32] epoch 17/20 | train_f1=0.9751 | val_f1=0.5779 | val_bal_acc=0.5765
[image_age_encoded32] epoch 18/20 | train_f1=0.9771 | val_f1=0.5636 | val_bal_acc=0.5782
[image_age_encoded32] epoch 19/20 | train_f1=0.9799 | val_f1=0.5770 | val_bal_acc=0.5709
[image_age_encoded32] epoch 20/20 | train_f1=0.9849 | val_f1=0.5786 | val_bal_acc=0.5702

===== Running image_age_encoded64 =====
[image_age_encoded64] epoch 1/20 | train_f1=0.4880 | val_f1=0.4827 | val_bal_acc=0.6001
[image_age_encoded64] epoch 2/20 | train_f1=0.7025 | val_f1=0.5486 | val_bal_acc=0.5965
[image_age_encoded64] epoch 3/20 | train_f1=0.7846 | val_f1=0.5670 | val_bal_acc=0.5885
[image_age_encoded64] epoch 4/20 | train_f1=0.8173 | val_f1=0.5494 | val_bal_acc=0.5948
[image_age_encoded64] epoch 5/20 | train_f1=0.8534 | val_f1=0.5551 | val_bal_acc=0.5919
[image_age_encoded64] epoch 6/20 | train_f1=0.8818 | val_f1=0.5653 | val_bal_acc=0.5940
[image_age_encoded64] epoch 7/20 | train_f1=0.8920 | val_f1=0.5930 | val_bal_acc=0.6022
[image_age_encoded64] epoch 8/20 | train_f1=0.9047 | val_f1=0.5896 | val_bal_acc=0.5983
[image_age_encoded64] epoch 9/20 | train_f1=0.9243 | val_f1=0.5777 | val_bal_acc=0.5893
[image_age_encoded64] epoch 10/20 | train_f1=0.9302 | val_f1=0.5889 | val_bal_acc=0.5746
[image_age_encoded64] epoch 11/20 | train_f1=0.9426 | val_f1=0.5819 | val_bal_acc=0.5919
[image_age_encoded64] epoch 12/20 | train_f1=0.9467 | val_f1=0.5828 | val_bal_acc=0.5739
[image_age_encoded64] epoch 13/20 | train_f1=0.9515 | val_f1=0.5797 | val_bal_acc=0.5744
[image_age_encoded64] epoch 14/20 | train_f1=0.9596 | val_f1=0.5736 | val_bal_acc=0.5709
[image_age_encoded64] epoch 15/20 | train_f1=0.9683 | val_f1=0.5768 | val_bal_acc=0.5537
[image_age_encoded64] epoch 16/20 | train_f1=0.9723 | val_f1=0.5789 | val_bal_acc=0.5704
[image_age_encoded64] epoch 17/20 | train_f1=0.9749 | val_f1=0.5804 | val_bal_acc=0.5758
[image_age_encoded64] epoch 18/20 | train_f1=0.9746 | val_f1=0.5930 | val_bal_acc=0.5802
[image_age_encoded64] epoch 19/20 | train_f1=0.9802 | val_f1=0.5858 | val_bal_acc=0.5545
[image_age_encoded64] epoch 20/20 | train_f1=0.9843 | val_f1=0.5802 | val_bal_acc=0.5650

===== Running image_age_encoded128 =====
[image_age_encoded128] epoch 1/20 | train_f1=0.4682 | val_f1=0.4859 | val_bal_acc=0.5945
[image_age_encoded128] epoch 2/20 | train_f1=0.7088 | val_f1=0.5272 | val_bal_acc=0.6037
[image_age_encoded128] epoch 3/20 | train_f1=0.7794 | val_f1=0.5501 | val_bal_acc=0.6021
[image_age_encoded128] epoch 4/20 | train_f1=0.8314 | val_f1=0.5784 | val_bal_acc=0.6013
[image_age_encoded128] epoch 5/20 | train_f1=0.8519 | val_f1=0.5756 | val_bal_acc=0.6042
[image_age_encoded128] epoch 6/20 | train_f1=0.8798 | val_f1=0.5926 | val_bal_acc=0.6113
[image_age_encoded128] epoch 7/20 | train_f1=0.8953 | val_f1=0.5824 | val_bal_acc=0.5810
[image_age_encoded128] epoch 8/20 | train_f1=0.9025 | val_f1=0.5630 | val_bal_acc=0.5984
[image_age_encoded128] epoch 9/20 | train_f1=0.9178 | val_f1=0.5610 | val_bal_acc=0.5808
[image_age_encoded128] epoch 10/20 | train_f1=0.9317 | val_f1=0.5796 | val_bal_acc=0.5738
[image_age_encoded128] epoch 11/20 | train_f1=0.9469 | val_f1=0.5919 | val_bal_acc=0.6007
[image_age_encoded128] epoch 12/20 | train_f1=0.9511 | val_f1=0.5767 | val_bal_acc=0.5925
[image_age_encoded128] epoch 13/20 | train_f1=0.9531 | val_f1=0.5833 | val_bal_acc=0.5697
[image_age_encoded128] epoch 14/20 | train_f1=0.9588 | val_f1=0.5776 | val_bal_acc=0.5841
[image_age_encoded128] epoch 15/20 | train_f1=0.9653 | val_f1=0.5810 | val_bal_acc=0.5664
[image_age_encoded128] epoch 16/20 | train_f1=0.9714 | val_f1=0.5880 | val_bal_acc=0.5827
[image_age_encoded128] epoch 17/20 | train_f1=0.9707 | val_f1=0.5876 | val_bal_acc=0.5621
[image_age_encoded128] epoch 18/20 | train_f1=0.9767 | val_f1=0.5661 | val_bal_acc=0.5614
[image_age_encoded128] epoch 19/20 | train_f1=0.9808 | val_f1=0.5833 | val_bal_acc=0.5631
[image_age_encoded128] epoch 20/20 | train_f1=0.9840 | val_f1=0.5944 | val_bal_acc=0.5836

===== Running image_sex_encoded32 =====
[image_sex_encoded32] epoch 1/20 | train_f1=0.4535 | val_f1=0.4711 | val_bal_acc=0.5526
[image_sex_encoded32] epoch 2/20 | train_f1=0.6985 | val_f1=0.4981 | val_bal_acc=0.5877
[image_sex_encoded32] epoch 3/20 | train_f1=0.7783 | val_f1=0.5326 | val_bal_acc=0.6021
[image_sex_encoded32] epoch 4/20 | train_f1=0.8187 | val_f1=0.5524 | val_bal_acc=0.6044
[image_sex_encoded32] epoch 5/20 | train_f1=0.8498 | val_f1=0.5500 | val_bal_acc=0.5601
[image_sex_encoded32] epoch 6/20 | train_f1=0.8754 | val_f1=0.5804 | val_bal_acc=0.5999
[image_sex_encoded32] epoch 7/20 | train_f1=0.8969 | val_f1=0.5740 | val_bal_acc=0.5953
[image_sex_encoded32] epoch 8/20 | train_f1=0.9019 | val_f1=0.5963 | val_bal_acc=0.5900
[image_sex_encoded32] epoch 9/20 | train_f1=0.9200 | val_f1=0.5802 | val_bal_acc=0.5861
[image_sex_encoded32] epoch 10/20 | train_f1=0.9343 | val_f1=0.5700 | val_bal_acc=0.5778
[image_sex_encoded32] epoch 11/20 | train_f1=0.9420 | val_f1=0.5676 | val_bal_acc=0.5785
[image_sex_encoded32] epoch 12/20 | train_f1=0.9463 | val_f1=0.5851 | val_bal_acc=0.5635
[image_sex_encoded32] epoch 13/20 | train_f1=0.9549 | val_f1=0.5831 | val_bal_acc=0.5786
[image_sex_encoded32] epoch 14/20 | train_f1=0.9638 | val_f1=0.5859 | val_bal_acc=0.5730
[image_sex_encoded32] epoch 15/20 | train_f1=0.9667 | val_f1=0.5742 | val_bal_acc=0.5551
[image_sex_encoded32] epoch 16/20 | train_f1=0.9728 | val_f1=0.5777 | val_bal_acc=0.5670
[image_sex_encoded32] epoch 17/20 | train_f1=0.9799 | val_f1=0.5800 | val_bal_acc=0.5602
[image_sex_encoded32] epoch 18/20 | train_f1=0.9801 | val_f1=0.5814 | val_bal_acc=0.5573
[image_sex_encoded32] epoch 19/20 | train_f1=0.9804 | val_f1=0.5776 | val_bal_acc=0.5723
[image_sex_encoded32] epoch 20/20 | train_f1=0.9863 | val_f1=0.5761 | val_bal_acc=0.5844

===== Running image_sex_encoded64 =====
[image_sex_encoded64] epoch 1/20 | train_f1=0.5009 | val_f1=0.4709 | val_bal_acc=0.5811
[image_sex_encoded64] epoch 2/20 | train_f1=0.7024 | val_f1=0.5091 | val_bal_acc=0.6132
[image_sex_encoded64] epoch 3/20 | train_f1=0.7710 | val_f1=0.5399 | val_bal_acc=0.6015
[image_sex_encoded64] epoch 4/20 | train_f1=0.8254 | val_f1=0.5515 | val_bal_acc=0.6053
[image_sex_encoded64] epoch 5/20 | train_f1=0.8545 | val_f1=0.5712 | val_bal_acc=0.5978
[image_sex_encoded64] epoch 6/20 | train_f1=0.8806 | val_f1=0.5726 | val_bal_acc=0.5973
[image_sex_encoded64] epoch 7/20 | train_f1=0.8970 | val_f1=0.5844 | val_bal_acc=0.6056
[image_sex_encoded64] epoch 8/20 | train_f1=0.9071 | val_f1=0.5847 | val_bal_acc=0.5861
[image_sex_encoded64] epoch 9/20 | train_f1=0.9180 | val_f1=0.6009 | val_bal_acc=0.5900
[image_sex_encoded64] epoch 10/20 | train_f1=0.9278 | val_f1=0.5859 | val_bal_acc=0.5916
[image_sex_encoded64] epoch 11/20 | train_f1=0.9429 | val_f1=0.5631 | val_bal_acc=0.5952
[image_sex_encoded64] epoch 12/20 | train_f1=0.9490 | val_f1=0.5943 | val_bal_acc=0.5885
[image_sex_encoded64] epoch 13/20 | train_f1=0.9585 | val_f1=0.5644 | val_bal_acc=0.5686
[image_sex_encoded64] epoch 14/20 | train_f1=0.9571 | val_f1=0.5921 | val_bal_acc=0.5863
[image_sex_encoded64] epoch 15/20 | train_f1=0.9671 | val_f1=0.5722 | val_bal_acc=0.5923
[image_sex_encoded64] epoch 16/20 | train_f1=0.9667 | val_f1=0.5869 | val_bal_acc=0.5664
[image_sex_encoded64] epoch 17/20 | train_f1=0.9750 | val_f1=0.5937 | val_bal_acc=0.5841
[image_sex_encoded64] epoch 18/20 | train_f1=0.9779 | val_f1=0.5857 | val_bal_acc=0.5595
[image_sex_encoded64] epoch 19/20 | train_f1=0.9767 | val_f1=0.5762 | val_bal_acc=0.5619
[image_sex_encoded64] epoch 20/20 | train_f1=0.9826 | val_f1=0.5865 | val_bal_acc=0.5608

===== Running image_sex_encoded128 =====
[image_sex_encoded128] epoch 1/20 | train_f1=0.4794 | val_f1=0.4701 | val_bal_acc=0.5340
[image_sex_encoded128] epoch 2/20 | train_f1=0.6931 | val_f1=0.5186 | val_bal_acc=0.5833
[image_sex_encoded128] epoch 3/20 | train_f1=0.7738 | val_f1=0.5780 | val_bal_acc=0.6091
[image_sex_encoded128] epoch 4/20 | train_f1=0.8119 | val_f1=0.5794 | val_bal_acc=0.5973
[image_sex_encoded128] epoch 5/20 | train_f1=0.8381 | val_f1=0.5595 | val_bal_acc=0.6125
[image_sex_encoded128] epoch 6/20 | train_f1=0.8702 | val_f1=0.5840 | val_bal_acc=0.5837
[image_sex_encoded128] epoch 7/20 | train_f1=0.8939 | val_f1=0.5753 | val_bal_acc=0.5864
[image_sex_encoded128] epoch 8/20 | train_f1=0.8988 | val_f1=0.5803 | val_bal_acc=0.5885
[image_sex_encoded128] epoch 9/20 | train_f1=0.9153 | val_f1=0.5893 | val_bal_acc=0.5808
[image_sex_encoded128] epoch 10/20 | train_f1=0.9173 | val_f1=0.5699 | val_bal_acc=0.5659
[image_sex_encoded128] epoch 11/20 | train_f1=0.9357 | val_f1=0.5893 | val_bal_acc=0.5705
[image_sex_encoded128] epoch 12/20 | train_f1=0.9464 | val_f1=0.5834 | val_bal_acc=0.5542
[image_sex_encoded128] epoch 13/20 | train_f1=0.9517 | val_f1=0.5717 | val_bal_acc=0.5743
[image_sex_encoded128] epoch 14/20 | train_f1=0.9572 | val_f1=0.5844 | val_bal_acc=0.5837
[image_sex_encoded128] epoch 15/20 | train_f1=0.9623 | val_f1=0.5908 | val_bal_acc=0.5916
[image_sex_encoded128] epoch 16/20 | train_f1=0.9661 | val_f1=0.5919 | val_bal_acc=0.5820
[image_sex_encoded128] epoch 17/20 | train_f1=0.9697 | val_f1=0.5994 | val_bal_acc=0.5952
[image_sex_encoded128] epoch 18/20 | train_f1=0.9724 | val_f1=0.5588 | val_bal_acc=0.5797
[image_sex_encoded128] epoch 19/20 | train_f1=0.9624 | val_f1=0.5802 | val_bal_acc=0.5632
[image_sex_encoded128] epoch 20/20 | train_f1=0.9738 | val_f1=0.5999 | val_bal_acc=0.5940

===== Running image_location_encoded32 =====
[image_location_encoded32] epoch 1/20 | train_f1=0.4579 | val_f1=0.4864 | val_bal_acc=0.6128
[image_location_encoded32] epoch 2/20 | train_f1=0.7034 | val_f1=0.5277 | val_bal_acc=0.5757
[image_location_encoded32] epoch 3/20 | train_f1=0.7886 | val_f1=0.5330 | val_bal_acc=0.6059
[image_location_encoded32] epoch 4/20 | train_f1=0.8299 | val_f1=0.5726 | val_bal_acc=0.6025
[image_location_encoded32] epoch 5/20 | train_f1=0.8606 | val_f1=0.5744 | val_bal_acc=0.5952
[image_location_encoded32] epoch 6/20 | train_f1=0.8896 | val_f1=0.5594 | val_bal_acc=0.5920
[image_location_encoded32] epoch 7/20 | train_f1=0.8974 | val_f1=0.5770 | val_bal_acc=0.5920
[image_location_encoded32] epoch 8/20 | train_f1=0.9136 | val_f1=0.5857 | val_bal_acc=0.5757
[image_location_encoded32] epoch 9/20 | train_f1=0.9200 | val_f1=0.5958 | val_bal_acc=0.5645
[image_location_encoded32] epoch 10/20 | train_f1=0.9315 | val_f1=0.5734 | val_bal_acc=0.5721
[image_location_encoded32] epoch 11/20 | train_f1=0.9447 | val_f1=0.5773 | val_bal_acc=0.5868
[image_location_encoded32] epoch 12/20 | train_f1=0.9520 | val_f1=0.5745 | val_bal_acc=0.5726
[image_location_encoded32] epoch 13/20 | train_f1=0.9586 | val_f1=0.5862 | val_bal_acc=0.5731
[image_location_encoded32] epoch 14/20 | train_f1=0.9671 | val_f1=0.6038 | val_bal_acc=0.5840
[image_location_encoded32] epoch 15/20 | train_f1=0.9700 | val_f1=0.5913 | val_bal_acc=0.5850
[image_location_encoded32] epoch 16/20 | train_f1=0.9749 | val_f1=0.5884 | val_bal_acc=0.5727
[image_location_encoded32] epoch 17/20 | train_f1=0.9787 | val_f1=0.5928 | val_bal_acc=0.5926
[image_location_encoded32] epoch 18/20 | train_f1=0.9826 | val_f1=0.6032 | val_bal_acc=0.5780
[image_location_encoded32] epoch 19/20 | train_f1=0.9858 | val_f1=0.5924 | val_bal_acc=0.5864
[image_location_encoded32] epoch 20/20 | train_f1=0.9827 | val_f1=0.5973 | val_bal_acc=0.5947

===== Running image_location_encoded64 =====
[image_location_encoded64] epoch 1/20 | train_f1=0.4689 | val_f1=0.5023 | val_bal_acc=0.6010
[image_location_encoded64] epoch 2/20 | train_f1=0.7020 | val_f1=0.5531 | val_bal_acc=0.6163
[image_location_encoded64] epoch 3/20 | train_f1=0.7709 | val_f1=0.5425 | val_bal_acc=0.5960
[image_location_encoded64] epoch 4/20 | train_f1=0.8226 | val_f1=0.5790 | val_bal_acc=0.6052
[image_location_encoded64] epoch 5/20 | train_f1=0.8481 | val_f1=0.5676 | val_bal_acc=0.6002
[image_location_encoded64] epoch 6/20 | train_f1=0.8720 | val_f1=0.5493 | val_bal_acc=0.5789
[image_location_encoded64] epoch 7/20 | train_f1=0.8942 | val_f1=0.5855 | val_bal_acc=0.5978
[image_location_encoded64] epoch 8/20 | train_f1=0.9134 | val_f1=0.5829 | val_bal_acc=0.6007
[image_location_encoded64] epoch 9/20 | train_f1=0.9224 | val_f1=0.5784 | val_bal_acc=0.6040
[image_location_encoded64] epoch 10/20 | train_f1=0.9282 | val_f1=0.6006 | val_bal_acc=0.5873
[image_location_encoded64] epoch 11/20 | train_f1=0.9417 | val_f1=0.5885 | val_bal_acc=0.5741
[image_location_encoded64] epoch 12/20 | train_f1=0.9433 | val_f1=0.5863 | val_bal_acc=0.5949
[image_location_encoded64] epoch 13/20 | train_f1=0.9493 | val_f1=0.5943 | val_bal_acc=0.5793
[image_location_encoded64] epoch 14/20 | train_f1=0.9616 | val_f1=0.5922 | val_bal_acc=0.5690
[image_location_encoded64] epoch 15/20 | train_f1=0.9647 | val_f1=0.5855 | val_bal_acc=0.5760
[image_location_encoded64] epoch 16/20 | train_f1=0.9672 | val_f1=0.5903 | val_bal_acc=0.5563
[image_location_encoded64] epoch 17/20 | train_f1=0.9739 | val_f1=0.5879 | val_bal_acc=0.5783
[image_location_encoded64] epoch 18/20 | train_f1=0.9784 | val_f1=0.5957 | val_bal_acc=0.5739
[image_location_encoded64] epoch 19/20 | train_f1=0.9808 | val_f1=0.5925 | val_bal_acc=0.5883
[image_location_encoded64] epoch 20/20 | train_f1=0.9850 | val_f1=0.5933 | val_bal_acc=0.5890

===== Running image_location_encoded128 =====
[image_location_encoded128] epoch 1/20 | train_f1=0.4894 | val_f1=0.5122 | val_bal_acc=0.5849
[image_location_encoded128] epoch 2/20 | train_f1=0.6916 | val_f1=0.5155 | val_bal_acc=0.5876
[image_location_encoded128] epoch 3/20 | train_f1=0.7741 | val_f1=0.5798 | val_bal_acc=0.6252
[image_location_encoded128] epoch 4/20 | train_f1=0.8169 | val_f1=0.5835 | val_bal_acc=0.6057
[image_location_encoded128] epoch 5/20 | train_f1=0.8525 | val_f1=0.5923 | val_bal_acc=0.5972
[image_location_encoded128] epoch 6/20 | train_f1=0.8746 | val_f1=0.5501 | val_bal_acc=0.5958
[image_location_encoded128] epoch 7/20 | train_f1=0.8901 | val_f1=0.5672 | val_bal_acc=0.5917
[image_location_encoded128] epoch 8/20 | train_f1=0.9090 | val_f1=0.5770 | val_bal_acc=0.5991
[image_location_encoded128] epoch 9/20 | train_f1=0.9166 | val_f1=0.5794 | val_bal_acc=0.5811
[image_location_encoded128] epoch 10/20 | train_f1=0.9279 | val_f1=0.5913 | val_bal_acc=0.5945
[image_location_encoded128] epoch 11/20 | train_f1=0.9375 | val_f1=0.5728 | val_bal_acc=0.5748
[image_location_encoded128] epoch 12/20 | train_f1=0.9438 | val_f1=0.5598 | val_bal_acc=0.5897
[image_location_encoded128] epoch 13/20 | train_f1=0.9534 | val_f1=0.5780 | val_bal_acc=0.5516
[image_location_encoded128] epoch 14/20 | train_f1=0.9612 | val_f1=0.5808 | val_bal_acc=0.5656
[image_location_encoded128] epoch 15/20 | train_f1=0.9623 | val_f1=0.5882 | val_bal_acc=0.5775
[image_location_encoded128] epoch 16/20 | train_f1=0.9704 | val_f1=0.5852 | val_bal_acc=0.5885
[image_location_encoded128] epoch 17/20 | train_f1=0.9728 | val_f1=0.5956 | val_bal_acc=0.5682
[image_location_encoded128] epoch 18/20 | train_f1=0.9761 | val_f1=0.5926 | val_bal_acc=0.5742
[image_location_encoded128] epoch 19/20 | train_f1=0.9788 | val_f1=0.5784 | val_bal_acc=0.5669
[image_location_encoded128] epoch 20/20 | train_f1=0.9847 | val_f1=0.5879 | val_bal_acc=0.5701

===== Running image_age_sex_encoded32 =====
[image_age_sex_encoded32] epoch 1/20 | train_f1=0.5106 | val_f1=0.5170 | val_bal_acc=0.5650
[image_age_sex_encoded32] epoch 2/20 | train_f1=0.7100 | val_f1=0.5387 | val_bal_acc=0.6245
[image_age_sex_encoded32] epoch 3/20 | train_f1=0.7831 | val_f1=0.5449 | val_bal_acc=0.5966
[image_age_sex_encoded32] epoch 4/20 | train_f1=0.8323 | val_f1=0.5759 | val_bal_acc=0.5947
[image_age_sex_encoded32] epoch 5/20 | train_f1=0.8732 | val_f1=0.5494 | val_bal_acc=0.5988
[image_age_sex_encoded32] epoch 6/20 | train_f1=0.8767 | val_f1=0.5577 | val_bal_acc=0.5767
[image_age_sex_encoded32] epoch 7/20 | train_f1=0.9000 | val_f1=0.5880 | val_bal_acc=0.6048
[image_age_sex_encoded32] epoch 8/20 | train_f1=0.9113 | val_f1=0.5778 | val_bal_acc=0.5858
[image_age_sex_encoded32] epoch 9/20 | train_f1=0.9225 | val_f1=0.5724 | val_bal_acc=0.6044
[image_age_sex_encoded32] epoch 10/20 | train_f1=0.9326 | val_f1=0.5894 | val_bal_acc=0.5692
[image_age_sex_encoded32] epoch 11/20 | train_f1=0.9364 | val_f1=0.6031 | val_bal_acc=0.5931
[image_age_sex_encoded32] epoch 12/20 | train_f1=0.9478 | val_f1=0.5767 | val_bal_acc=0.5517
[image_age_sex_encoded32] epoch 13/20 | train_f1=0.9563 | val_f1=0.5955 | val_bal_acc=0.5801
[image_age_sex_encoded32] epoch 14/20 | train_f1=0.9615 | val_f1=0.5855 | val_bal_acc=0.5814
[image_age_sex_encoded32] epoch 15/20 | train_f1=0.9666 | val_f1=0.5931 | val_bal_acc=0.5863
[image_age_sex_encoded32] epoch 16/20 | train_f1=0.9709 | val_f1=0.5859 | val_bal_acc=0.5690
[image_age_sex_encoded32] epoch 17/20 | train_f1=0.9750 | val_f1=0.5712 | val_bal_acc=0.5591
[image_age_sex_encoded32] epoch 18/20 | train_f1=0.9767 | val_f1=0.5693 | val_bal_acc=0.5451
[image_age_sex_encoded32] epoch 19/20 | train_f1=0.9798 | val_f1=0.5989 | val_bal_acc=0.5760
[image_age_sex_encoded32] epoch 20/20 | train_f1=0.9830 | val_f1=0.5870 | val_bal_acc=0.5670

===== Running image_age_sex_encoded64 =====
[image_age_sex_encoded64] epoch 1/20 | train_f1=0.4740 | val_f1=0.5141 | val_bal_acc=0.5560
[image_age_sex_encoded64] epoch 2/20 | train_f1=0.7063 | val_f1=0.5257 | val_bal_acc=0.6104
[image_age_sex_encoded64] epoch 3/20 | train_f1=0.7723 | val_f1=0.5421 | val_bal_acc=0.6029
[image_age_sex_encoded64] epoch 4/20 | train_f1=0.8282 | val_f1=0.5427 | val_bal_acc=0.6004
[image_age_sex_encoded64] epoch 5/20 | train_f1=0.8543 | val_f1=0.5758 | val_bal_acc=0.6038
[image_age_sex_encoded64] epoch 6/20 | train_f1=0.8735 | val_f1=0.5831 | val_bal_acc=0.6076
[image_age_sex_encoded64] epoch 7/20 | train_f1=0.8950 | val_f1=0.5957 | val_bal_acc=0.5852
[image_age_sex_encoded64] epoch 8/20 | train_f1=0.9070 | val_f1=0.6040 | val_bal_acc=0.5850
[image_age_sex_encoded64] epoch 9/20 | train_f1=0.9194 | val_f1=0.5712 | val_bal_acc=0.5728
[image_age_sex_encoded64] epoch 10/20 | train_f1=0.9245 | val_f1=0.5914 | val_bal_acc=0.5796
[image_age_sex_encoded64] epoch 11/20 | train_f1=0.9367 | val_f1=0.5870 | val_bal_acc=0.5915
[image_age_sex_encoded64] epoch 12/20 | train_f1=0.9527 | val_f1=0.5960 | val_bal_acc=0.5956
[image_age_sex_encoded64] epoch 13/20 | train_f1=0.9566 | val_f1=0.5962 | val_bal_acc=0.5964
[image_age_sex_encoded64] epoch 14/20 | train_f1=0.9642 | val_f1=0.5969 | val_bal_acc=0.5912
[image_age_sex_encoded64] epoch 15/20 | train_f1=0.9639 | val_f1=0.5901 | val_bal_acc=0.5904
[image_age_sex_encoded64] epoch 16/20 | train_f1=0.9678 | val_f1=0.5845 | val_bal_acc=0.5668
[image_age_sex_encoded64] epoch 17/20 | train_f1=0.9719 | val_f1=0.5647 | val_bal_acc=0.5655
[image_age_sex_encoded64] epoch 18/20 | train_f1=0.9736 | val_f1=0.5411 | val_bal_acc=0.6115
[image_age_sex_encoded64] epoch 19/20 | train_f1=0.9765 | val_f1=0.5793 | val_bal_acc=0.5673
[image_age_sex_encoded64] epoch 20/20 | train_f1=0.9818 | val_f1=0.5937 | val_bal_acc=0.5897

===== Running image_age_sex_encoded128 =====
[image_age_sex_encoded128] epoch 1/20 | train_f1=0.4801 | val_f1=0.4978 | val_bal_acc=0.6072
[image_age_sex_encoded128] epoch 2/20 | train_f1=0.6938 | val_f1=0.5327 | val_bal_acc=0.6106
[image_age_sex_encoded128] epoch 3/20 | train_f1=0.7719 | val_f1=0.5367 | val_bal_acc=0.6211
[image_age_sex_encoded128] epoch 4/20 | train_f1=0.8226 | val_f1=0.5425 | val_bal_acc=0.6273
[image_age_sex_encoded128] epoch 5/20 | train_f1=0.8493 | val_f1=0.5425 | val_bal_acc=0.5955
[image_age_sex_encoded128] epoch 6/20 | train_f1=0.8802 | val_f1=0.5931 | val_bal_acc=0.5965
[image_age_sex_encoded128] epoch 7/20 | train_f1=0.8971 | val_f1=0.5638 | val_bal_acc=0.5864
[image_age_sex_encoded128] epoch 8/20 | train_f1=0.9075 | val_f1=0.5811 | val_bal_acc=0.5819
[image_age_sex_encoded128] epoch 9/20 | train_f1=0.9127 | val_f1=0.6021 | val_bal_acc=0.6025
[image_age_sex_encoded128] epoch 10/20 | train_f1=0.9221 | val_f1=0.5754 | val_bal_acc=0.5805
[image_age_sex_encoded128] epoch 11/20 | train_f1=0.9412 | val_f1=0.5932 | val_bal_acc=0.5932
[image_age_sex_encoded128] epoch 12/20 | train_f1=0.9450 | val_f1=0.5804 | val_bal_acc=0.5765
[image_age_sex_encoded128] epoch 13/20 | train_f1=0.9495 | val_f1=0.5794 | val_bal_acc=0.5848
[image_age_sex_encoded128] epoch 14/20 | train_f1=0.9576 | val_f1=0.5881 | val_bal_acc=0.5807
[image_age_sex_encoded128] epoch 15/20 | train_f1=0.9670 | val_f1=0.5895 | val_bal_acc=0.5721
[image_age_sex_encoded128] epoch 16/20 | train_f1=0.9673 | val_f1=0.5905 | val_bal_acc=0.5898
[image_age_sex_encoded128] epoch 17/20 | train_f1=0.9723 | val_f1=0.6027 | val_bal_acc=0.5906
[image_age_sex_encoded128] epoch 18/20 | train_f1=0.9715 | val_f1=0.5853 | val_bal_acc=0.5653
[image_age_sex_encoded128] epoch 19/20 | train_f1=0.9742 | val_f1=0.5792 | val_bal_acc=0.5741
[image_age_sex_encoded128] epoch 20/20 | train_f1=0.9809 | val_f1=0.5754 | val_bal_acc=0.5677

===== Running image_age_location_encoded32 =====
[image_age_location_encoded32] epoch 1/20 | train_f1=0.4830 | val_f1=0.5283 | val_bal_acc=0.6027
[image_age_location_encoded32] epoch 2/20 | train_f1=0.7254 | val_f1=0.5244 | val_bal_acc=0.5915
[image_age_location_encoded32] epoch 3/20 | train_f1=0.7870 | val_f1=0.5626 | val_bal_acc=0.5977
[image_age_location_encoded32] epoch 4/20 | train_f1=0.8391 | val_f1=0.5450 | val_bal_acc=0.5803
[image_age_location_encoded32] epoch 5/20 | train_f1=0.8640 | val_f1=0.5738 | val_bal_acc=0.5906
[image_age_location_encoded32] epoch 6/20 | train_f1=0.8884 | val_f1=0.5906 | val_bal_acc=0.6047
[image_age_location_encoded32] epoch 7/20 | train_f1=0.8955 | val_f1=0.5892 | val_bal_acc=0.6056
[image_age_location_encoded32] epoch 8/20 | train_f1=0.9128 | val_f1=0.6022 | val_bal_acc=0.6066
[image_age_location_encoded32] epoch 9/20 | train_f1=0.9291 | val_f1=0.5892 | val_bal_acc=0.5852
[image_age_location_encoded32] epoch 10/20 | train_f1=0.9377 | val_f1=0.5875 | val_bal_acc=0.5632
[image_age_location_encoded32] epoch 11/20 | train_f1=0.9441 | val_f1=0.5891 | val_bal_acc=0.5823
[image_age_location_encoded32] epoch 12/20 | train_f1=0.9565 | val_f1=0.6031 | val_bal_acc=0.5872
[image_age_location_encoded32] epoch 13/20 | train_f1=0.9635 | val_f1=0.5898 | val_bal_acc=0.5917
[image_age_location_encoded32] epoch 14/20 | train_f1=0.9625 | val_f1=0.6045 | val_bal_acc=0.5851
[image_age_location_encoded32] epoch 15/20 | train_f1=0.9676 | val_f1=0.5650 | val_bal_acc=0.5921
[image_age_location_encoded32] epoch 16/20 | train_f1=0.9704 | val_f1=0.5863 | val_bal_acc=0.5790
[image_age_location_encoded32] epoch 17/20 | train_f1=0.9770 | val_f1=0.5887 | val_bal_acc=0.5706
[image_age_location_encoded32] epoch 18/20 | train_f1=0.9746 | val_f1=0.5928 | val_bal_acc=0.5809
[image_age_location_encoded32] epoch 19/20 | train_f1=0.9831 | val_f1=0.5890 | val_bal_acc=0.5970
[image_age_location_encoded32] epoch 20/20 | train_f1=0.9847 | val_f1=0.5976 | val_bal_acc=0.5815

===== Running image_age_location_encoded64 =====
[image_age_location_encoded64] epoch 1/20 | train_f1=0.4855 | val_f1=0.5385 | val_bal_acc=0.5974
[image_age_location_encoded64] epoch 2/20 | train_f1=0.7095 | val_f1=0.5459 | val_bal_acc=0.6065
[image_age_location_encoded64] epoch 3/20 | train_f1=0.7730 | val_f1=0.5723 | val_bal_acc=0.5852
[image_age_location_encoded64] epoch 4/20 | train_f1=0.8275 | val_f1=0.5843 | val_bal_acc=0.6092
[image_age_location_encoded64] epoch 5/20 | train_f1=0.8622 | val_f1=0.5560 | val_bal_acc=0.5951
[image_age_location_encoded64] epoch 6/20 | train_f1=0.8896 | val_f1=0.5924 | val_bal_acc=0.5971
[image_age_location_encoded64] epoch 7/20 | train_f1=0.9107 | val_f1=0.5917 | val_bal_acc=0.6013
[image_age_location_encoded64] epoch 8/20 | train_f1=0.9143 | val_f1=0.5972 | val_bal_acc=0.6078
[image_age_location_encoded64] epoch 9/20 | train_f1=0.9243 | val_f1=0.5710 | val_bal_acc=0.5926
[image_age_location_encoded64] epoch 10/20 | train_f1=0.9357 | val_f1=0.5849 | val_bal_acc=0.5874
[image_age_location_encoded64] epoch 11/20 | train_f1=0.9477 | val_f1=0.5990 | val_bal_acc=0.5803
[image_age_location_encoded64] epoch 12/20 | train_f1=0.9526 | val_f1=0.5753 | val_bal_acc=0.5658
[image_age_location_encoded64] epoch 13/20 | train_f1=0.9605 | val_f1=0.5892 | val_bal_acc=0.5525
[image_age_location_encoded64] epoch 14/20 | train_f1=0.9631 | val_f1=0.5648 | val_bal_acc=0.5769
[image_age_location_encoded64] epoch 15/20 | train_f1=0.9668 | val_f1=0.5976 | val_bal_acc=0.5747
[image_age_location_encoded64] epoch 16/20 | train_f1=0.9724 | val_f1=0.5973 | val_bal_acc=0.5725
[image_age_location_encoded64] epoch 17/20 | train_f1=0.9743 | val_f1=0.5868 | val_bal_acc=0.5863
[image_age_location_encoded64] epoch 18/20 | train_f1=0.9794 | val_f1=0.5905 | val_bal_acc=0.5746
[image_age_location_encoded64] epoch 19/20 | train_f1=0.9832 | val_f1=0.5954 | val_bal_acc=0.5735
[image_age_location_encoded64] epoch 20/20 | train_f1=0.9860 | val_f1=0.5959 | val_bal_acc=0.5616

===== Running image_age_location_encoded128 =====
[image_age_location_encoded128] epoch 1/20 | train_f1=0.4778 | val_f1=0.4868 | val_bal_acc=0.6182
[image_age_location_encoded128] epoch 2/20 | train_f1=0.6964 | val_f1=0.5242 | val_bal_acc=0.6132
[image_age_location_encoded128] epoch 3/20 | train_f1=0.7921 | val_f1=0.5660 | val_bal_acc=0.6268
[image_age_location_encoded128] epoch 4/20 | train_f1=0.8167 | val_f1=0.5304 | val_bal_acc=0.5859
[image_age_location_encoded128] epoch 5/20 | train_f1=0.8568 | val_f1=0.5606 | val_bal_acc=0.5898
[image_age_location_encoded128] epoch 6/20 | train_f1=0.8787 | val_f1=0.5698 | val_bal_acc=0.5943
[image_age_location_encoded128] epoch 7/20 | train_f1=0.8998 | val_f1=0.5775 | val_bal_acc=0.6002
[image_age_location_encoded128] epoch 8/20 | train_f1=0.9151 | val_f1=0.5624 | val_bal_acc=0.5806
[image_age_location_encoded128] epoch 9/20 | train_f1=0.9249 | val_f1=0.6040 | val_bal_acc=0.5864
[image_age_location_encoded128] epoch 10/20 | train_f1=0.9334 | val_f1=0.5929 | val_bal_acc=0.5830
[image_age_location_encoded128] epoch 11/20 | train_f1=0.9412 | val_f1=0.5795 | val_bal_acc=0.5959
[image_age_location_encoded128] epoch 12/20 | train_f1=0.9452 | val_f1=0.5883 | val_bal_acc=0.5926
[image_age_location_encoded128] epoch 13/20 | train_f1=0.9541 | val_f1=0.5706 | val_bal_acc=0.5817
[image_age_location_encoded128] epoch 14/20 | train_f1=0.9648 | val_f1=0.5996 | val_bal_acc=0.5824
[image_age_location_encoded128] epoch 15/20 | train_f1=0.9673 | val_f1=0.6093 | val_bal_acc=0.5984
[image_age_location_encoded128] epoch 16/20 | train_f1=0.9732 | val_f1=0.5980 | val_bal_acc=0.5895
[image_age_location_encoded128] epoch 17/20 | train_f1=0.9769 | val_f1=0.5956 | val_bal_acc=0.5815
[image_age_location_encoded128] epoch 18/20 | train_f1=0.9757 | val_f1=0.5915 | val_bal_acc=0.5667
[image_age_location_encoded128] epoch 19/20 | train_f1=0.9804 | val_f1=0.6062 | val_bal_acc=0.5890
[image_age_location_encoded128] epoch 20/20 | train_f1=0.9810 | val_f1=0.6023 | val_bal_acc=0.5837

===== Running image_sex_location_encoded32 =====
[image_sex_location_encoded32] epoch 1/20 | train_f1=0.4890 | val_f1=0.5574 | val_bal_acc=0.6146
[image_sex_location_encoded32] epoch 2/20 | train_f1=0.7173 | val_f1=0.5409 | val_bal_acc=0.6004
[image_sex_location_encoded32] epoch 3/20 | train_f1=0.8031 | val_f1=0.5374 | val_bal_acc=0.5791
[image_sex_location_encoded32] epoch 4/20 | train_f1=0.8265 | val_f1=0.5526 | val_bal_acc=0.6012
[image_sex_location_encoded32] epoch 5/20 | train_f1=0.8504 | val_f1=0.5729 | val_bal_acc=0.5927
[image_sex_location_encoded32] epoch 6/20 | train_f1=0.8754 | val_f1=0.6057 | val_bal_acc=0.5885
[image_sex_location_encoded32] epoch 7/20 | train_f1=0.8973 | val_f1=0.5808 | val_bal_acc=0.6099
[image_sex_location_encoded32] epoch 8/20 | train_f1=0.9118 | val_f1=0.5928 | val_bal_acc=0.5912
[image_sex_location_encoded32] epoch 9/20 | train_f1=0.9312 | val_f1=0.5862 | val_bal_acc=0.5921
[image_sex_location_encoded32] epoch 10/20 | train_f1=0.9294 | val_f1=0.5951 | val_bal_acc=0.5940
[image_sex_location_encoded32] epoch 11/20 | train_f1=0.9438 | val_f1=0.5729 | val_bal_acc=0.6025
[image_sex_location_encoded32] epoch 12/20 | train_f1=0.9520 | val_f1=0.5895 | val_bal_acc=0.6009
[image_sex_location_encoded32] epoch 13/20 | train_f1=0.9503 | val_f1=0.5883 | val_bal_acc=0.5781
[image_sex_location_encoded32] epoch 14/20 | train_f1=0.9624 | val_f1=0.6013 | val_bal_acc=0.5705
[image_sex_location_encoded32] epoch 15/20 | train_f1=0.9688 | val_f1=0.5870 | val_bal_acc=0.5686
[image_sex_location_encoded32] epoch 16/20 | train_f1=0.9691 | val_f1=0.5898 | val_bal_acc=0.5495
[image_sex_location_encoded32] epoch 17/20 | train_f1=0.9756 | val_f1=0.5960 | val_bal_acc=0.5883
[image_sex_location_encoded32] epoch 18/20 | train_f1=0.9780 | val_f1=0.5769 | val_bal_acc=0.5847
[image_sex_location_encoded32] epoch 19/20 | train_f1=0.9752 | val_f1=0.5991 | val_bal_acc=0.5799
[image_sex_location_encoded32] epoch 20/20 | train_f1=0.9816 | val_f1=0.5980 | val_bal_acc=0.5595

===== Running image_sex_location_encoded64 =====
[image_sex_location_encoded64] epoch 1/20 | train_f1=0.4979 | val_f1=0.5109 | val_bal_acc=0.5874
[image_sex_location_encoded64] epoch 2/20 | train_f1=0.7054 | val_f1=0.5586 | val_bal_acc=0.6199
[image_sex_location_encoded64] epoch 3/20 | train_f1=0.7800 | val_f1=0.5460 | val_bal_acc=0.6050
[image_sex_location_encoded64] epoch 4/20 | train_f1=0.8249 | val_f1=0.5652 | val_bal_acc=0.6114
[image_sex_location_encoded64] epoch 5/20 | train_f1=0.8469 | val_f1=0.5552 | val_bal_acc=0.5900
[image_sex_location_encoded64] epoch 6/20 | train_f1=0.8782 | val_f1=0.5904 | val_bal_acc=0.6008
[image_sex_location_encoded64] epoch 7/20 | train_f1=0.8895 | val_f1=0.5937 | val_bal_acc=0.5821
[image_sex_location_encoded64] epoch 8/20 | train_f1=0.9091 | val_f1=0.5854 | val_bal_acc=0.5953
[image_sex_location_encoded64] epoch 9/20 | train_f1=0.9196 | val_f1=0.5754 | val_bal_acc=0.5810
[image_sex_location_encoded64] epoch 10/20 | train_f1=0.9319 | val_f1=0.5824 | val_bal_acc=0.5950
[image_sex_location_encoded64] epoch 11/20 | train_f1=0.9409 | val_f1=0.5950 | val_bal_acc=0.5842
[image_sex_location_encoded64] epoch 12/20 | train_f1=0.9514 | val_f1=0.5966 | val_bal_acc=0.5899
[image_sex_location_encoded64] epoch 13/20 | train_f1=0.9587 | val_f1=0.5830 | val_bal_acc=0.5867
[image_sex_location_encoded64] epoch 14/20 | train_f1=0.9622 | val_f1=0.6011 | val_bal_acc=0.5811
[image_sex_location_encoded64] epoch 15/20 | train_f1=0.9618 | val_f1=0.5883 | val_bal_acc=0.5646
[image_sex_location_encoded64] epoch 16/20 | train_f1=0.9735 | val_f1=0.5813 | val_bal_acc=0.5815
[image_sex_location_encoded64] epoch 17/20 | train_f1=0.9713 | val_f1=0.6125 | val_bal_acc=0.5907
[image_sex_location_encoded64] epoch 18/20 | train_f1=0.9814 | val_f1=0.5778 | val_bal_acc=0.5980
[image_sex_location_encoded64] epoch 19/20 | train_f1=0.9779 | val_f1=0.5768 | val_bal_acc=0.5793
[image_sex_location_encoded64] epoch 20/20 | train_f1=0.9859 | val_f1=0.6074 | val_bal_acc=0.5725

===== Running image_sex_location_encoded128 =====
[image_sex_location_encoded128] epoch 1/20 | train_f1=0.4777 | val_f1=0.4964 | val_bal_acc=0.6232
[image_sex_location_encoded128] epoch 2/20 | train_f1=0.6931 | val_f1=0.5550 | val_bal_acc=0.6458
[image_sex_location_encoded128] epoch 3/20 | train_f1=0.7765 | val_f1=0.5703 | val_bal_acc=0.6463
[image_sex_location_encoded128] epoch 4/20 | train_f1=0.8177 | val_f1=0.5527 | val_bal_acc=0.6006
[image_sex_location_encoded128] epoch 5/20 | train_f1=0.8547 | val_f1=0.5748 | val_bal_acc=0.5960
[image_sex_location_encoded128] epoch 6/20 | train_f1=0.8717 | val_f1=0.5563 | val_bal_acc=0.5875
[image_sex_location_encoded128] epoch 7/20 | train_f1=0.8887 | val_f1=0.5859 | val_bal_acc=0.5815
[image_sex_location_encoded128] epoch 8/20 | train_f1=0.9063 | val_f1=0.5678 | val_bal_acc=0.5716
[image_sex_location_encoded128] epoch 9/20 | train_f1=0.9207 | val_f1=0.5816 | val_bal_acc=0.5830
[image_sex_location_encoded128] epoch 10/20 | train_f1=0.9300 | val_f1=0.6055 | val_bal_acc=0.5963
[image_sex_location_encoded128] epoch 11/20 | train_f1=0.9372 | val_f1=0.5656 | val_bal_acc=0.5860
[image_sex_location_encoded128] epoch 12/20 | train_f1=0.9411 | val_f1=0.5852 | val_bal_acc=0.5881
[image_sex_location_encoded128] epoch 13/20 | train_f1=0.9526 | val_f1=0.5943 | val_bal_acc=0.5652
[image_sex_location_encoded128] epoch 14/20 | train_f1=0.9635 | val_f1=0.6032 | val_bal_acc=0.5789
[image_sex_location_encoded128] epoch 15/20 | train_f1=0.9651 | val_f1=0.6066 | val_bal_acc=0.5754
[image_sex_location_encoded128] epoch 16/20 | train_f1=0.9675 | val_f1=0.5946 | val_bal_acc=0.5661
[image_sex_location_encoded128] epoch 17/20 | train_f1=0.9729 | val_f1=0.5890 | val_bal_acc=0.5803
[image_sex_location_encoded128] epoch 18/20 | train_f1=0.9757 | val_f1=0.5989 | val_bal_acc=0.5765
[image_sex_location_encoded128] epoch 19/20 | train_f1=0.9755 | val_f1=0.5978 | val_bal_acc=0.5770
[image_sex_location_encoded128] epoch 20/20 | train_f1=0.9750 | val_f1=0.5865 | val_bal_acc=0.5556
Out[25]:
18
In [26]:
# Cell C: 整理成总表,并挑每个组合的最佳维度
import pandas as pd

all_search_df = pd.DataFrame(all_search_results)

display(
    all_search_df[
        [
            "method",
            "base_method",
            "features",
            "metadata_dim",
            "metadata_embed_dim",
            "best_val_macro_f1",
            "test_accuracy",
            "test_balanced_accuracy",
            "test_macro_f1",
        ]
    ].sort_values(["base_method", "metadata_embed_dim"])
)

best_per_combination_df = (
    all_search_df
    .sort_values(["base_method", "best_val_macro_f1"], ascending=[True, False])
    .groupby("base_method", as_index=False)
    .first()
)

display(
    best_per_combination_df[
        [
            "method",
            "base_method",
            "features",
            "metadata_dim",
            "metadata_embed_dim",
            "best_val_macro_f1",
            "test_accuracy",
            "test_balanced_accuracy",
            "test_macro_f1",
        ]
    ].sort_values("base_method")
)
method base_method features metadata_dim metadata_embed_dim best_val_macro_f1 test_accuracy test_balanced_accuracy test_macro_f1
0 image_age_encoded32 image_age age 1 32 0.592534 0.770425 0.578116 0.576380
1 image_age_encoded64 image_age age 1 64 0.593023 0.781904 0.601570 0.584162
2 image_age_encoded128 image_age age 1 128 0.594447 0.766374 0.549634 0.557388
12 image_age_location_encoded32 image_age_location age,location 16 32 0.604529 0.799460 0.578341 0.588440
13 image_age_location_encoded64 image_age_location age,location 16 64 0.599031 0.788656 0.568477 0.575066
14 image_age_location_encoded128 image_age_location age,location 16 128 0.609298 0.783255 0.602824 0.607910
9 image_age_sex_encoded32 image_age_sex age,sex 4 32 0.603107 0.777178 0.587462 0.590630
10 image_age_sex_encoded64 image_age_sex age,sex 4 64 0.604039 0.793383 0.582281 0.581666
11 image_age_sex_encoded128 image_age_sex age,sex 4 128 0.602657 0.774477 0.572578 0.576931
6 image_location_encoded32 image_location location 15 32 0.603778 0.754220 0.558853 0.564874
7 image_location_encoded64 image_location location 15 64 0.600608 0.781904 0.574504 0.580142
8 image_location_encoded128 image_location location 15 128 0.595590 0.784605 0.543721 0.561448
3 image_sex_encoded32 image_sex sex 3 32 0.596308 0.759622 0.575438 0.576118
4 image_sex_encoded64 image_sex sex 3 64 0.600935 0.786631 0.580348 0.579055
5 image_sex_encoded128 image_sex sex 3 128 0.599878 0.766374 0.568703 0.569694
15 image_sex_location_encoded32 image_sex_location sex,location 18 32 0.605750 0.791357 0.611511 0.615742
16 image_sex_location_encoded64 image_sex_location sex,location 18 64 0.612467 0.792708 0.598458 0.596497
17 image_sex_location_encoded128 image_sex_location sex,location 18 128 0.606623 0.796759 0.585301 0.596229
method base_method features metadata_dim metadata_embed_dim best_val_macro_f1 test_accuracy test_balanced_accuracy test_macro_f1
0 image_age_encoded128 image_age age 1 128 0.594447 0.766374 0.549634 0.557388
1 image_age_location_encoded128 image_age_location age,location 16 128 0.609298 0.783255 0.602824 0.607910
2 image_age_sex_encoded64 image_age_sex age,sex 4 64 0.604039 0.793383 0.582281 0.581666
3 image_location_encoded32 image_location location 15 32 0.603778 0.754220 0.558853 0.564874
4 image_sex_encoded64 image_sex sex 3 64 0.600935 0.786631 0.580348 0.579055
5 image_sex_location_encoded64 image_sex_location sex,location 18 64 0.612467 0.792708 0.598458 0.596497
In [27]:
# Cell D: 保存结果
from datetime import datetime
import json
from pathlib import Path

save_dir = Path("/Users/applesues01/Documents/Medical_Agent/supports")
save_dir.mkdir(parents=True, exist_ok=True)

date_tag = datetime.now().strftime("%Y-%m-%d")
time_tag = datetime.now().strftime("%H%M%S")

full_csv = save_dir / f"{date_tag}_{time_tag}_metadata_embed_search_full.csv"
best_csv = save_dir / f"{date_tag}_{time_tag}_metadata_embed_search_best.csv"
summary_json = save_dir / f"{date_tag}_{time_tag}_metadata_embed_search_summary.json"

all_search_df.to_csv(full_csv, index=False)
best_per_combination_df.to_csv(best_csv, index=False)

payload = {
    "generated_at": datetime.now().isoformat(),
    "purpose": "best metadata embedding dimension search for 6 metadata combinations",
    "embed_dims": EMBED_DIMS,
    "num_epochs": 5,
    "full_results": all_search_results,
    "best_results": best_per_combination_df.to_dict(orient="records"),
}

def to_jsonable(x):
    import numpy as np
    if isinstance(x, dict):
        return {k: to_jsonable(v) for k, v in x.items()}
    if isinstance(x, list):
        return [to_jsonable(v) for v in x]
    if isinstance(x, tuple):
        return [to_jsonable(v) for v in x]
    if isinstance(x, np.integer):
        return int(x)
    if isinstance(x, np.floating):
        return float(x)
    if isinstance(x, np.ndarray):
        return x.tolist()
    return x

with open(summary_json, "w", encoding="utf-8") as f:
    json.dump(to_jsonable(payload), f, ensure_ascii=False, indent=2)

print("Saved:")
print(full_csv)
print(best_csv)
print(summary_json)
Saved:
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-03_165248_metadata_embed_search_full.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-03_165248_metadata_embed_search_best.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-03_165248_metadata_embed_search_summary.json