10 Save All Metadata Models¶

Train and save the best metadata-combination models so the case-level agent can call the corresponding model after each question.

In [1]:
import os
import copy
import json
import numpy as np
import pandas as pd
from pathlib import Path
from datetime import datetime

import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
from torchvision.models import resnet50
from sklearn.metrics import accuracy_score, balanced_accuracy_score, f1_score
from sklearn.utils.class_weight import compute_class_weight
from PIL import Image
In [2]:
PROJECT_DIR = Path('/Users/applesues01/Documents/Medical_Agent')
DATA_DIR = PROJECT_DIR / 'data' / 'HAM10000'
SPLIT_DIR = DATA_DIR / 'splits'
CHECKPOINT_DIR = PROJECT_DIR / 'checkpoints'
SUPPORT_DIR = PROJECT_DIR / 'supports'

IMAGE_DIR1 = DATA_DIR / 'HAM10000_images_part_1'
IMAGE_DIR2 = DATA_DIR / 'HAM10000_images_part_2'

CHECKPOINT_DIR.mkdir(parents=True, exist_ok=True)
SUPPORT_DIR.mkdir(parents=True, exist_ok=True)

train_df = pd.read_csv(SPLIT_DIR / 'train.csv')
val_df = pd.read_csv(SPLIT_DIR / 'val.csv')
test_df = pd.read_csv(SPLIT_DIR / 'test.csv')

len(train_df), len(val_df), len(test_df)
Out[2]:
(7002, 1532, 1481)
In [3]:
device = torch.device('mps' if torch.backends.mps.is_available() else 'cpu')
device
Out[3]:
device(type='mps')
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],
}

IMAGE_SIZE = 224
BATCH_SIZE = 16
NUM_EPOCHS = 20

TRAINED_BACKBONE_PATH = CHECKPOINT_DIR / 'resnet50_image_only_finetuned_best.pth'
In [5]:
BEST_EXPERIMENTS = [
    {'name': 'image_age', 'features': ['age'], 'metadata_embed_dim': 128},
    {'name': 'image_sex', 'features': ['sex'], 'metadata_embed_dim': 64},
    {'name': 'image_location', 'features': ['location'], 'metadata_embed_dim': 32},
    {'name': 'image_age_sex', 'features': ['age', 'sex'], 'metadata_embed_dim': 64},
    {'name': 'image_age_location', 'features': ['age', 'location'], 'metadata_embed_dim': 128},
    {'name': 'image_sex_location', 'features': ['sex', 'location'], 'metadata_embed_dim': 64},
    {'name': 'image_all_metadata', 'features': ['age', 'sex', 'location'], 'metadata_embed_dim': 128},
]

pd.DataFrame(BEST_EXPERIMENTS)
Out[5]:
name features metadata_embed_dim
0 image_age [age] 128
1 image_sex [sex] 64
2 image_location [location] 32
3 image_age_sex [age, sex] 64
4 image_age_location [age, location] 128
5 image_sex_location [sex, location] 64
6 image_all_metadata [age, sex, location] 128
In [6]:
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 [7]:
def resolve_image_path(image_id: str):
    filename = f'{image_id}.jpg'
    path1 = IMAGE_DIR1 / filename
    path2 = IMAGE_DIR2 / filename
    if path1.exists():
        return path1
    return path2


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 [8]:
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 [9]:
image_only_model = resnet50(weights=None)
image_only_model.fc = nn.Linear(image_only_model.fc.in_features, 7)
image_only_model.load_state_dict(torch.load(TRAINED_BACKBONE_PATH, map_location=device))
image_only_model = image_only_model.to(device)
image_only_model.eval()
print(TRAINED_BACKBONE_PATH)
/Users/applesues01/Documents/Medical_Agent/checkpoints/resnet50_image_only_finetuned_best.pth
In [10]:
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)
        return torch.flatten(x, 1)


feature_extractor = ResNet50FeatureExtractor(image_only_model).to(device)
feature_extractor.eval()
for param in feature_extractor.parameters():
    param.requires_grad = False
In [11]:
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)


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)

train_image_features.shape, val_image_features.shape, test_image_features.shape
Out[11]:
(torch.Size([7002, 2048]), torch.Size([1532, 2048]), torch.Size([1481, 2048]))
In [12]:
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)


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 [13]:
train_age_mean = train_df['age'].mean()

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 [14]:
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 [15]:
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),
        'preds': preds,
        'truths': truths,
    }


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 [16]:
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)
class_weights
Out[16]:
tensor([ 4.3491,  2.7330,  1.2924, 13.1617,  1.2857,  0.2138, 10.1039],
       device='mps:0')
