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