-
Notifications
You must be signed in to change notification settings - Fork 4
Expand file tree
/
Copy pathtrainer.py
More file actions
87 lines (69 loc) · 2.81 KB
/
Copy pathtrainer.py
File metadata and controls
87 lines (69 loc) · 2.81 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
"""Train the Region Proposal Network on Pascal VOC data."""
from __future__ import annotations
import tensorflow as tf
from tensorflow.keras.callbacks import ModelCheckpoint
from utils import bbox_utils, data_utils, io_utils, train_utils
def main() -> None:
"""Run RPN training from the command line.
Returns:
None: The model is trained and checkpoints are written to disk.
"""
args = io_utils.handle_args()
if args.handle_gpu:
io_utils.handle_gpu_compatibility()
batch_size = 8
epochs = 50
load_weights = False
with_voc_2012 = False
backbone = args.backbone
get_model = io_utils.get_rpn_model_builder(backbone)
hyper_params = train_utils.get_hyper_params(backbone)
train_data, dataset_info = data_utils.get_dataset("voc/2007", "train+validation")
val_data, _ = data_utils.get_dataset("voc/2007", "test")
train_total_items = data_utils.get_total_item_size(dataset_info, "train+validation")
val_total_items = data_utils.get_total_item_size(dataset_info, "test")
if with_voc_2012:
voc_2012_data, voc_2012_info = data_utils.get_dataset("voc/2012", "train+validation")
voc_2012_total_items = data_utils.get_total_item_size(voc_2012_info, "train+validation")
train_total_items += voc_2012_total_items
train_data = train_data.concatenate(voc_2012_data)
labels = data_utils.get_labels(dataset_info)
hyper_params["total_labels"] = len(labels) + 1
img_size = hyper_params["img_size"]
train_data = data_utils.build_dataset(
train_data,
img_size,
img_size,
batch_size,
apply_augmentation=True
)
val_data = data_utils.build_dataset(val_data, img_size, img_size, batch_size)
anchors = bbox_utils.generate_anchors(hyper_params)
rpn_train_feed = train_utils.rpn_generator(train_data, anchors, hyper_params)
rpn_val_feed = train_utils.rpn_generator(val_data, anchors, hyper_params)
rpn_model, _ = get_model(hyper_params)
rpn_model.compile(
optimizer=tf.optimizers.Adam(learning_rate=1e-5),
loss=[train_utils.reg_loss, train_utils.cls_loss]
)
rpn_model_path = io_utils.get_model_path("rpn", backbone)
if load_weights:
rpn_model.load_weights(rpn_model_path)
checkpoint_callback = ModelCheckpoint(
rpn_model_path,
monitor="val_loss",
save_best_only=True,
save_weights_only=True
)
step_size_train = train_utils.get_step_size(train_total_items, batch_size)
step_size_val = train_utils.get_step_size(val_total_items, batch_size)
rpn_model.fit(
rpn_train_feed,
steps_per_epoch=step_size_train,
validation_data=rpn_val_feed,
validation_steps=step_size_val,
epochs=epochs,
callbacks=[checkpoint_callback]
)
if __name__ == "__main__":
main()