Skip to content

Commit 2d4bb72

Browse files
No public description
PiperOrigin-RevId: 800720755
1 parent 52b93bf commit 2d4bb72

1 file changed

Lines changed: 31 additions & 1 deletion

File tree

  • official/projects/waste_identification_ml/fine_tuning/Detectron2 Mask RCNN

‎official/projects/waste_identification_ml/fine_tuning/Detectron2 Mask RCNN/README.md‎

Lines changed: 31 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -100,7 +100,7 @@ model zoo, ResNet-50 pretrained on ImageNet).
100100
- `ROI_HEADS.NUM_CLASSES`: Number of classes in your custom dataset (excluding background).
101101
- `BACKBONE.FREEZE_AT`: Freezes the initial layers up to this stage in the backbone. 0 means no layers are frozen (i.e., all layers are trainable).
102102
- `MAX_ITER`: Total number of training iterations.
103-
- `BASE_LR`: Base learning rate for training. Try, base_lr = 0.001 × (batch_size / 16).
103+
- `BASE_LR`: Base learning rate for training. Try, base_lr = (0.02 or 0.001) × (batch_size / 16).
104104
- `IMS_PER_BATCH`: Number of images per training batch (i.e., batch size).
105105
- `CHECKPOINT_PERIOD`: Save model checkpoints after this many iterations.
106106
- `WARMUP_ITERS`: Number of warmup iterations for learning rate scheduling.
@@ -118,6 +118,36 @@ Images larger than this will be resized down.
118118
- `OUTPUT_DIR`: Directory path where all model outputs
119119
(checkpoints, logs, predictions) will be saved.
120120
121+
Calculated the parameters using the formula below, but its subjective -
122+
123+
```python
124+
dataset_size = 347 # replace with your actual number
125+
IMS_PER_BATCH = 32 # total across all GPUs
126+
epochs = 300
127+
checkpoint_every_n_epochs = 50
128+
129+
130+
# Derived values
131+
iters_per_epoch = dataset_size / IMS_PER_BATCH
132+
MAX_ITER = int(iters_per_epoch * epochs)
133+
134+
STEP1 = int(MAX_ITER * 0.6)
135+
STEP2 = int(MAX_ITER * 0.8)
136+
STEP3 = int(MAX_ITER * 0.9)
137+
138+
WARMUP_ITERS = int(MAX_ITER * 0.05)
139+
BASE_LR = 0.001 * (IMS_PER_BATCH / 16)
140+
CHECKPOINT_PERIOD = int(checkpoint_every_n_epochs * iters_per_epoch)
141+
142+
print(f"MAX_ITER: {MAX_ITER}")
143+
print(f"WARMUP_ITERS: {WARMUP_ITERS}")
144+
print(f"BASE_LR: {BASE_LR}")
145+
print(f"CHECKPOINT_PERIOD: {CHECKPOINT_PERIOD}")
146+
print(f"STEP1: {STEP1}")
147+
print(f"STEP2: {STEP2}")
148+
print(f"STEP3: {STEP3}")
149+
```
150+
121151
```yaml
122152
_BASE_: "../Base-RCNN-FPN.yaml"
123153

0 commit comments

Comments
 (0)