|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| """Configuration definitions for EfficientNet losses, learning rates, and optimizers."""
|
| from __future__ import absolute_import
|
| from __future__ import division
|
| from __future__ import print_function
|
| import dataclasses
|
| from official.legacy.image_classification.configs import base_configs
|
| from official.modeling.hyperparams import base_config
|
|
|
|
|
| @dataclasses.dataclass
|
| class EfficientNetModelConfig(base_configs.ModelConfig):
|
| """Configuration for the EfficientNet model.
|
|
|
| This configuration will default to settings used for training efficientnet-b0
|
| on a v3-8 TPU on ImageNet.
|
|
|
| Attributes:
|
| name: The name of the model. Defaults to 'EfficientNet'.
|
| num_classes: The number of classes in the model.
|
| model_params: A dictionary that represents the parameters of the
|
| EfficientNet model. These will be passed in to the "from_name" function.
|
| loss: The configuration for loss. Defaults to a categorical cross entropy
|
| implementation.
|
| optimizer: The configuration for optimizations. Defaults to an RMSProp
|
| configuration.
|
| learning_rate: The configuration for learning rate. Defaults to an
|
| exponential configuration.
|
| """
|
| name: str = 'EfficientNet'
|
| num_classes: int = 1000
|
| model_params: base_config.Config = dataclasses.field(
|
| default_factory=lambda: {
|
| 'model_name': 'efficientnet-b0',
|
| 'model_weights_path': '',
|
| 'weights_format': 'saved_model',
|
| 'overrides': {
|
| 'batch_norm': 'default',
|
| 'rescale_input': True,
|
| 'num_classes': 1000,
|
| 'activation': 'swish',
|
| 'dtype': 'float32',
|
| }
|
| })
|
| loss: base_configs.LossConfig = dataclasses.field(
|
| default_factory=lambda: base_configs.LossConfig(
|
| name='categorical_crossentropy', label_smoothing=0.1
|
| )
|
| )
|
| optimizer: base_configs.OptimizerConfig = dataclasses.field(
|
| default_factory=lambda: base_configs.OptimizerConfig(
|
| name='rmsprop',
|
| decay=0.9,
|
| epsilon=0.001,
|
| momentum=0.9,
|
| moving_average_decay=None,
|
| )
|
| )
|
| learning_rate: base_configs.LearningRateConfig = dataclasses.field(
|
| default_factory=lambda: base_configs.LearningRateConfig(
|
| name='exponential',
|
| initial_lr=0.008,
|
| decay_epochs=2.4,
|
| decay_rate=0.97,
|
| warmup_epochs=5,
|
| scale_by_batch_size=1.0 / 128.0,
|
| staircase=True,
|
| )
|
| )
|
|
|