In [17]:
def train_and_save_experiment(experiment_name, selected_features, metadata_embed_dim, num_epochs=NUM_EPOCHS):
    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}] 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)

    checkpoint_path = CHECKPOINT_DIR / f'{experiment_name}_best.pth'
    torch.save(best_state, checkpoint_path)

    history_path = SUPPORT_DIR / f'{experiment_name}_history.csv'
    pd.DataFrame(history).to_csv(history_path, index=False)

    return {
        'Method': experiment_name,
        'Features': ','.join(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'],
        'Checkpoint Path': str(checkpoint_path),
        'History Path': str(history_path),
    }
In [18]:
all_saved_results = []

for exp in BEST_EXPERIMENTS:
    result = train_and_save_experiment(
        experiment_name=exp['name'],
        selected_features=exp['features'],
        metadata_embed_dim=exp['metadata_embed_dim'],
        num_epochs=NUM_EPOCHS,
    )
    all_saved_results.append(result)

saved_models_df = pd.DataFrame(all_saved_results)
saved_models_df
[image_age] epoch 1/20 | train_f1=0.4904 | val_f1=0.4790 | val_bal_acc=0.5880
[image_age] epoch 2/20 | train_f1=0.6993 | val_f1=0.5486 | val_bal_acc=0.6002
[image_age] epoch 3/20 | train_f1=0.7844 | val_f1=0.5397 | val_bal_acc=0.6181
[image_age] epoch 4/20 | train_f1=0.8212 | val_f1=0.5389 | val_bal_acc=0.5800
[image_age] epoch 5/20 | train_f1=0.8466 | val_f1=0.5356 | val_bal_acc=0.5732
[image_age] epoch 6/20 | train_f1=0.8813 | val_f1=0.5759 | val_bal_acc=0.5940
[image_age] epoch 7/20 | train_f1=0.8933 | val_f1=0.5707 | val_bal_acc=0.5907
[image_age] epoch 8/20 | train_f1=0.9005 | val_f1=0.5885 | val_bal_acc=0.5951
[image_age] epoch 9/20 | train_f1=0.9132 | val_f1=0.5926 | val_bal_acc=0.5978
[image_age] epoch 10/20 | train_f1=0.9267 | val_f1=0.5592 | val_bal_acc=0.5656
[image_age] epoch 11/20 | train_f1=0.9360 | val_f1=0.5426 | val_bal_acc=0.5514
[image_age] epoch 12/20 | train_f1=0.9382 | val_f1=0.5885 | val_bal_acc=0.5821
[image_age] epoch 13/20 | train_f1=0.9542 | val_f1=0.5668 | val_bal_acc=0.5652
[image_age] epoch 14/20 | train_f1=0.9499 | val_f1=0.5789 | val_bal_acc=0.5539
[image_age] epoch 15/20 | train_f1=0.9626 | val_f1=0.5921 | val_bal_acc=0.5814
[image_age] epoch 16/20 | train_f1=0.9716 | val_f1=0.5740 | val_bal_acc=0.5637
[image_age] epoch 17/20 | train_f1=0.9702 | val_f1=0.5750 | val_bal_acc=0.5627
[image_age] epoch 18/20 | train_f1=0.9778 | val_f1=0.5654 | val_bal_acc=0.5527
[image_age] epoch 19/20 | train_f1=0.9752 | val_f1=0.5727 | val_bal_acc=0.5651
[image_age] epoch 20/20 | train_f1=0.9788 | val_f1=0.5691 | val_bal_acc=0.5550
[image_sex] epoch 1/20 | train_f1=0.4455 | val_f1=0.4822 | val_bal_acc=0.5255
[image_sex] epoch 2/20 | train_f1=0.6983 | val_f1=0.5077 | val_bal_acc=0.6344
[image_sex] epoch 3/20 | train_f1=0.7740 | val_f1=0.5544 | val_bal_acc=0.5940
[image_sex] epoch 4/20 | train_f1=0.8239 | val_f1=0.5616 | val_bal_acc=0.5945
[image_sex] epoch 5/20 | train_f1=0.8499 | val_f1=0.5641 | val_bal_acc=0.5782
[image_sex] epoch 6/20 | train_f1=0.8726 | val_f1=0.5531 | val_bal_acc=0.5884
[image_sex] epoch 7/20 | train_f1=0.8901 | val_f1=0.5622 | val_bal_acc=0.6003
[image_sex] epoch 8/20 | train_f1=0.9052 | val_f1=0.5871 | val_bal_acc=0.5824
[image_sex] epoch 9/20 | train_f1=0.9171 | val_f1=0.5622 | val_bal_acc=0.5766
[image_sex] epoch 10/20 | train_f1=0.9293 | val_f1=0.5706 | val_bal_acc=0.5672
[image_sex] epoch 11/20 | train_f1=0.9383 | val_f1=0.5634 | val_bal_acc=0.5502
[image_sex] epoch 12/20 | train_f1=0.9495 | val_f1=0.5464 | val_bal_acc=0.5303
[image_sex] epoch 13/20 | train_f1=0.9489 | val_f1=0.5796 | val_bal_acc=0.5715
[image_sex] epoch 14/20 | train_f1=0.9580 | val_f1=0.5952 | val_bal_acc=0.5819
[image_sex] epoch 15/20 | train_f1=0.9628 | val_f1=0.5959 | val_bal_acc=0.5920
[image_sex] epoch 16/20 | train_f1=0.9684 | val_f1=0.5930 | val_bal_acc=0.5837
[image_sex] epoch 17/20 | train_f1=0.9738 | val_f1=0.5892 | val_bal_acc=0.5646
[image_sex] epoch 18/20 | train_f1=0.9761 | val_f1=0.5744 | val_bal_acc=0.5665
[image_sex] epoch 19/20 | train_f1=0.9814 | val_f1=0.5867 | val_bal_acc=0.5647
[image_sex] epoch 20/20 | train_f1=0.9809 | val_f1=0.5767 | val_bal_acc=0.5509
[image_location] epoch 1/20 | train_f1=0.4741 | val_f1=0.4715 | val_bal_acc=0.5593
[image_location] epoch 2/20 | train_f1=0.7140 | val_f1=0.5347 | val_bal_acc=0.5951
[image_location] epoch 3/20 | train_f1=0.7816 | val_f1=0.5564 | val_bal_acc=0.5958
[image_location] epoch 4/20 | train_f1=0.8328 | val_f1=0.5516 | val_bal_acc=0.5888
[image_location] epoch 5/20 | train_f1=0.8598 | val_f1=0.5516 | val_bal_acc=0.5978
[image_location] epoch 6/20 | train_f1=0.8706 | val_f1=0.5633 | val_bal_acc=0.5820
[image_location] epoch 7/20 | train_f1=0.8904 | val_f1=0.5897 | val_bal_acc=0.5999
[image_location] epoch 8/20 | train_f1=0.9060 | val_f1=0.5960 | val_bal_acc=0.6041
[image_location] epoch 9/20 | train_f1=0.9255 | val_f1=0.5739 | val_bal_acc=0.5642
[image_location] epoch 10/20 | train_f1=0.9345 | val_f1=0.6000 | val_bal_acc=0.5816
[image_location] epoch 11/20 | train_f1=0.9424 | val_f1=0.5888 | val_bal_acc=0.5832
[image_location] epoch 12/20 | train_f1=0.9483 | val_f1=0.5907 | val_bal_acc=0.5740
[image_location] epoch 13/20 | train_f1=0.9560 | val_f1=0.6003 | val_bal_acc=0.5849
[image_location] epoch 14/20 | train_f1=0.9634 | val_f1=0.5891 | val_bal_acc=0.5836
[image_location] epoch 15/20 | train_f1=0.9607 | val_f1=0.5885 | val_bal_acc=0.5962
[image_location] epoch 16/20 | train_f1=0.9700 | val_f1=0.6042 | val_bal_acc=0.5733
[image_location] epoch 17/20 | train_f1=0.9729 | val_f1=0.5859 | val_bal_acc=0.5751
[image_location] epoch 18/20 | train_f1=0.9749 | val_f1=0.5970 | val_bal_acc=0.5997
[image_location] epoch 19/20 | train_f1=0.9791 | val_f1=0.5749 | val_bal_acc=0.5442
[image_location] epoch 20/20 | train_f1=0.9825 | val_f1=0.5897 | val_bal_acc=0.5693
[image_age_sex] epoch 1/20 | train_f1=0.4324 | val_f1=0.5045 | val_bal_acc=0.5779
[image_age_sex] epoch 2/20 | train_f1=0.7107 | val_f1=0.5466 | val_bal_acc=0.6200
[image_age_sex] epoch 3/20 | train_f1=0.7831 | val_f1=0.5368 | val_bal_acc=0.5904
[image_age_sex] epoch 4/20 | train_f1=0.8349 | val_f1=0.5615 | val_bal_acc=0.5875
[image_age_sex] epoch 5/20 | train_f1=0.8582 | val_f1=0.5798 | val_bal_acc=0.5856
[image_age_sex] epoch 6/20 | train_f1=0.8698 | val_f1=0.5819 | val_bal_acc=0.6061
[image_age_sex] epoch 7/20 | train_f1=0.8934 | val_f1=0.5845 | val_bal_acc=0.6144
[image_age_sex] epoch 8/20 | train_f1=0.9154 | val_f1=0.5781 | val_bal_acc=0.5945
[image_age_sex] epoch 9/20 | train_f1=0.9242 | val_f1=0.5747 | val_bal_acc=0.5886
[image_age_sex] epoch 10/20 | train_f1=0.9349 | val_f1=0.5965 | val_bal_acc=0.5911
[image_age_sex] epoch 11/20 | train_f1=0.9483 | val_f1=0.5880 | val_bal_acc=0.5920
[image_age_sex] epoch 12/20 | train_f1=0.9503 | val_f1=0.5921 | val_bal_acc=0.5878
[image_age_sex] epoch 13/20 | train_f1=0.9584 | val_f1=0.5692 | val_bal_acc=0.5865
[image_age_sex] epoch 14/20 | train_f1=0.9644 | val_f1=0.5627 | val_bal_acc=0.5553
[image_age_sex] epoch 15/20 | train_f1=0.9695 | val_f1=0.5908 | val_bal_acc=0.5809
[image_age_sex] epoch 16/20 | train_f1=0.9705 | val_f1=0.5879 | val_bal_acc=0.5717
[image_age_sex] epoch 17/20 | train_f1=0.9781 | val_f1=0.5808 | val_bal_acc=0.6036
[image_age_sex] epoch 18/20 | train_f1=0.9734 | val_f1=0.5845 | val_bal_acc=0.5731
[image_age_sex] epoch 19/20 | train_f1=0.9797 | val_f1=0.5467 | val_bal_acc=0.5458
[image_age_sex] epoch 20/20 | train_f1=0.9826 | val_f1=0.5685 | val_bal_acc=0.5554
[image_age_location] epoch 1/20 | train_f1=0.5016 | val_f1=0.5030 | val_bal_acc=0.6419
[image_age_location] epoch 2/20 | train_f1=0.7073 | val_f1=0.5661 | val_bal_acc=0.6328
[image_age_location] epoch 3/20 | train_f1=0.7842 | val_f1=0.6027 | val_bal_acc=0.6237
[image_age_location] epoch 4/20 | train_f1=0.8161 | val_f1=0.5629 | val_bal_acc=0.6149
[image_age_location] epoch 5/20 | train_f1=0.8610 | val_f1=0.5966 | val_bal_acc=0.5896
[image_age_location] epoch 6/20 | train_f1=0.8784 | val_f1=0.5906 | val_bal_acc=0.6031
[image_age_location] epoch 7/20 | train_f1=0.8958 | val_f1=0.5783 | val_bal_acc=0.5913
[image_age_location] epoch 8/20 | train_f1=0.9159 | val_f1=0.5779 | val_bal_acc=0.5942
[image_age_location] epoch 9/20 | train_f1=0.9205 | val_f1=0.5765 | val_bal_acc=0.5764
[image_age_location] epoch 10/20 | train_f1=0.9326 | val_f1=0.5989 | val_bal_acc=0.5870
[image_age_location] epoch 11/20 | train_f1=0.9445 | val_f1=0.5811 | val_bal_acc=0.5892
[image_age_location] epoch 12/20 | train_f1=0.9540 | val_f1=0.5734 | val_bal_acc=0.5785
[image_age_location] epoch 13/20 | train_f1=0.9539 | val_f1=0.5841 | val_bal_acc=0.5509
[image_age_location] epoch 14/20 | train_f1=0.9604 | val_f1=0.5949 | val_bal_acc=0.5848
[image_age_location] epoch 15/20 | train_f1=0.9662 | val_f1=0.6004 | val_bal_acc=0.5989
[image_age_location] epoch 16/20 | train_f1=0.9724 | val_f1=0.6069 | val_bal_acc=0.5959
[image_age_location] epoch 17/20 | train_f1=0.9766 | val_f1=0.6019 | val_bal_acc=0.5881
[image_age_location] epoch 18/20 | train_f1=0.9771 | val_f1=0.6095 | val_bal_acc=0.5921
[image_age_location] epoch 19/20 | train_f1=0.9764 | val_f1=0.5962 | val_bal_acc=0.5759
[image_age_location] epoch 20/20 | train_f1=0.9810 | val_f1=0.5935 | val_bal_acc=0.5911
[image_sex_location] epoch 1/20 | train_f1=0.4832 | val_f1=0.5004 | val_bal_acc=0.5867
[image_sex_location] epoch 2/20 | train_f1=0.7093 | val_f1=0.5697 | val_bal_acc=0.6061
[image_sex_location] epoch 3/20 | train_f1=0.7760 | val_f1=0.5738 | val_bal_acc=0.6010
[image_sex_location] epoch 4/20 | train_f1=0.8255 | val_f1=0.5686 | val_bal_acc=0.5947
[image_sex_location] epoch 5/20 | train_f1=0.8480 | val_f1=0.5791 | val_bal_acc=0.6153
[image_sex_location] epoch 6/20 | train_f1=0.8750 | val_f1=0.5985 | val_bal_acc=0.5856
[image_sex_location] epoch 7/20 | train_f1=0.8999 | val_f1=0.6060 | val_bal_acc=0.5919
[image_sex_location] epoch 8/20 | train_f1=0.9075 | val_f1=0.5700 | val_bal_acc=0.5963
[image_sex_location] epoch 9/20 | train_f1=0.9054 | val_f1=0.5781 | val_bal_acc=0.5873
[image_sex_location] epoch 10/20 | train_f1=0.9266 | val_f1=0.5724 | val_bal_acc=0.5890
[image_sex_location] epoch 11/20 | train_f1=0.9409 | val_f1=0.5981 | val_bal_acc=0.5828
[image_sex_location] epoch 12/20 | train_f1=0.9464 | val_f1=0.5806 | val_bal_acc=0.5883
[image_sex_location] epoch 13/20 | train_f1=0.9514 | val_f1=0.6035 | val_bal_acc=0.5754
[image_sex_location] epoch 14/20 | train_f1=0.9595 | val_f1=0.5972 | val_bal_acc=0.5848
[image_sex_location] epoch 15/20 | train_f1=0.9633 | val_f1=0.5798 | val_bal_acc=0.5857
[image_sex_location] epoch 16/20 | train_f1=0.9663 | val_f1=0.5929 | val_bal_acc=0.5880
[image_sex_location] epoch 17/20 | train_f1=0.9737 | val_f1=0.5994 | val_bal_acc=0.5708
[image_sex_location] epoch 18/20 | train_f1=0.9798 | val_f1=0.6014 | val_bal_acc=0.5904
[image_sex_location] epoch 19/20 | train_f1=0.9822 | val_f1=0.5865 | val_bal_acc=0.5809
[image_sex_location] epoch 20/20 | train_f1=0.9844 | val_f1=0.6020 | val_bal_acc=0.5758
[image_all_metadata] epoch 1/20 | train_f1=0.4725 | val_f1=0.5428 | val_bal_acc=0.5697
[image_all_metadata] epoch 2/20 | train_f1=0.6982 | val_f1=0.5902 | val_bal_acc=0.6312
[image_all_metadata] epoch 3/20 | train_f1=0.7815 | val_f1=0.5555 | val_bal_acc=0.6135
[image_all_metadata] epoch 4/20 | train_f1=0.8206 | val_f1=0.5671 | val_bal_acc=0.5891
[image_all_metadata] epoch 5/20 | train_f1=0.8521 | val_f1=0.5907 | val_bal_acc=0.5944
[image_all_metadata] epoch 6/20 | train_f1=0.8800 | val_f1=0.5993 | val_bal_acc=0.6155
[image_all_metadata] epoch 7/20 | train_f1=0.8960 | val_f1=0.5607 | val_bal_acc=0.5914
[image_all_metadata] epoch 8/20 | train_f1=0.9159 | val_f1=0.5823 | val_bal_acc=0.5896
[image_all_metadata] epoch 9/20 | train_f1=0.9260 | val_f1=0.6065 | val_bal_acc=0.5895
[image_all_metadata] epoch 10/20 | train_f1=0.9345 | val_f1=0.6075 | val_bal_acc=0.5908
[image_all_metadata] epoch 11/20 | train_f1=0.9385 | val_f1=0.6010 | val_bal_acc=0.6021
[image_all_metadata] epoch 12/20 | train_f1=0.9500 | val_f1=0.5806 | val_bal_acc=0.5955
[image_all_metadata] epoch 13/20 | train_f1=0.9557 | val_f1=0.6037 | val_bal_acc=0.5742
[image_all_metadata] epoch 14/20 | train_f1=0.9579 | val_f1=0.5913 | val_bal_acc=0.5759
[image_all_metadata] epoch 15/20 | train_f1=0.9662 | val_f1=0.5916 | val_bal_acc=0.5711
[image_all_metadata] epoch 16/20 | train_f1=0.9745 | val_f1=0.6002 | val_bal_acc=0.5832
[image_all_metadata] epoch 17/20 | train_f1=0.9720 | val_f1=0.5920 | val_bal_acc=0.5962
[image_all_metadata] epoch 18/20 | train_f1=0.9710 | val_f1=0.6053 | val_bal_acc=0.5876
[image_all_metadata] epoch 19/20 | train_f1=0.9770 | val_f1=0.6112 | val_bal_acc=0.5896
[image_all_metadata] epoch 20/20 | train_f1=0.9767 | val_f1=0.6051 | val_bal_acc=0.5835
Out[18]:
Method Features Metadata Dim Metadata Embed Dim Best Val Macro-F1 Test Accuracy Test Balanced Accuracy Test Macro-F1 Checkpoint Path History Path
0 image_age age 1 128 0.592628 0.778528 0.601930 0.592185 /Users/applesues01/Documents/Medical_Agent/che... /Users/applesues01/Documents/Medical_Agent/sup...
1 image_sex sex 3 64 0.595947 0.762323 0.565402 0.565863 /Users/applesues01/Documents/Medical_Agent/che... /Users/applesues01/Documents/Medical_Agent/sup...
2 image_location location 15 32 0.604193 0.790007 0.552086 0.566736 /Users/applesues01/Documents/Medical_Agent/che... /Users/applesues01/Documents/Medical_Agent/sup...
3 image_age_sex age,sex 4 64 0.596464 0.788656 0.585681 0.583705 /Users/applesues01/Documents/Medical_Agent/che... /Users/applesues01/Documents/Medical_Agent/sup...
4 image_age_location age,location 16 128 0.609501 0.777178 0.573000 0.580037 /Users/applesues01/Documents/Medical_Agent/che... /Users/applesues01/Documents/Medical_Agent/sup...
5 image_sex_location sex,location 18 64 0.605986 0.800135 0.609610 0.611445 /Users/applesues01/Documents/Medical_Agent/che... /Users/applesues01/Documents/Medical_Agent/sup...
6 image_all_metadata age,sex,location 19 128 0.611203 0.803511 0.590618 0.595896 /Users/applesues01/Documents/Medical_Agent/che... /Users/applesues01/Documents/Medical_Agent/sup...
In [19]:
date_tag = datetime.now().strftime('%Y-%m-%d')
time_tag = datetime.now().strftime('%H%M%S')

summary_csv_path = SUPPORT_DIR / f'{date_tag}_{time_tag}_all_saved_metadata_models.csv'
summary_json_path = SUPPORT_DIR / f'{date_tag}_{time_tag}_all_saved_metadata_models.json'

saved_models_df.to_csv(summary_csv_path, index=False)

with open(summary_json_path, 'w', encoding='utf-8') as f:
    json.dump(all_saved_results, f, ensure_ascii=False, indent=2)

print(summary_csv_path)
print(summary_json_path)
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-03_194704_all_saved_metadata_models.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-03_194704_all_saved_metadata_models.json

我觉得这样子还是不太严谨,所以我决定全部重做,把权重的测算再算一遍,就是raw,32,64,128,256,因为虽然说之前全加进去128最好,但是单独也许就说不定了呢?连续5轮不上升就结束当前训练换下一个¶

In [20]:
SEARCH_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"]},
]

