neurovlm.training.MLPBrainToTextRetrievalTrainConfig

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_size

brain_dim

brain_hidden_dim

device

early_stopping_min_delta

early_stopping_patience

epochs

eval_batch_size

gradient_clip

initialize_text_from_mse

learning_rate

max_eval_batches

max_train_batches

metric_direction

num_workers

output_root

preset

primary_metric

resume

run_id

seed

shared_dim

temperature

text_dim

text_hidden_dim

weight_decay