neurovlm.training.MLPBrainToTextRetrievalTrainConfig#
- class neurovlm.training.MLPBrainToTextRetrievalTrainConfig(output_root='runs', run_id=None, seed=42, device='auto', epochs=100, batch_size=256, eval_batch_size=None, num_workers=0, learning_rate=5e-05, weight_decay=0.0, gradient_clip=None, early_stopping_patience=None, early_stopping_min_delta=0.0, max_train_batches=None, max_eval_batches=None, resume=None, preset='retained', text_dim=768, text_hidden_dim=512, brain_dim=384, brain_hidden_dim=384, shared_dim=384, temperature=0.07, initialize_text_from_mse=True, primary_metric='val_i2t_normalized_k_recall_curve_auc', metric_direction=MetricDirection.MAX)[source]#
Train the retained contrastive heads for brain-to-text retrieval.
The optimization objective remains the symmetric retained InfoNCE loss. Only checkpoint selection changes: this task selects the image/brain-to-text (
i2t) full recall-curve AUC instead of validation loss.- Parameters:
output_root (str | Path)
run_id (str | None)
seed (int)
device (str)
epochs (int)
batch_size (int)
eval_batch_size (int | None)
num_workers (int)
learning_rate (float)
weight_decay (float)
gradient_clip (float | None)
early_stopping_patience (int | None)
early_stopping_min_delta (float)
max_train_batches (int | None)
max_eval_batches (int | None)
resume (str | Path | None)
preset (str)
text_dim (int)
text_hidden_dim (int)
brain_dim (int)
brain_hidden_dim (int)
shared_dim (int)
temperature (float)
initialize_text_from_mse (bool)
primary_metric (str)
metric_direction (MetricDirection)
- __init__(output_root='runs', run_id=None, seed=42, device='auto', epochs=100, batch_size=256, eval_batch_size=None, num_workers=0, learning_rate=5e-05, weight_decay=0.0, gradient_clip=None, early_stopping_patience=None, early_stopping_min_delta=0.0, max_train_batches=None, max_eval_batches=None, resume=None, preset='retained', text_dim=768, text_hidden_dim=512, brain_dim=384, brain_hidden_dim=384, shared_dim=384, temperature=0.07, initialize_text_from_mse=True, primary_metric='val_i2t_normalized_k_recall_curve_auc', metric_direction=MetricDirection.MAX)#
- Parameters:
output_root (str | Path)
run_id (str | None)
seed (int)
device (str)
epochs (int)
batch_size (int)
eval_batch_size (int | None)
num_workers (int)
learning_rate (float)
weight_decay (float)
gradient_clip (float | None)
early_stopping_patience (int | None)
early_stopping_min_delta (float)
max_train_batches (int | None)
max_eval_batches (int | None)
resume (str | Path | None)
preset (str)
text_dim (int)
text_hidden_dim (int)
brain_dim (int)
brain_hidden_dim (int)
shared_dim (int)
temperature (float)
initialize_text_from_mse (bool)
primary_metric (str)
metric_direction (MetricDirection)
- Return type:
None
Methods
__init__([output_root, run_id, seed, ...])architecture()Attributes
batch_sizebrain_dimbrain_hidden_dimdeviceearly_stopping_min_deltaearly_stopping_patienceepochseval_batch_sizegradient_clipinitialize_text_from_mselearning_ratemax_eval_batchesmax_train_batchesmetric_directionnum_workersoutput_rootpresetprimary_metricresumerun_idseedshared_dimtemperaturetext_dimtext_hidden_dimweight_decay