SEARCH_SETTINGS = [
    {"mode": "raw", "embed_dim": None},
    {"mode": "encoded", "embed_dim": 32},
    {"mode": "encoded", "embed_dim": 64},
    {"mode": "encoded", "embed_dim": 128},
    {"mode": "encoded", "embed_dim": 256},
]
In [21]:
class RawFusionClassifier(nn.Module):
    def __init__(self, metadata_input_dim, num_classes=7):
        super().__init__()
        self.classifier = nn.Sequential(
            nn.Linear(2048 + metadata_input_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 [22]:
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 [23]:
def train_one_model_with_early_stopping(
    model,
    train_loader,
    val_loader,
    criterion,
    optimizer,
    device,
    max_epochs=30,
    patience=5,
    min_delta=1e-4
):
    best_val_f1 = -1.0
    best_state = None
    wait = 0
    history = []

    for epoch in range(max_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"epoch {epoch+1}/{max_epochs} | "
            f"train_f1={train_metrics['macro_f1']:.4f} | "
            f"val_f1={val_metrics['macro_f1']:.4f} | "
            f"wait={wait}"
        )

        if val_metrics["macro_f1"] > best_val_f1 + min_delta:
            best_val_f1 = val_metrics["macro_f1"]
            best_state = copy.deepcopy(model.state_dict())
            wait = 0
        else:
            wait += 1

        if wait >= patience:
            print(f"Early stop at epoch {epoch+1}")
            break

    model.load_state_dict(best_state)
    return best_val_f1, best_state, history
In [24]:
def run_single_trial(
    experiment_name,
    selected_features,
    mode,
    embed_dim,
    max_epochs=30,
    patience=5
):
    train_loader, val_loader, test_loader, metadata_dim = build_cached_loaders_for_experiment(selected_features)

    if mode == "raw":
        model = RawFusionClassifier(metadata_input_dim=metadata_dim, num_classes=7).to(device)
    else:
        model = MetadataEncoderFusionClassifier(
            metadata_input_dim=metadata_dim,
            metadata_embed_dim=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, best_state, history = train_one_model_with_early_stopping(
        model,
        train_loader,
        val_loader,
        criterion,
        optimizer,
        device,
        max_epochs=max_epochs,
        patience=patience
    )

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

    return {
        "method": experiment_name,
        "features": ",".join(selected_features),
        "mode": mode,
        "metadata_dim": metadata_dim,
        "metadata_embed_dim": 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"],
        "best_state": best_state,
        "history": history,
    }
In [25]:
all_best_results = []
all_trial_results = []

for exp in SEARCH_EXPERIMENTS:
    print(f"\n========== {exp['name']} ==========")

    best_trial = None

    for setting in SEARCH_SETTINGS:
        mode = setting["mode"]
        embed_dim = setting["embed_dim"]

        tag = "raw" if mode == "raw" else f"encoded{embed_dim}"
        print(f"\n--- Trial: {exp['name']} | {tag} ---")

        result = run_single_trial(
            experiment_name=exp["name"],
            selected_features=exp["features"],
            mode=mode,
            embed_dim=embed_dim,
            max_epochs=30,
            patience=5
        )

        all_trial_results.append(result)

        if (best_trial is None) or (result["best_val_macro_f1"] > best_trial["best_val_macro_f1"]):
            best_trial = result

    checkpoint_name = f"{exp['name']}_best.pth"
    checkpoint_path = CHECKPOINT_DIR / checkpoint_name
    torch.save(best_trial["best_state"], checkpoint_path)

    history_path = SUPPORT_DIR / f"{exp['name']}_best_history.csv"
    pd.DataFrame(best_trial["history"]).to_csv(history_path, index=False)

    best_trial["checkpoint_path"] = str(checkpoint_path)
    best_trial["history_path"] = str(history_path)

    all_best_results.append(best_trial)
========== image_age ==========

--- Trial: image_age | raw ---
epoch 1/30 | train_f1=0.4626 | val_f1=0.5187 | wait=0
epoch 2/30 | train_f1=0.7053 | val_f1=0.5085 | wait=0
epoch 3/30 | train_f1=0.7749 | val_f1=0.5781 | wait=1
epoch 4/30 | train_f1=0.8205 | val_f1=0.5383 | wait=0
epoch 5/30 | train_f1=0.8589 | val_f1=0.5476 | wait=1
epoch 6/30 | train_f1=0.8812 | val_f1=0.5759 | wait=2
epoch 7/30 | train_f1=0.8960 | val_f1=0.5541 | wait=3
epoch 8/30 | train_f1=0.9079 | val_f1=0.5762 | wait=4
Early stop at epoch 8

--- Trial: image_age | encoded32 ---
epoch 1/30 | train_f1=0.5132 | val_f1=0.4613 | wait=0
epoch 2/30 | train_f1=0.6926 | val_f1=0.5264 | wait=0
epoch 3/30 | train_f1=0.7738 | val_f1=0.5422 | wait=0
epoch 4/30 | train_f1=0.8315 | val_f1=0.5604 | wait=0
epoch 5/30 | train_f1=0.8541 | val_f1=0.5623 | wait=0
epoch 6/30 | train_f1=0.8836 | val_f1=0.5414 | wait=0
epoch 7/30 | train_f1=0.9015 | val_f1=0.5625 | wait=1
epoch 8/30 | train_f1=0.9081 | val_f1=0.5887 | wait=0
epoch 9/30 | train_f1=0.9217 | val_f1=0.5780 | wait=0
epoch 10/30 | train_f1=0.9311 | val_f1=0.5925 | wait=1
epoch 11/30 | train_f1=0.9427 | val_f1=0.5798 | wait=0
epoch 12/30 | train_f1=0.9436 | val_f1=0.5835 | wait=1
epoch 13/30 | train_f1=0.9517 | val_f1=0.5903 | wait=2
epoch 14/30 | train_f1=0.9643 | val_f1=0.5768 | wait=3
epoch 15/30 | train_f1=0.9694 | val_f1=0.5725 | wait=4
Early stop at epoch 15

--- Trial: image_age | encoded64 ---
epoch 1/30 | train_f1=0.4601 | val_f1=0.5235 | wait=0
epoch 2/30 | train_f1=0.7106 | val_f1=0.5306 | wait=0
epoch 3/30 | train_f1=0.7840 | val_f1=0.5502 | wait=0
epoch 4/30 | train_f1=0.8313 | val_f1=0.5923 | wait=0
epoch 5/30 | train_f1=0.8543 | val_f1=0.5773 | wait=0
epoch 6/30 | train_f1=0.8787 | val_f1=0.5516 | wait=1
epoch 7/30 | train_f1=0.8922 | val_f1=0.5498 | wait=2
epoch 8/30 | train_f1=0.9096 | val_f1=0.5782 | wait=3
epoch 9/30 | train_f1=0.9258 | val_f1=0.5648 | wait=4
Early stop at epoch 9

--- Trial: image_age | encoded128 ---
epoch 1/30 | train_f1=0.4635 | val_f1=0.5223 | wait=0
epoch 2/30 | train_f1=0.7072 | val_f1=0.5146 | wait=0
epoch 3/30 | train_f1=0.7805 | val_f1=0.5346 | wait=1
epoch 4/30 | train_f1=0.8156 | val_f1=0.5693 | wait=0
epoch 5/30 | train_f1=0.8535 | val_f1=0.5814 | wait=0
epoch 6/30 | train_f1=0.8705 | val_f1=0.5680 | wait=0
epoch 7/30 | train_f1=0.8959 | val_f1=0.5711 | wait=1
epoch 8/30 | train_f1=0.9033 | val_f1=0.5731 | wait=2
epoch 9/30 | train_f1=0.9273 | val_f1=0.5657 | wait=3
epoch 10/30 | train_f1=0.9319 | val_f1=0.5879 | wait=4
epoch 11/30 | train_f1=0.9422 | val_f1=0.5683 | wait=0
epoch 12/30 | train_f1=0.9439 | val_f1=0.5657 | wait=1
epoch 13/30 | train_f1=0.9562 | val_f1=0.5695 | wait=2
epoch 14/30 | train_f1=0.9614 | val_f1=0.5949 | wait=3
epoch 15/30 | train_f1=0.9631 | val_f1=0.5824 | wait=0
epoch 16/30 | train_f1=0.9648 | val_f1=0.5614 | wait=1
epoch 17/30 | train_f1=0.9681 | val_f1=0.5663 | wait=2
epoch 18/30 | train_f1=0.9740 | val_f1=0.5745 | wait=3
epoch 19/30 | train_f1=0.9768 | val_f1=0.5785 | wait=4
Early stop at epoch 19

--- Trial: image_age | encoded256 ---
epoch 1/30 | train_f1=0.4573 | val_f1=0.5050 | wait=0
epoch 2/30 | train_f1=0.6879 | val_f1=0.5495 | wait=0
epoch 3/30 | train_f1=0.7735 | val_f1=0.5336 | wait=0
epoch 4/30 | train_f1=0.8188 | val_f1=0.5376 | wait=1
epoch 5/30 | train_f1=0.8460 | val_f1=0.5569 | wait=2
epoch 6/30 | train_f1=0.8745 | val_f1=0.5827 | wait=0
epoch 7/30 | train_f1=0.8839 | val_f1=0.5642 | wait=0
epoch 8/30 | train_f1=0.9045 | val_f1=0.5832 | wait=1
epoch 9/30 | train_f1=0.9188 | val_f1=0.5840 | wait=0
epoch 10/30 | train_f1=0.9260 | val_f1=0.5737 | wait=0
epoch 11/30 | train_f1=0.9316 | val_f1=0.5658 | wait=1
epoch 12/30 | train_f1=0.9437 | val_f1=0.5849 | wait=2
epoch 13/30 | train_f1=0.9487 | val_f1=0.5865 | wait=0
epoch 14/30 | train_f1=0.9479 | val_f1=0.6024 | wait=0
epoch 15/30 | train_f1=0.9583 | val_f1=0.5621 | wait=0
epoch 16/30 | train_f1=0.9562 | val_f1=0.5873 | wait=1
epoch 17/30 | train_f1=0.9656 | val_f1=0.5723 | wait=2
epoch 18/30 | train_f1=0.9642 | val_f1=0.5662 | wait=3
epoch 19/30 | train_f1=0.9712 | val_f1=0.5686 | wait=4
Early stop at epoch 19

========== image_sex ==========

--- Trial: image_sex | raw ---
epoch 1/30 | train_f1=0.4829 | val_f1=0.4927 | wait=0
epoch 2/30 | train_f1=0.6992 | val_f1=0.4894 | wait=0
epoch 3/30 | train_f1=0.7745 | val_f1=0.5170 | wait=1
epoch 4/30 | train_f1=0.8206 | val_f1=0.5462 | wait=0
epoch 5/30 | train_f1=0.8473 | val_f1=0.5452 | wait=0
epoch 6/30 | train_f1=0.8794 | val_f1=0.5835 | wait=1
epoch 7/30 | train_f1=0.8988 | val_f1=0.5699 | wait=0
epoch 8/30 | train_f1=0.9078 | val_f1=0.5711 | wait=1
epoch 9/30 | train_f1=0.9230 | val_f1=0.5827 | wait=2
epoch 10/30 | train_f1=0.9333 | val_f1=0.5709 | wait=3
epoch 11/30 | train_f1=0.9419 | val_f1=0.5878 | wait=4
epoch 12/30 | train_f1=0.9470 | val_f1=0.5757 | wait=0
epoch 13/30 | train_f1=0.9600 | val_f1=0.5937 | wait=1
epoch 14/30 | train_f1=0.9641 | val_f1=0.5734 | wait=0
epoch 15/30 | train_f1=0.9686 | val_f1=0.5850 | wait=1
epoch 16/30 | train_f1=0.9701 | val_f1=0.5773 | wait=2
epoch 17/30 | train_f1=0.9752 | val_f1=0.5799 | wait=3
epoch 18/30 | train_f1=0.9812 | val_f1=0.5866 | wait=4
Early stop at epoch 18

--- Trial: image_sex | encoded32 ---
epoch 1/30 | train_f1=0.4757 | val_f1=0.4932 | wait=0
epoch 2/30 | train_f1=0.7000 | val_f1=0.5231 | wait=0
epoch 3/30 | train_f1=0.7753 | val_f1=0.5559 | wait=0
epoch 4/30 | train_f1=0.8209 | val_f1=0.5470 | wait=0
epoch 5/30 | train_f1=0.8594 | val_f1=0.5762 | wait=1
epoch 6/30 | train_f1=0.8713 | val_f1=0.5845 | wait=0
epoch 7/30 | train_f1=0.8821 | val_f1=0.5740 | wait=0
epoch 8/30 | train_f1=0.9121 | val_f1=0.5824 | wait=1
epoch 9/30 | train_f1=0.9189 | val_f1=0.5699 | wait=2
epoch 10/30 | train_f1=0.9290 | val_f1=0.5645 | wait=3
epoch 11/30 | train_f1=0.9438 | val_f1=0.5847 | wait=4
epoch 12/30 | train_f1=0.9517 | val_f1=0.5767 | wait=0
epoch 13/30 | train_f1=0.9526 | val_f1=0.5747 | wait=1
epoch 14/30 | train_f1=0.9609 | val_f1=0.5804 | wait=2
epoch 15/30 | train_f1=0.9650 | val_f1=0.5852 | wait=3
epoch 16/30 | train_f1=0.9720 | val_f1=0.5692 | wait=0
epoch 17/30 | train_f1=0.9774 | val_f1=0.5789 | wait=1
epoch 18/30 | train_f1=0.9795 | val_f1=0.5892 | wait=2
epoch 19/30 | train_f1=0.9839 | val_f1=0.5834 | wait=0
epoch 20/30 | train_f1=0.9847 | val_f1=0.5811 | wait=1
epoch 21/30 | train_f1=0.9875 | val_f1=0.5894 | wait=2
epoch 22/30 | train_f1=0.9904 | val_f1=0.5940 | wait=0
epoch 23/30 | train_f1=0.9865 | val_f1=0.5689 | wait=0
epoch 24/30 | train_f1=0.9911 | val_f1=0.5860 | wait=1
epoch 25/30 | train_f1=0.9915 | val_f1=0.5921 | wait=2
epoch 26/30 | train_f1=0.9926 | val_f1=0.5760 | wait=3
epoch 27/30 | train_f1=0.9899 | val_f1=0.5843 | wait=4
Early stop at epoch 27

--- Trial: image_sex | encoded64 ---
epoch 1/30 | train_f1=0.4577 | val_f1=0.5059 | wait=0
epoch 2/30 | train_f1=0.6876 | val_f1=0.5242 | wait=0
epoch 3/30 | train_f1=0.7691 | val_f1=0.5229 | wait=0
epoch 4/30 | train_f1=0.8053 | val_f1=0.5486 | wait=1
epoch 5/30 | train_f1=0.8468 | val_f1=0.5711 | wait=0
epoch 6/30 | train_f1=0.8601 | val_f1=0.5815 | wait=0
epoch 7/30 | train_f1=0.8893 | val_f1=0.5675 | wait=0
epoch 8/30 | train_f1=0.9018 | val_f1=0.5645 | wait=1
epoch 9/30 | train_f1=0.9129 | val_f1=0.5702 | wait=2
epoch 10/30 | train_f1=0.9267 | val_f1=0.5836 | wait=3
epoch 11/30 | train_f1=0.9358 | val_f1=0.5841 | wait=0
epoch 12/30 | train_f1=0.9428 | val_f1=0.5929 | wait=0
epoch 13/30 | train_f1=0.9534 | val_f1=0.5635 | wait=0
epoch 14/30 | train_f1=0.9554 | val_f1=0.5400 | wait=1
epoch 15/30 | train_f1=0.9584 | val_f1=0.5793 | wait=2
epoch 16/30 | train_f1=0.9685 | val_f1=0.5759 | wait=3
epoch 17/30 | train_f1=0.9759 | val_f1=0.5821 | wait=4
Early stop at epoch 17

--- Trial: image_sex | encoded128 ---
epoch 1/30 | train_f1=0.4586 | val_f1=0.4841 | wait=0
epoch 2/30 | train_f1=0.6943 | val_f1=0.5322 | wait=0
epoch 3/30 | train_f1=0.7690 | val_f1=0.5707 | wait=0
epoch 4/30 | train_f1=0.8213 | val_f1=0.5591 | wait=0
epoch 5/30 | train_f1=0.8537 | val_f1=0.5568 | wait=1
epoch 6/30 | train_f1=0.8724 | val_f1=0.5631 | wait=2
epoch 7/30 | train_f1=0.8882 | val_f1=0.5759 | wait=3
epoch 8/30 | train_f1=0.9033 | val_f1=0.5507 | wait=0
epoch 9/30 | train_f1=0.9157 | val_f1=0.5769 | wait=1
epoch 10/30 | train_f1=0.9261 | val_f1=0.5815 | wait=0
epoch 11/30 | train_f1=0.9373 | val_f1=0.5849 | wait=0
epoch 12/30 | train_f1=0.9427 | val_f1=0.5718 | wait=0
epoch 13/30 | train_f1=0.9487 | val_f1=0.5932 | wait=1
epoch 14/30 | train_f1=0.9567 | val_f1=0.5727 | wait=0
epoch 15/30 | train_f1=0.9618 | val_f1=0.5758 | wait=1
epoch 16/30 | train_f1=0.9676 | val_f1=0.5757 | wait=2
epoch 17/30 | train_f1=0.9686 | val_f1=0.5685 | wait=3
epoch 18/30 | train_f1=0.9781 | val_f1=0.5878 | wait=4
Early stop at epoch 18

--- Trial: image_sex | encoded256 ---
epoch 1/30 | train_f1=0.3915 | val_f1=0.4578 | wait=0
epoch 2/30 | train_f1=0.6626 | val_f1=0.5414 | wait=0
epoch 3/30 | train_f1=0.7569 | val_f1=0.5279 | wait=0
epoch 4/30 | train_f1=0.8005 | val_f1=0.5367 | wait=1
epoch 5/30 | train_f1=0.8328 | val_f1=0.5601 | wait=2
epoch 6/30 | train_f1=0.8624 | val_f1=0.5724 | wait=0
epoch 7/30 | train_f1=0.8739 | val_f1=0.5788 | wait=0
epoch 8/30 | train_f1=0.8864 | val_f1=0.5928 | wait=0
epoch 9/30 | train_f1=0.8985 | val_f1=0.5710 | wait=0
epoch 10/30 | train_f1=0.9138 | val_f1=0.5719 | wait=1
epoch 11/30 | train_f1=0.9283 | val_f1=0.5914 | wait=2
epoch 12/30 | train_f1=0.9306 | val_f1=0.5758 | wait=3
epoch 13/30 | train_f1=0.9396 | val_f1=0.5768 | wait=4
Early stop at epoch 13

========== image_location ==========

--- Trial: image_location | raw ---
epoch 1/30 | train_f1=0.4784 | val_f1=0.5012 | wait=0
epoch 2/30 | train_f1=0.7062 | val_f1=0.5656 | wait=0
epoch 3/30 | train_f1=0.7933 | val_f1=0.5412 | wait=0
epoch 4/30 | train_f1=0.8393 | val_f1=0.5476 | wait=1
epoch 5/30 | train_f1=0.8595 | val_f1=0.5839 | wait=2
epoch 6/30 | train_f1=0.8772 | val_f1=0.5711 | wait=0
epoch 7/30 | train_f1=0.9000 | val_f1=0.5828 | wait=1
epoch 8/30 | train_f1=0.9119 | val_f1=0.5744 | wait=2
epoch 9/30 | train_f1=0.9250 | val_f1=0.5556 | wait=3
epoch 10/30 | train_f1=0.9332 | val_f1=0.5760 | wait=4
Early stop at epoch 10

--- Trial: image_location | encoded32 ---
epoch 1/30 | train_f1=0.5021 | val_f1=0.5144 | wait=0
epoch 2/30 | train_f1=0.7180 | val_f1=0.5463 | wait=0
epoch 3/30 | train_f1=0.7824 | val_f1=0.5345 | wait=0
epoch 4/30 | train_f1=0.8255 | val_f1=0.5603 | wait=1
epoch 5/30 | train_f1=0.8532 | val_f1=0.5646 | wait=0
epoch 6/30 | train_f1=0.8703 | val_f1=0.5775 | wait=0
epoch 7/30 | train_f1=0.8939 | val_f1=0.5883 | wait=0
epoch 8/30 | train_f1=0.9107 | val_f1=0.5818 | wait=0
epoch 9/30 | train_f1=0.9222 | val_f1=0.5728 | wait=1
epoch 10/30 | train_f1=0.9294 | val_f1=0.5727 | wait=2
epoch 11/30 | train_f1=0.9474 | val_f1=0.5888 | wait=3
epoch 12/30 | train_f1=0.9504 | val_f1=0.5603 | wait=0
epoch 13/30 | train_f1=0.9533 | val_f1=0.5675 | wait=1
epoch 14/30 | train_f1=0.9589 | val_f1=0.5755 | wait=2
epoch 15/30 | train_f1=0.9629 | val_f1=0.5760 | wait=3
epoch 16/30 | train_f1=0.9675 | val_f1=0.5900 | wait=4
epoch 17/30 | train_f1=0.9700 | val_f1=0.5945 | wait=0
epoch 18/30 | train_f1=0.9796 | val_f1=0.5898 | wait=0
epoch 19/30 | train_f1=0.9846 | val_f1=0.5801 | wait=1
epoch 20/30 | train_f1=0.9833 | val_f1=0.5826 | wait=2
epoch 21/30 | train_f1=0.9825 | val_f1=0.5782 | wait=3
epoch 22/30 | train_f1=0.9902 | val_f1=0.5746 | wait=4
Early stop at epoch 22

--- Trial: image_location | encoded64 ---
epoch 1/30 | train_f1=0.4983 | val_f1=0.5289 | wait=0
epoch 2/30 | train_f1=0.7102 | val_f1=0.5201 | wait=0
epoch 3/30 | train_f1=0.7784 | val_f1=0.5256 | wait=1
epoch 4/30 | train_f1=0.8228 | val_f1=0.5602 | wait=2
epoch 5/30 | train_f1=0.8532 | val_f1=0.5716 | wait=0
epoch 6/30 | train_f1=0.8816 | val_f1=0.5605 | wait=0
epoch 7/30 | train_f1=0.8964 | val_f1=0.5451 | wait=1
epoch 8/30 | train_f1=0.9069 | val_f1=0.5776 | wait=2
epoch 9/30 | train_f1=0.9212 | val_f1=0.5780 | wait=0
epoch 10/30 | train_f1=0.9265 | val_f1=0.5694 | wait=0
epoch 11/30 | train_f1=0.9406 | val_f1=0.5875 | wait=1
epoch 12/30 | train_f1=0.9454 | val_f1=0.5941 | wait=0
epoch 13/30 | train_f1=0.9534 | val_f1=0.5935 | wait=0
epoch 14/30 | train_f1=0.9585 | val_f1=0.6029 | wait=1
epoch 15/30 | train_f1=0.9683 | val_f1=0.5936 | wait=0
epoch 16/30 | train_f1=0.9708 | val_f1=0.5863 | wait=1
epoch 17/30 | train_f1=0.9708 | val_f1=0.5686 | wait=2
epoch 18/30 | train_f1=0.9789 | val_f1=0.5989 | wait=3
epoch 19/30 | train_f1=0.9803 | val_f1=0.6029 | wait=4
Early stop at epoch 19

--- Trial: image_location | encoded128 ---
epoch 1/30 | train_f1=0.4688 | val_f1=0.5062 | wait=0
epoch 2/30 | train_f1=0.6948 | val_f1=0.5494 | wait=0
epoch 3/30 | train_f1=0.7897 | val_f1=0.5701 | wait=0
epoch 4/30 | train_f1=0.8195 | val_f1=0.5313 | wait=0
epoch 5/30 | train_f1=0.8411 | val_f1=0.5567 | wait=1
epoch 6/30 | train_f1=0.8708 | val_f1=0.5829 | wait=2
epoch 7/30 | train_f1=0.8873 | val_f1=0.5765 | wait=0
epoch 8/30 | train_f1=0.9028 | val_f1=0.5911 | wait=1
epoch 9/30 | train_f1=0.9160 | val_f1=0.5994 | wait=0
epoch 10/30 | train_f1=0.9316 | val_f1=0.5953 | wait=0
epoch 11/30 | train_f1=0.9432 | val_f1=0.5717 | wait=1
epoch 12/30 | train_f1=0.9512 | val_f1=0.5907 | wait=2
epoch 13/30 | train_f1=0.9522 | val_f1=0.5857 | wait=3
epoch 14/30 | train_f1=0.9528 | val_f1=0.5871 | wait=4
Early stop at epoch 14

--- Trial: image_location | encoded256 ---
epoch 1/30 | train_f1=0.4480 | val_f1=0.4708 | wait=0
epoch 2/30 | train_f1=0.6808 | val_f1=0.5177 | wait=0
epoch 3/30 | train_f1=0.7657 | val_f1=0.5660 | wait=0
epoch 4/30 | train_f1=0.8002 | val_f1=0.5732 | wait=0
epoch 5/30 | train_f1=0.8328 | val_f1=0.5709 | wait=0
epoch 6/30 | train_f1=0.8605 | val_f1=0.5729 | wait=1
epoch 7/30 | train_f1=0.8835 | val_f1=0.5807 | wait=2
epoch 8/30 | train_f1=0.9029 | val_f1=0.5674 | wait=0
epoch 9/30 | train_f1=0.9090 | val_f1=0.5860 | wait=1
epoch 10/30 | train_f1=0.9190 | val_f1=0.5882 | wait=0
epoch 11/30 | train_f1=0.9232 | val_f1=0.5845 | wait=0
epoch 12/30 | train_f1=0.9446 | val_f1=0.5860 | wait=1
epoch 13/30 | train_f1=0.9487 | val_f1=0.5882 | wait=2
epoch 14/30 | train_f1=0.9529 | val_f1=0.5791 | wait=3
epoch 15/30 | train_f1=0.9602 | val_f1=0.5792 | wait=4
Early stop at epoch 15

========== image_age_sex ==========

--- Trial: image_age_sex | raw ---
epoch 1/30 | train_f1=0.4623 | val_f1=0.4909 | wait=0
epoch 2/30 | train_f1=0.7022 | val_f1=0.5436 | wait=0
epoch 3/30 | train_f1=0.7819 | val_f1=0.5803 | wait=0
epoch 4/30 | train_f1=0.8141 | val_f1=0.5441 | wait=0
epoch 5/30 | train_f1=0.8499 | val_f1=0.5609 | wait=1
epoch 6/30 | train_f1=0.8724 | val_f1=0.5755 | wait=2
epoch 7/30 | train_f1=0.8947 | val_f1=0.5697 | wait=3
epoch 8/30 | train_f1=0.9106 | val_f1=0.5886 | wait=4
epoch 9/30 | train_f1=0.9238 | val_f1=0.5601 | wait=0
epoch 10/30 | train_f1=0.9306 | val_f1=0.5811 | wait=1
epoch 11/30 | train_f1=0.9455 | val_f1=0.5840 | wait=2
epoch 12/30 | train_f1=0.9490 | val_f1=0.5751 | wait=3
epoch 13/30 | train_f1=0.9517 | val_f1=0.5827 | wait=4
Early stop at epoch 13

--- Trial: image_age_sex | encoded32 ---
epoch 1/30 | train_f1=0.4853 | val_f1=0.4914 | wait=0
epoch 2/30 | train_f1=0.6917 | val_f1=0.5338 | wait=0
epoch 3/30 | train_f1=0.7869 | val_f1=0.5600 | wait=0
epoch 4/30 | train_f1=0.8216 | val_f1=0.5453 | wait=0
epoch 5/30 | train_f1=0.8453 | val_f1=0.5584 | wait=1
epoch 6/30 | train_f1=0.8816 | val_f1=0.5514 | wait=2
epoch 7/30 | train_f1=0.8994 | val_f1=0.5825 | wait=3
epoch 8/30 | train_f1=0.9100 | val_f1=0.5797 | wait=0
epoch 9/30 | train_f1=0.9209 | val_f1=0.5728 | wait=1
epoch 10/30 | train_f1=0.9373 | val_f1=0.5940 | wait=2
epoch 11/30 | train_f1=0.9452 | val_f1=0.5939 | wait=0
epoch 12/30 | train_f1=0.9523 | val_f1=0.5831 | wait=1
epoch 13/30 | train_f1=0.9486 | val_f1=0.5961 | wait=2
epoch 14/30 | train_f1=0.9584 | val_f1=0.5792 | wait=0
epoch 15/30 | train_f1=0.9676 | val_f1=0.5790 | wait=1
epoch 16/30 | train_f1=0.9730 | val_f1=0.5824 | wait=2
epoch 17/30 | train_f1=0.9768 | val_f1=0.5887 | wait=3
epoch 18/30 | train_f1=0.9773 | val_f1=0.5860 | wait=4
Early stop at epoch 18

--- Trial: image_age_sex | encoded64 ---
epoch 1/30 | train_f1=0.4863 | val_f1=0.4959 | wait=0
epoch 2/30 | train_f1=0.7007 | val_f1=0.5478 | wait=0
epoch 3/30 | train_f1=0.7897 | val_f1=0.5415 | wait=0
epoch 4/30 | train_f1=0.8298 | val_f1=0.5563 | wait=1
epoch 5/30 | train_f1=0.8521 | val_f1=0.5783 | wait=0
epoch 6/30 | train_f1=0.8790 | val_f1=0.5871 | wait=0
epoch 7/30 | train_f1=0.9013 | val_f1=0.5735 | wait=0
epoch 8/30 | train_f1=0.9144 | val_f1=0.5805 | wait=1
epoch 9/30 | train_f1=0.9264 | val_f1=0.5737 | wait=2
epoch 10/30 | train_f1=0.9354 | val_f1=0.5822 | wait=3
epoch 11/30 | train_f1=0.9324 | val_f1=0.5788 | wait=4
Early stop at epoch 11

--- Trial: image_age_sex | encoded128 ---
epoch 1/30 | train_f1=0.4559 | val_f1=0.5135 | wait=0
epoch 2/30 | train_f1=0.6891 | val_f1=0.5264 | wait=0
epoch 3/30 | train_f1=0.7722 | val_f1=0.5790 | wait=0
epoch 4/30 | train_f1=0.8237 | val_f1=0.5540 | wait=0
epoch 5/30 | train_f1=0.8541 | val_f1=0.5740 | wait=1
epoch 6/30 | train_f1=0.8712 | val_f1=0.5678 | wait=2
epoch 7/30 | train_f1=0.8787 | val_f1=0.5879 | wait=3
epoch 8/30 | train_f1=0.9084 | val_f1=0.5361 | wait=0
epoch 9/30 | train_f1=0.9186 | val_f1=0.5614 | wait=1
epoch 10/30 | train_f1=0.9298 | val_f1=0.5852 | wait=2
epoch 11/30 | train_f1=0.9396 | val_f1=0.5955 | wait=3
epoch 12/30 | train_f1=0.9443 | val_f1=0.5910 | wait=0
epoch 13/30 | train_f1=0.9526 | val_f1=0.5862 | wait=1
epoch 14/30 | train_f1=0.9535 | val_f1=0.5908 | wait=2
epoch 15/30 | train_f1=0.9618 | val_f1=0.5808 | wait=3
epoch 16/30 | train_f1=0.9684 | val_f1=0.5778 | wait=4
Early stop at epoch 16

--- Trial: image_age_sex | encoded256 ---
epoch 1/30 | train_f1=0.4567 | val_f1=0.4435 | wait=0
epoch 2/30 | train_f1=0.6912 | val_f1=0.5330 | wait=0
epoch 3/30 | train_f1=0.7619 | val_f1=0.5436 | wait=0
epoch 4/30 | train_f1=0.8098 | val_f1=0.5290 | wait=0
epoch 5/30 | train_f1=0.8319 | val_f1=0.5611 | wait=1
epoch 6/30 | train_f1=0.8682 | val_f1=0.5775 | wait=0
epoch 7/30 | train_f1=0.8881 | val_f1=0.5767 | wait=0
epoch 8/30 | train_f1=0.9068 | val_f1=0.5955 | wait=1
epoch 9/30 | train_f1=0.9059 | val_f1=0.5822 | wait=0
epoch 10/30 | train_f1=0.9260 | val_f1=0.5824 | wait=1
epoch 11/30 | train_f1=0.9335 | val_f1=0.5690 | wait=2
epoch 12/30 | train_f1=0.9398 | val_f1=0.5934 | wait=3
epoch 13/30 | train_f1=0.9455 | val_f1=0.5908 | wait=4
Early stop at epoch 13

========== image_age_location ==========

--- Trial: image_age_location | raw ---
epoch 1/30 | train_f1=0.4748 | val_f1=0.5098 | wait=0
epoch 2/30 | train_f1=0.7119 | val_f1=0.5480 | wait=0
epoch 3/30 | train_f1=0.7960 | val_f1=0.5632 | wait=0
epoch 4/30 | train_f1=0.8299 | val_f1=0.5685 | wait=0
epoch 5/30 | train_f1=0.8481 | val_f1=0.5700 | wait=0
epoch 6/30 | train_f1=0.8781 | val_f1=0.5816 | wait=0
epoch 7/30 | train_f1=0.9037 | val_f1=0.6066 | wait=0
epoch 8/30 | train_f1=0.9199 | val_f1=0.5966 | wait=0
epoch 9/30 | train_f1=0.9225 | val_f1=0.5758 | wait=1
epoch 10/30 | train_f1=0.9406 | val_f1=0.5958 | wait=2
epoch 11/30 | train_f1=0.9430 | val_f1=0.5846 | wait=3
epoch 12/30 | train_f1=0.9547 | val_f1=0.6025 | wait=4
Early stop at epoch 12

--- Trial: image_age_location | encoded32 ---
epoch 1/30 | train_f1=0.4802 | val_f1=0.5103 | wait=0
epoch 2/30 | train_f1=0.7094 | val_f1=0.5460 | wait=0
epoch 3/30 | train_f1=0.7906 | val_f1=0.5620 | wait=0
epoch 4/30 | train_f1=0.8348 | val_f1=0.5796 | wait=0
epoch 5/30 | train_f1=0.8601 | val_f1=0.5674 | wait=0
epoch 6/30 | train_f1=0.8860 | val_f1=0.5820 | wait=1
epoch 7/30 | train_f1=0.8991 | val_f1=0.5904 | wait=0
epoch 8/30 | train_f1=0.9174 | val_f1=0.6010 | wait=0
epoch 9/30 | train_f1=0.9280 | val_f1=0.5830 | wait=0
epoch 10/30 | train_f1=0.9418 | val_f1=0.5926 | wait=1
epoch 11/30 | train_f1=0.9440 | val_f1=0.5852 | wait=2
epoch 12/30 | train_f1=0.9493 | val_f1=0.5838 | wait=3
epoch 13/30 | train_f1=0.9584 | val_f1=0.5917 | wait=4
Early stop at epoch 13

--- Trial: image_age_location | encoded64 ---
epoch 1/30 | train_f1=0.4926 | val_f1=0.5616 | wait=0
epoch 2/30 | train_f1=0.7249 | val_f1=0.5538 | wait=0
epoch 3/30 | train_f1=0.8025 | val_f1=0.5507 | wait=1
epoch 4/30 | train_f1=0.8263 | val_f1=0.5545 | wait=2
epoch 5/30 | train_f1=0.8609 | val_f1=0.5686 | wait=3
epoch 6/30 | train_f1=0.8898 | val_f1=0.5729 | wait=0
epoch 7/30 | train_f1=0.8972 | val_f1=0.5754 | wait=0
epoch 8/30 | train_f1=0.9122 | val_f1=0.5542 | wait=0
epoch 9/30 | train_f1=0.9288 | val_f1=0.5876 | wait=1
epoch 10/30 | train_f1=0.9363 | val_f1=0.5993 | wait=0
epoch 11/30 | train_f1=0.9468 | val_f1=0.5872 | wait=0
epoch 12/30 | train_f1=0.9564 | val_f1=0.5829 | wait=1
epoch 13/30 | train_f1=0.9563 | val_f1=0.5735 | wait=2
epoch 14/30 | train_f1=0.9660 | val_f1=0.5833 | wait=3
epoch 15/30 | train_f1=0.9671 | val_f1=0.6024 | wait=4
epoch 16/30 | train_f1=0.9747 | val_f1=0.5887 | wait=0
epoch 17/30 | train_f1=0.9768 | val_f1=0.5961 | wait=1
epoch 18/30 | train_f1=0.9762 | val_f1=0.5954 | wait=2
epoch 19/30 | train_f1=0.9835 | val_f1=0.6034 | wait=3
epoch 20/30 | train_f1=0.9832 | val_f1=0.5775 | wait=0
epoch 21/30 | train_f1=0.9833 | val_f1=0.5970 | wait=1
epoch 22/30 | train_f1=0.9881 | val_f1=0.5955 | wait=2
epoch 23/30 | train_f1=0.9861 | val_f1=0.5936 | wait=3
epoch 24/30 | train_f1=0.9938 | val_f1=0.5810 | wait=4
Early stop at epoch 24

--- Trial: image_age_location | encoded128 ---
epoch 1/30 | train_f1=0.4832 | val_f1=0.4966 | wait=0
epoch 2/30 | train_f1=0.7062 | val_f1=0.5480 | wait=0
epoch 3/30 | train_f1=0.7866 | val_f1=0.5356 | wait=0
epoch 4/30 | train_f1=0.8255 | val_f1=0.5682 | wait=1
epoch 5/30 | train_f1=0.8539 | val_f1=0.5717 | wait=0
epoch 6/30 | train_f1=0.8641 | val_f1=0.5892 | wait=0
epoch 7/30 | train_f1=0.8998 | val_f1=0.5762 | wait=0
epoch 8/30 | train_f1=0.9068 | val_f1=0.5845 | wait=1
epoch 9/30 | train_f1=0.9296 | val_f1=0.5813 | wait=2
epoch 10/30 | train_f1=0.9214 | val_f1=0.5851 | wait=3
epoch 11/30 | train_f1=0.9444 | val_f1=0.5769 | wait=4
Early stop at epoch 11

--- Trial: image_age_location | encoded256 ---
epoch 1/30 | train_f1=0.4743 | val_f1=0.4662 | wait=0
epoch 2/30 | train_f1=0.6807 | val_f1=0.5566 | wait=0
epoch 3/30 | train_f1=0.7754 | val_f1=0.5539 | wait=0
epoch 4/30 | train_f1=0.8187 | val_f1=0.5908 | wait=1
epoch 5/30 | train_f1=0.8489 | val_f1=0.5699 | wait=0
epoch 6/30 | train_f1=0.8679 | val_f1=0.5592 | wait=1
epoch 7/30 | train_f1=0.8842 | val_f1=0.5610 | wait=2
epoch 8/30 | train_f1=0.9164 | val_f1=0.5741 | wait=3
epoch 9/30 | train_f1=0.9269 | val_f1=0.5870 | wait=4
Early stop at epoch 9

========== image_sex_location ==========

--- Trial: image_sex_location | raw ---
epoch 1/30 | train_f1=0.4777 | val_f1=0.4859 | wait=0
epoch 2/30 | train_f1=0.6974 | val_f1=0.5402 | wait=0
epoch 3/30 | train_f1=0.7849 | val_f1=0.5939 | wait=0
epoch 4/30 | train_f1=0.8272 | val_f1=0.5606 | wait=0
epoch 5/30 | train_f1=0.8574 | val_f1=0.5733 | wait=1
epoch 6/30 | train_f1=0.8798 | val_f1=0.5569 | wait=2
epoch 7/30 | train_f1=0.8987 | val_f1=0.5869 | wait=3
epoch 8/30 | train_f1=0.9099 | val_f1=0.5896 | wait=4
Early stop at epoch 8

--- Trial: image_sex_location | encoded32 ---
epoch 1/30 | train_f1=0.4661 | val_f1=0.5195 | wait=0
epoch 2/30 | train_f1=0.7095 | val_f1=0.5494 | wait=0
epoch 3/30 | train_f1=0.7739 | val_f1=0.5472 | wait=0
epoch 4/30 | train_f1=0.8321 | val_f1=0.5611 | wait=1
epoch 5/30 | train_f1=0.8574 | val_f1=0.5748 | wait=0
epoch 6/30 | train_f1=0.8785 | val_f1=0.5578 | wait=0
epoch 7/30 | train_f1=0.9014 | val_f1=0.5755 | wait=1
epoch 8/30 | train_f1=0.9190 | val_f1=0.5986 | wait=0
epoch 9/30 | train_f1=0.9312 | val_f1=0.5911 | wait=0
epoch 10/30 | train_f1=0.9394 | val_f1=0.5764 | wait=1
epoch 11/30 | train_f1=0.9487 | val_f1=0.5750 | wait=2
epoch 12/30 | train_f1=0.9563 | val_f1=0.5780 | wait=3
epoch 13/30 | train_f1=0.9584 | val_f1=0.5797 | wait=4
Early stop at epoch 13

--- Trial: image_sex_location | encoded64 ---
epoch 1/30 | train_f1=0.4910 | val_f1=0.5517 | wait=0
epoch 2/30 | train_f1=0.7138 | val_f1=0.5722 | wait=0
epoch 3/30 | train_f1=0.7767 | val_f1=0.5668 | wait=0
epoch 4/30 | train_f1=0.8282 | val_f1=0.5743 | wait=1
epoch 5/30 | train_f1=0.8432 | val_f1=0.5568 | wait=0
epoch 6/30 | train_f1=0.8747 | val_f1=0.5671 | wait=1
epoch 7/30 | train_f1=0.8923 | val_f1=0.5622 | wait=2
epoch 8/30 | train_f1=0.9086 | val_f1=0.5617 | wait=3
epoch 9/30 | train_f1=0.9196 | val_f1=0.5698 | wait=4
Early stop at epoch 9

--- Trial: image_sex_location | encoded128 ---
epoch 1/30 | train_f1=0.4941 | val_f1=0.4997 | wait=0
epoch 2/30 | train_f1=0.6924 | val_f1=0.5428 | wait=0
epoch 3/30 | train_f1=0.7702 | val_f1=0.5625 | wait=0
epoch 4/30 | train_f1=0.8215 | val_f1=0.5621 | wait=0
epoch 5/30 | train_f1=0.8533 | val_f1=0.6103 | wait=1
epoch 6/30 | train_f1=0.8664 | val_f1=0.5554 | wait=0
epoch 7/30 | train_f1=0.8886 | val_f1=0.5812 | wait=1
epoch 8/30 | train_f1=0.9089 | val_f1=0.6040 | wait=2
epoch 9/30 | train_f1=0.9202 | val_f1=0.6130 | wait=3
epoch 10/30 | train_f1=0.9192 | val_f1=0.6020 | wait=0
epoch 11/30 | train_f1=0.9336 | val_f1=0.5946 | wait=1
epoch 12/30 | train_f1=0.9433 | val_f1=0.5855 | wait=2
epoch 13/30 | train_f1=0.9545 | val_f1=0.5893 | wait=3
epoch 14/30 | train_f1=0.9582 | val_f1=0.5884 | wait=4
Early stop at epoch 14

--- Trial: image_sex_location | encoded256 ---
epoch 1/30 | train_f1=0.4936 | val_f1=0.5144 | wait=0
epoch 2/30 | train_f1=0.6825 | val_f1=0.5387 | wait=0
epoch 3/30 | train_f1=0.7695 | val_f1=0.5461 | wait=0
epoch 4/30 | train_f1=0.8191 | val_f1=0.5389 | wait=0
epoch 5/30 | train_f1=0.8432 | val_f1=0.5929 | wait=1
epoch 6/30 | train_f1=0.8697 | val_f1=0.5738 | wait=0
epoch 7/30 | train_f1=0.8880 | val_f1=0.5810 | wait=1
epoch 8/30 | train_f1=0.9013 | val_f1=0.5777 | wait=2
epoch 9/30 | train_f1=0.9030 | val_f1=0.6011 | wait=3
epoch 10/30 | train_f1=0.9221 | val_f1=0.5817 | wait=0
epoch 11/30 | train_f1=0.9324 | val_f1=0.5929 | wait=1
epoch 12/30 | train_f1=0.9311 | val_f1=0.5833 | wait=2
epoch 13/30 | train_f1=0.9503 | val_f1=0.5840 | wait=3
epoch 14/30 | train_f1=0.9535 | val_f1=0.5955 | wait=4
Early stop at epoch 14

========== image_all_metadata ==========

--- Trial: image_all_metadata | raw ---
epoch 1/30 | train_f1=0.4706 | val_f1=0.4907 | wait=0
epoch 2/30 | train_f1=0.7042 | val_f1=0.5276 | wait=0
epoch 3/30 | train_f1=0.7828 | val_f1=0.5346 | wait=0
epoch 4/30 | train_f1=0.8329 | val_f1=0.5408 | wait=0
epoch 5/30 | train_f1=0.8540 | val_f1=0.5666 | wait=0
epoch 6/30 | train_f1=0.8809 | val_f1=0.5556 | wait=0
epoch 7/30 | train_f1=0.8959 | val_f1=0.5717 | wait=1
epoch 8/30 | train_f1=0.9159 | val_f1=0.5794 | wait=0
epoch 9/30 | train_f1=0.9281 | val_f1=0.5846 | wait=0
epoch 10/30 | train_f1=0.9354 | val_f1=0.5896 | wait=0
epoch 11/30 | train_f1=0.9403 | val_f1=0.5917 | wait=0
epoch 12/30 | train_f1=0.9513 | val_f1=0.5998 | wait=0
epoch 13/30 | train_f1=0.9618 | val_f1=0.5935 | wait=0
epoch 14/30 | train_f1=0.9630 | val_f1=0.5939 | wait=1
epoch 15/30 | train_f1=0.9663 | val_f1=0.5892 | wait=2
epoch 16/30 | train_f1=0.9691 | val_f1=0.5833 | wait=3
epoch 17/30 | train_f1=0.9788 | val_f1=0.6003 | wait=4
epoch 18/30 | train_f1=0.9793 | val_f1=0.5971 | wait=0
epoch 19/30 | train_f1=0.9826 | val_f1=0.5750 | wait=1
epoch 20/30 | train_f1=0.9831 | val_f1=0.5723 | wait=2
epoch 21/30 | train_f1=0.9857 | val_f1=0.5844 | wait=3
epoch 22/30 | train_f1=0.9884 | val_f1=0.5920 | wait=4
Early stop at epoch 22

--- Trial: image_all_metadata | encoded32 ---
epoch 1/30 | train_f1=0.4950 | val_f1=0.5364 | wait=0
epoch 2/30 | train_f1=0.6997 | val_f1=0.5511 | wait=0
epoch 3/30 | train_f1=0.7856 | val_f1=0.5505 | wait=0
epoch 4/30 | train_f1=0.8253 | val_f1=0.5432 | wait=1
epoch 5/30 | train_f1=0.8530 | val_f1=0.5610 | wait=2
epoch 6/30 | train_f1=0.8842 | val_f1=0.5919 | wait=0
epoch 7/30 | train_f1=0.8966 | val_f1=0.6052 | wait=0
epoch 8/30 | train_f1=0.9112 | val_f1=0.5978 | wait=0
epoch 9/30 | train_f1=0.9269 | val_f1=0.5981 | wait=1
epoch 10/30 | train_f1=0.9336 | val_f1=0.6090 | wait=2
epoch 11/30 | train_f1=0.9400 | val_f1=0.6029 | wait=0
epoch 12/30 | train_f1=0.9459 | val_f1=0.6038 | wait=1
epoch 13/30 | train_f1=0.9580 | val_f1=0.6065 | wait=2
epoch 14/30 | train_f1=0.9656 | val_f1=0.6046 | wait=3
epoch 15/30 | train_f1=0.9699 | val_f1=0.6109 | wait=4
epoch 16/30 | train_f1=0.9758 | val_f1=0.5920 | wait=0
epoch 17/30 | train_f1=0.9743 | val_f1=0.5893 | wait=1
epoch 18/30 | train_f1=0.9797 | val_f1=0.6181 | wait=2
epoch 19/30 | train_f1=0.9820 | val_f1=0.6021 | wait=0
epoch 20/30 | train_f1=0.9878 | val_f1=0.5873 | wait=1
epoch 21/30 | train_f1=0.9882 | val_f1=0.5893 | wait=2
epoch 22/30 | train_f1=0.9876 | val_f1=0.5961 | wait=3
epoch 23/30 | train_f1=0.9895 | val_f1=0.5925 | wait=4
Early stop at epoch 23

--- Trial: image_all_metadata | encoded64 ---
epoch 1/30 | train_f1=0.4823 | val_f1=0.5134 | wait=0
epoch 2/30 | train_f1=0.7103 | val_f1=0.5400 | wait=0
epoch 3/30 | train_f1=0.7845 | val_f1=0.5476 | wait=0
epoch 4/30 | train_f1=0.8382 | val_f1=0.5625 | wait=0
epoch 5/30 | train_f1=0.8580 | val_f1=0.5712 | wait=0
epoch 6/30 | train_f1=0.8845 | val_f1=0.5726 | wait=0
epoch 7/30 | train_f1=0.8973 | val_f1=0.5785 | wait=0
epoch 8/30 | train_f1=0.9140 | val_f1=0.5888 | wait=0
epoch 9/30 | train_f1=0.9250 | val_f1=0.5832 | wait=0
epoch 10/30 | train_f1=0.9424 | val_f1=0.5825 | wait=1
epoch 11/30 | train_f1=0.9432 | val_f1=0.5737 | wait=2
epoch 12/30 | train_f1=0.9444 | val_f1=0.5898 | wait=3
epoch 13/30 | train_f1=0.9591 | val_f1=0.5936 | wait=0
epoch 14/30 | train_f1=0.9638 | val_f1=0.5881 | wait=0
epoch 15/30 | train_f1=0.9674 | val_f1=0.5897 | wait=1
epoch 16/30 | train_f1=0.9712 | val_f1=0.5864 | wait=2
epoch 17/30 | train_f1=0.9737 | val_f1=0.6109 | wait=3
epoch 18/30 | train_f1=0.9757 | val_f1=0.5888 | wait=0
epoch 19/30 | train_f1=0.9806 | val_f1=0.6012 | wait=1
epoch 20/30 | train_f1=0.9794 | val_f1=0.5938 | wait=2
epoch 21/30 | train_f1=0.9838 | val_f1=0.6047 | wait=3
epoch 22/30 | train_f1=0.9886 | val_f1=0.5878 | wait=4
Early stop at epoch 22

--- Trial: image_all_metadata | encoded128 ---
epoch 1/30 | train_f1=0.4868 | val_f1=0.5101 | wait=0
epoch 2/30 | train_f1=0.6932 | val_f1=0.5645 | wait=0
epoch 3/30 | train_f1=0.7871 | val_f1=0.5571 | wait=0
epoch 4/30 | train_f1=0.8258 | val_f1=0.5465 | wait=1
epoch 5/30 | train_f1=0.8540 | val_f1=0.5859 | wait=2
epoch 6/30 | train_f1=0.8784 | val_f1=0.5955 | wait=0
epoch 7/30 | train_f1=0.8954 | val_f1=0.5988 | wait=0
epoch 8/30 | train_f1=0.9124 | val_f1=0.5965 | wait=0
epoch 9/30 | train_f1=0.9227 | val_f1=0.5999 | wait=1
epoch 10/30 | train_f1=0.9383 | val_f1=0.5954 | wait=0
epoch 11/30 | train_f1=0.9420 | val_f1=0.5870 | wait=1
epoch 12/30 | train_f1=0.9459 | val_f1=0.5976 | wait=2
epoch 13/30 | train_f1=0.9538 | val_f1=0.5901 | wait=3
epoch 14/30 | train_f1=0.9612 | val_f1=0.5877 | wait=4
Early stop at epoch 14

--- Trial: image_all_metadata | encoded256 ---
epoch 1/30 | train_f1=0.5095 | val_f1=0.5221 | wait=0
epoch 2/30 | train_f1=0.6986 | val_f1=0.5781 | wait=0
epoch 3/30 | train_f1=0.7809 | val_f1=0.5984 | wait=0
epoch 4/30 | train_f1=0.8259 | val_f1=0.5570 | wait=0
epoch 5/30 | train_f1=0.8515 | val_f1=0.5875 | wait=1
epoch 6/30 | train_f1=0.8778 | val_f1=0.5779 | wait=2
epoch 7/30 | train_f1=0.8961 | val_f1=0.5739 | wait=3
epoch 8/30 | train_f1=0.9076 | val_f1=0.5782 | wait=4
Early stop at epoch 8

image_age -> encoded256¶

image_sex -> encoded32¶

image_location -> encoded64¶

image_age_sex -> encoded32¶

image_age_location -> raw¶

image_sex_location -> encoded128¶

image_all_metadata -> encoded32¶

In [26]:
# Cell X1: 只给 image_age 做更高维补充实验
extra_age_trials = [
    {"mode": "encoded", "embed_dim": 512},
    {"mode": "encoded", "embed_dim": 1024},
]

extra_age_trials
Out[26]:
[{'mode': 'encoded', 'embed_dim': 512}, {'mode': 'encoded', 'embed_dim': 1024}]
In [27]:
# Cell X2: 运行 image_age 的额外维度实验
extra_age_results = []

for setting in extra_age_trials:
    print(f"\n--- Extra Trial: image_age | encoded{setting['embed_dim']} ---")

    result = run_single_trial(
        experiment_name="image_age",
        selected_features=["age"],
        mode=setting["mode"],
        embed_dim=setting["embed_dim"],
        max_epochs=30,
        patience=5
    )

    extra_age_results.append(result)
--- Extra Trial: image_age | encoded512 ---
epoch 1/30 | train_f1=0.4113 | val_f1=0.5001 | wait=0
epoch 2/30 | train_f1=0.6570 | val_f1=0.5224 | wait=0
epoch 3/30 | train_f1=0.7394 | val_f1=0.4938 | wait=0
epoch 4/30 | train_f1=0.7902 | val_f1=0.5214 | wait=1
epoch 5/30 | train_f1=0.8148 | val_f1=0.5468 | wait=2
epoch 6/30 | train_f1=0.8559 | val_f1=0.5726 | wait=0
epoch 7/30 | train_f1=0.8723 | val_f1=0.5747 | wait=0
epoch 8/30 | train_f1=0.8773 | val_f1=0.5533 | wait=0
epoch 9/30 | train_f1=0.8944 | val_f1=0.5589 | wait=1
epoch 10/30 | train_f1=0.9055 | val_f1=0.5661 | wait=2
epoch 11/30 | train_f1=0.9072 | val_f1=0.5523 | wait=3
epoch 12/30 | train_f1=0.9272 | val_f1=0.5800 | wait=4
epoch 13/30 | train_f1=0.9126 | val_f1=0.5810 | wait=0
epoch 14/30 | train_f1=0.9450 | val_f1=0.5808 | wait=0
epoch 15/30 | train_f1=0.9436 | val_f1=0.5688 | wait=1
epoch 16/30 | train_f1=0.9486 | val_f1=0.5661 | wait=2
epoch 17/30 | train_f1=0.9455 | val_f1=0.5819 | wait=3
epoch 18/30 | train_f1=0.9566 | val_f1=0.5788 | wait=0
epoch 19/30 | train_f1=0.9541 | val_f1=0.5875 | wait=1
epoch 20/30 | train_f1=0.9586 | val_f1=0.5768 | wait=0
epoch 21/30 | train_f1=0.9667 | val_f1=0.5792 | wait=1
epoch 22/30 | train_f1=0.9729 | val_f1=0.5905 | wait=2
epoch 23/30 | train_f1=0.9664 | val_f1=0.5802 | wait=0
epoch 24/30 | train_f1=0.9756 | val_f1=0.5913 | wait=1
epoch 25/30 | train_f1=0.9651 | val_f1=0.5760 | wait=0
epoch 26/30 | train_f1=0.9759 | val_f1=0.5484 | wait=1
epoch 27/30 | train_f1=0.9745 | val_f1=0.5700 | wait=2
epoch 28/30 | train_f1=0.9765 | val_f1=0.5742 | wait=3
epoch 29/30 | train_f1=0.9807 | val_f1=0.5897 | wait=4
Early stop at epoch 29

--- Extra Trial: image_age | encoded1024 ---
epoch 1/30 | train_f1=0.3959 | val_f1=0.4322 | wait=0
epoch 2/30 | train_f1=0.6052 | val_f1=0.4816 | wait=0
epoch 3/30 | train_f1=0.6882 | val_f1=0.4969 | wait=0
epoch 4/30 | train_f1=0.7796 | val_f1=0.5490 | wait=0
epoch 5/30 | train_f1=0.7778 | val_f1=0.5309 | wait=0
epoch 6/30 | train_f1=0.8278 | val_f1=0.5524 | wait=1
epoch 7/30 | train_f1=0.8365 | val_f1=0.5553 | wait=0
epoch 8/30 | train_f1=0.8491 | val_f1=0.5656 | wait=0
epoch 9/30 | train_f1=0.8801 | val_f1=0.5716 | wait=0
epoch 10/30 | train_f1=0.8824 | val_f1=0.5317 | wait=0
epoch 11/30 | train_f1=0.8882 | val_f1=0.5692 | wait=1
epoch 12/30 | train_f1=0.8841 | val_f1=0.5777 | wait=2
epoch 13/30 | train_f1=0.9102 | val_f1=0.5746 | wait=0
epoch 14/30 | train_f1=0.9230 | val_f1=0.5801 | wait=1
epoch 15/30 | train_f1=0.9179 | val_f1=0.5904 | wait=0
epoch 16/30 | train_f1=0.9318 | val_f1=0.5806 | wait=0
epoch 17/30 | train_f1=0.9277 | val_f1=0.5880 | wait=1
epoch 18/30 | train_f1=0.9375 | val_f1=0.5765 | wait=2
epoch 19/30 | train_f1=0.9328 | val_f1=0.5317 | wait=3
epoch 20/30 | train_f1=0.9411 | val_f1=0.5625 | wait=4
Early stop at epoch 20
In [29]:
# Cell X3(推荐版)
image_age_old_df = pd.DataFrame([
    {
        "mode": r["mode"],
        "metadata_embed_dim": r["metadata_embed_dim"],
        "best_val_macro_f1": r["best_val_macro_f1"],
        "test_accuracy": r["test_accuracy"],
        "test_balanced_accuracy": r["test_balanced_accuracy"],
        "test_macro_f1": r["test_macro_f1"],
    }
    for r in all_trial_results
    if r["method"] == "image_age"
])

image_age_extra_df = pd.DataFrame([
    {
        "mode": r["mode"],
        "metadata_embed_dim": r["metadata_embed_dim"],
        "best_val_macro_f1": r["best_val_macro_f1"],
        "test_accuracy": r["test_accuracy"],
        "test_balanced_accuracy": r["test_balanced_accuracy"],
        "test_macro_f1": r["test_macro_f1"],
    }
    for r in extra_age_results
])

image_age_compare_df = pd.concat(
    [image_age_old_df, image_age_extra_df],
    ignore_index=True
).sort_values(["mode", "metadata_embed_dim"], na_position="first")

image_age_compare_df
Out[29]:
mode metadata_embed_dim best_val_macro_f1 test_accuracy test_balanced_accuracy test_macro_f1
1 encoded 32.0 0.592539 0.776502 0.585846 0.592534
2 encoded 64.0 0.592251 0.778528 0.633168 0.619543
3 encoded 128.0 0.594872 0.789332 0.587897 0.597060
4 encoded 256.0 0.602405 0.788656 0.598717 0.590842
5 encoded 512.0 0.591276 0.787306 0.571431 0.582345
6 encoded 1024.0 0.590421 0.792708 0.618421 0.603831
0 raw NaN 0.578138 0.726536 0.602865 0.561773

选256,不能选1024,不然又是数据泄露¶

In [30]:
# Cell X4: 保存 image_age 扩展维度结果
extra_age_save_path = SUPPORT_DIR / f"{date_tag}_{time_tag}_image_age_extra_dims.csv"
image_age_compare_df.to_csv(extra_age_save_path, index=False)
print(extra_age_save_path)
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-03_194704_image_age_extra_dims.csv
In [ ]: