class
ContrastiveTrainer
PythonContrastiveTrainer( encoder: keras.Model, projector: keras.Model | tuple[keras.Model, keras.Model], augmenter: keras.Layer | tuple[keras.Layer, keras.Layer] | None = None, probe: keras.Layer | keras.Model | None = None,)Creates a self-supervised contrastive trainer for a model.
Parameters
| Name | Type | Default | Description |
|---|---|---|---|
encoder | keras.Model | Required | The encoder model to be trained. |
projector | keras.Model | tuple[keras.Model, keras.Model] | Required | The projector model to be trained. |
augmenter | keras.Layer | tuple[keras.Layer, keras.Layer] | None | None | The augmenter to be used for data augmentation. |
probe | keras.Layer | keras.Model | None | None | The probe model to be trained. If None, no probe is used. |
constant
SAMPLES
PythonSAMPLES = 'data'constant
LABELS
PythonLABELS = 'labels'constant
AUG_SAMPLES_0
PythonAUG_SAMPLES_0 = 'augmented_data_0'constant
AUG_SAMPLES_1
PythonAUG_SAMPLES_1 = 'augmented_data_1'attribute
augmenters
Pythonaugmenters: tuple[keras.Layer, keras.Layer]attribute
encoder
Pythonencoder: keras.Model = encoderattribute
projectors
Pythonprojectors = projector if else (projector, projector)attribute
probe
Pythonprobe = probeattribute
loss_metric
Pythonloss_metric = keras.metrics.Mean(name='loss')attribute
encoder_metrics
Pythonencoder_metrics = []attribute
probe_loss_metric
Pythonprobe_loss_metric = keras.metrics.Mean(name='probe_loss')attribute
probe_metrics
Pythonprobe_metrics = []attribute
metrics
Pythonmetricsmethod
compile
Pythoncompile( encoder_optimizer: keras.Optimizer, encoder_loss: keras.Loss, encoder_metrics: list[keras.Metric] | None = None, probe_optimizer: keras.Optimizer | None = None, probe_loss: keras.Loss | None = None, probe_metrics: list[keras.Metric] | None = None, **kwargs={},)method
fit
Pythonfit(x=None, y=None, sample_weight=None, batch_size=None, validation_data=None, **kwargs={})method
run_augmenters
Pythonrun_augmenters(x, y=None)method
train_step
Pythontrain_step(data)method
test_step
Pythontest_step(data)method
call
Pythoncall(inputs)method
linear_probe
Pythonlinear_probe(num_classes, **kwargs={})staticmethod
method
save
PythonWe only save the encoder model
save(filepath, overwrite=True, zipped=True, **kwargs={})We only save the encoder model