segment-everything 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- segment_everything/__init__.py +5 -0
- segment_everything/augmentation/albumentations_helper.py +0 -0
- segment_everything/detect_and_segment.py +131 -0
- segment_everything/napari_helper.py +15 -0
- segment_everything/prompt_generator.py +188 -0
- segment_everything/py.typed +5 -0
- segment_everything/stacked_label_dataset.py +113 -0
- segment_everything/stacked_labels.py +428 -0
- segment_everything/vendored/PromptGuidedDecoder/Prompt_guided_Mask_Decoder.pt +0 -0
- segment_everything/vendored/__init__.py +5 -0
- segment_everything/vendored/dice.py +158 -0
- segment_everything/vendored/efficientvit/__init__.py +0 -0
- segment_everything/vendored/efficientvit/apps/__init__.py +0 -0
- segment_everything/vendored/efficientvit/apps/data_provider/__init__.py +7 -0
- segment_everything/vendored/efficientvit/apps/data_provider/augment/__init__.py +6 -0
- segment_everything/vendored/efficientvit/apps/data_provider/augment/bbox.py +30 -0
- segment_everything/vendored/efficientvit/apps/data_provider/augment/color_aug.py +78 -0
- segment_everything/vendored/efficientvit/apps/data_provider/base.py +254 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/__init__.py +6 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/_data_loader.py +1538 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/_data_worker.py +357 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/controller.py +100 -0
- segment_everything/vendored/efficientvit/apps/setup.py +150 -0
- segment_everything/vendored/efficientvit/apps/trainer/__init__.py +6 -0
- segment_everything/vendored/efficientvit/apps/trainer/base.py +318 -0
- segment_everything/vendored/efficientvit/apps/trainer/run_config.py +129 -0
- segment_everything/vendored/efficientvit/apps/utils/__init__.py +12 -0
- segment_everything/vendored/efficientvit/apps/utils/dist.py +32 -0
- segment_everything/vendored/efficientvit/apps/utils/ema.py +52 -0
- segment_everything/vendored/efficientvit/apps/utils/export.py +45 -0
- segment_everything/vendored/efficientvit/apps/utils/init.py +66 -0
- segment_everything/vendored/efficientvit/apps/utils/lr.py +52 -0
- segment_everything/vendored/efficientvit/apps/utils/metric.py +43 -0
- segment_everything/vendored/efficientvit/apps/utils/misc.py +101 -0
- segment_everything/vendored/efficientvit/apps/utils/opt.py +28 -0
- segment_everything/vendored/efficientvit/cls_model_zoo.py +79 -0
- segment_everything/vendored/efficientvit/clscore/__init__.py +0 -0
- segment_everything/vendored/efficientvit/clscore/data_provider/__init__.py +5 -0
- segment_everything/vendored/efficientvit/clscore/data_provider/imagenet.py +142 -0
- segment_everything/vendored/efficientvit/clscore/trainer/__init__.py +6 -0
- segment_everything/vendored/efficientvit/clscore/trainer/cls_run_config.py +18 -0
- segment_everything/vendored/efficientvit/clscore/trainer/cls_trainer.py +265 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/__init__.py +7 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/label_smooth.py +18 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/metric.py +23 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/mixup.py +67 -0
- segment_everything/vendored/efficientvit/models/__init__.py +0 -0
- segment_everything/vendored/efficientvit/models/efficientvit/__init__.py +8 -0
- segment_everything/vendored/efficientvit/models/efficientvit/backbone.py +380 -0
- segment_everything/vendored/efficientvit/models/efficientvit/cls.py +188 -0
- segment_everything/vendored/efficientvit/models/efficientvit/sam.py +181 -0
- segment_everything/vendored/efficientvit/models/efficientvit/seg.py +373 -0
- segment_everything/vendored/efficientvit/models/nn/__init__.py +8 -0
- segment_everything/vendored/efficientvit/models/nn/act.py +30 -0
- segment_everything/vendored/efficientvit/models/nn/drop.py +104 -0
- segment_everything/vendored/efficientvit/models/nn/norm.py +164 -0
- segment_everything/vendored/efficientvit/models/nn/ops.py +597 -0
- segment_everything/vendored/efficientvit/models/utils/__init__.py +7 -0
- segment_everything/vendored/efficientvit/models/utils/list.py +53 -0
- segment_everything/vendored/efficientvit/models/utils/network.py +73 -0
- segment_everything/vendored/efficientvit/models/utils/random.py +65 -0
- segment_everything/vendored/efficientvit/sam_model_zoo.py +45 -0
- segment_everything/vendored/efficientvit/seg_model_zoo.py +70 -0
- segment_everything/vendored/get_object_aware.py +26 -0
- segment_everything/vendored/mobilesamv2/__init__.py +16 -0
- segment_everything/vendored/mobilesamv2/automatic_mask_generator.py +415 -0
- segment_everything/vendored/mobilesamv2/build_sam.py +246 -0
- segment_everything/vendored/mobilesamv2/modeling/__init__.py +11 -0
- segment_everything/vendored/mobilesamv2/modeling/common.py +43 -0
- segment_everything/vendored/mobilesamv2/modeling/image_encoder.py +394 -0
- segment_everything/vendored/mobilesamv2/modeling/mask_decoder.py +213 -0
- segment_everything/vendored/mobilesamv2/modeling/prompt_encoder.py +217 -0
- segment_everything/vendored/mobilesamv2/modeling/sam.py +203 -0
- segment_everything/vendored/mobilesamv2/modeling/transformer.py +240 -0
- segment_everything/vendored/mobilesamv2/predictor.py +384 -0
- segment_everything/vendored/mobilesamv2/utils/__init__.py +5 -0
- segment_everything/vendored/mobilesamv2/utils/amg.py +347 -0
- segment_everything/vendored/mobilesamv2/utils/onnx.py +144 -0
- segment_everything/vendored/mobilesamv2/utils/transforms.py +103 -0
- segment_everything/vendored/object_detection/__init__.py +0 -0
- segment_everything/vendored/object_detection/ultralytics/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/nn/__init__.py +9 -0
- segment_everything/vendored/object_detection/ultralytics/nn/autobackend.py +658 -0
- segment_everything/vendored/object_detection/ultralytics/nn/autoshape.py +397 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/__init__.py +110 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/block.py +304 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/conv.py +297 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/head.py +468 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/transformer.py +378 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/utils.py +78 -0
- segment_everything/vendored/object_detection/ultralytics/nn/tasks.py +1049 -0
- segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/__init__.py +6 -0
- segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/model.py +104 -0
- segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/predict.py +95 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/cfg/__init__.py +588 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/cfg/default.yaml +117 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/__init__.py +9 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/annotator.py +53 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/augment.py +899 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/base.py +286 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/build.py +213 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/converter.py +358 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataloaders/__init__.py +0 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataloaders/stream_loaders.py +459 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataset.py +274 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataset_wrappers.py +53 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/download_weights.sh +18 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_coco.sh +60 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_coco128.sh +17 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_imagenet.sh +51 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/utils.py +716 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/__init__.py +0 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/exporter.py +1214 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/model.py +641 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/predictor.py +461 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/results.py +741 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/__init__.py +893 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/autobatch.py +108 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/callbacks/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/callbacks/base.py +212 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/checks.py +547 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/dist.py +67 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/downloads.py +353 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/errors.py +12 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/files.py +100 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/instance.py +391 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/loss.py +579 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/metrics.py +1189 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/ops.py +870 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/patches.py +45 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/plotting.py +767 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/tal.py +276 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/torch_utils.py +684 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/tuner.py +54 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/v8/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/v8/detect/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/v8/detect/predict.py +69 -0
- segment_everything/vendored/tinyvit/__init__.py +2 -0
- segment_everything/vendored/tinyvit/tiny_vit.py +867 -0
- segment_everything/weights_helper.py +124 -0
- segment_everything-0.1.0.dist-info/METADATA +53 -0
- segment_everything-0.1.0.dist-info/RECORD +145 -0
- segment_everything-0.1.0.dist-info/WHEEL +4 -0
- segment_everything-0.1.0.dist-info/licenses/LICENSE +28 -0
|
@@ -0,0 +1,108 @@
|
|
|
1
|
+
# Ultralytics YOLO 🚀, AGPL-3.0 license
|
|
2
|
+
"""
|
|
3
|
+
Functions for estimating the best YOLO batch size to use a fraction of the available CUDA memory in PyTorch.
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
from copy import deepcopy
|
|
7
|
+
|
|
8
|
+
import numpy as np
|
|
9
|
+
import torch
|
|
10
|
+
|
|
11
|
+
from ..utils import DEFAULT_CFG, LOGGER, colorstr
|
|
12
|
+
from ..utils.torch_utils import profile
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def check_train_batch_size(model, imgsz=640, amp=True):
|
|
16
|
+
"""
|
|
17
|
+
Check YOLO training batch size using the autobatch() function.
|
|
18
|
+
|
|
19
|
+
Args:
|
|
20
|
+
model (torch.nn.Module): YOLO model to check batch size for.
|
|
21
|
+
imgsz (int): Image size used for training.
|
|
22
|
+
amp (bool): If True, use automatic mixed precision (AMP) for training.
|
|
23
|
+
|
|
24
|
+
Returns:
|
|
25
|
+
(int): Optimal batch size computed using the autobatch() function.
|
|
26
|
+
"""
|
|
27
|
+
|
|
28
|
+
with torch.cuda.amp.autocast(amp):
|
|
29
|
+
return autobatch(
|
|
30
|
+
deepcopy(model).train(), imgsz
|
|
31
|
+
) # compute optimal batch size
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def autobatch(model, imgsz=640, fraction=0.67, batch_size=DEFAULT_CFG.batch):
|
|
35
|
+
"""
|
|
36
|
+
Automatically estimate the best YOLO batch size to use a fraction of the available CUDA memory.
|
|
37
|
+
|
|
38
|
+
Args:
|
|
39
|
+
model (torch.nn.module): YOLO model to compute batch size for.
|
|
40
|
+
imgsz (int, optional): The image size used as input for the YOLO model. Defaults to 640.
|
|
41
|
+
fraction (float, optional): The fraction of available CUDA memory to use. Defaults to 0.67.
|
|
42
|
+
batch_size (int, optional): The default batch size to use if an error is detected. Defaults to 16.
|
|
43
|
+
|
|
44
|
+
Returns:
|
|
45
|
+
(int): The optimal batch size.
|
|
46
|
+
"""
|
|
47
|
+
|
|
48
|
+
# Check device
|
|
49
|
+
prefix = colorstr("AutoBatch: ")
|
|
50
|
+
LOGGER.info(f"{prefix}Computing optimal batch size for imgsz={imgsz}")
|
|
51
|
+
device = next(model.parameters()).device # get model device
|
|
52
|
+
if device.type == "cpu":
|
|
53
|
+
LOGGER.info(
|
|
54
|
+
f"{prefix}CUDA not detected, using default CPU batch-size {batch_size}"
|
|
55
|
+
)
|
|
56
|
+
return batch_size
|
|
57
|
+
if torch.backends.cudnn.benchmark:
|
|
58
|
+
LOGGER.info(
|
|
59
|
+
f"{prefix} ⚠️ Requires torch.backends.cudnn.benchmark=False, using default batch-size {batch_size}"
|
|
60
|
+
)
|
|
61
|
+
return batch_size
|
|
62
|
+
|
|
63
|
+
# Inspect CUDA memory
|
|
64
|
+
gb = 1 << 30 # bytes to GiB (1024 ** 3)
|
|
65
|
+
d = str(device).upper() # 'CUDA:0'
|
|
66
|
+
properties = torch.cuda.get_device_properties(device) # device properties
|
|
67
|
+
t = properties.total_memory / gb # GiB total
|
|
68
|
+
r = torch.cuda.memory_reserved(device) / gb # GiB reserved
|
|
69
|
+
a = torch.cuda.memory_allocated(device) / gb # GiB allocated
|
|
70
|
+
f = t - (r + a) # GiB free
|
|
71
|
+
LOGGER.info(
|
|
72
|
+
f"{prefix}{d} ({properties.name}) {t:.2f}G total, {r:.2f}G reserved, {a:.2f}G allocated, {f:.2f}G free"
|
|
73
|
+
)
|
|
74
|
+
|
|
75
|
+
# Profile batch sizes
|
|
76
|
+
batch_sizes = [1, 2, 4, 8, 16]
|
|
77
|
+
try:
|
|
78
|
+
img = [torch.empty(b, 3, imgsz, imgsz) for b in batch_sizes]
|
|
79
|
+
results = profile(img, model, n=3, device=device)
|
|
80
|
+
|
|
81
|
+
# Fit a solution
|
|
82
|
+
y = [x[2] for x in results if x] # memory [2]
|
|
83
|
+
p = np.polyfit(
|
|
84
|
+
batch_sizes[: len(y)], y, deg=1
|
|
85
|
+
) # first degree polynomial fit
|
|
86
|
+
b = int(
|
|
87
|
+
(f * fraction - p[1]) / p[0]
|
|
88
|
+
) # y intercept (optimal batch size)
|
|
89
|
+
if None in results: # some sizes failed
|
|
90
|
+
i = results.index(None) # first fail index
|
|
91
|
+
if b >= batch_sizes[i]: # y intercept above failure point
|
|
92
|
+
b = batch_sizes[max(i - 1, 0)] # select prior safe point
|
|
93
|
+
if b < 1 or b > 1024: # b outside of safe range
|
|
94
|
+
b = batch_size
|
|
95
|
+
LOGGER.info(
|
|
96
|
+
f"{prefix}WARNING ⚠️ CUDA anomaly detected, using default batch-size {batch_size}."
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
fraction = (np.polyval(p, b) + r + a) / t # actual fraction predicted
|
|
100
|
+
LOGGER.info(
|
|
101
|
+
f"{prefix}Using batch-size {b} for {d} {t * fraction:.2f}G/{t:.2f}G ({fraction * 100:.0f}%) ✅"
|
|
102
|
+
)
|
|
103
|
+
return b
|
|
104
|
+
except Exception as e:
|
|
105
|
+
LOGGER.warning(
|
|
106
|
+
f"{prefix}WARNING ⚠️ error detected: {e}, using default batch-size {batch_size}."
|
|
107
|
+
)
|
|
108
|
+
return batch_size
|
|
@@ -0,0 +1,212 @@
|
|
|
1
|
+
# Ultralytics YOLO 🚀, AGPL-3.0 license
|
|
2
|
+
"""
|
|
3
|
+
Base callbacks
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
from collections import defaultdict
|
|
7
|
+
from copy import deepcopy
|
|
8
|
+
|
|
9
|
+
# Trainer callbacks ----------------------------------------------------------------------------------------------------
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def on_pretrain_routine_start(trainer):
|
|
13
|
+
"""Called before the pretraining routine starts."""
|
|
14
|
+
pass
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def on_pretrain_routine_end(trainer):
|
|
18
|
+
"""Called after the pretraining routine ends."""
|
|
19
|
+
pass
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def on_train_start(trainer):
|
|
23
|
+
"""Called when the training starts."""
|
|
24
|
+
pass
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def on_train_epoch_start(trainer):
|
|
28
|
+
"""Called at the start of each training epoch."""
|
|
29
|
+
pass
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def on_train_batch_start(trainer):
|
|
33
|
+
"""Called at the start of each training batch."""
|
|
34
|
+
pass
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def optimizer_step(trainer):
|
|
38
|
+
"""Called when the optimizer takes a step."""
|
|
39
|
+
pass
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def on_before_zero_grad(trainer):
|
|
43
|
+
"""Called before the gradients are set to zero."""
|
|
44
|
+
pass
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def on_train_batch_end(trainer):
|
|
48
|
+
"""Called at the end of each training batch."""
|
|
49
|
+
pass
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def on_train_epoch_end(trainer):
|
|
53
|
+
"""Called at the end of each training epoch."""
|
|
54
|
+
pass
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def on_fit_epoch_end(trainer):
|
|
58
|
+
"""Called at the end of each fit epoch (train + val)."""
|
|
59
|
+
pass
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def on_model_save(trainer):
|
|
63
|
+
"""Called when the model is saved."""
|
|
64
|
+
pass
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def on_train_end(trainer):
|
|
68
|
+
"""Called when the training ends."""
|
|
69
|
+
pass
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def on_params_update(trainer):
|
|
73
|
+
"""Called when the model parameters are updated."""
|
|
74
|
+
pass
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def teardown(trainer):
|
|
78
|
+
"""Called during the teardown of the training process."""
|
|
79
|
+
pass
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
# Validator callbacks --------------------------------------------------------------------------------------------------
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
def on_val_start(validator):
|
|
86
|
+
"""Called when the validation starts."""
|
|
87
|
+
pass
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
def on_val_batch_start(validator):
|
|
91
|
+
"""Called at the start of each validation batch."""
|
|
92
|
+
pass
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
def on_val_batch_end(validator):
|
|
96
|
+
"""Called at the end of each validation batch."""
|
|
97
|
+
pass
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def on_val_end(validator):
|
|
101
|
+
"""Called when the validation ends."""
|
|
102
|
+
pass
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
# Predictor callbacks --------------------------------------------------------------------------------------------------
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def on_predict_start(predictor):
|
|
109
|
+
"""Called when the prediction starts."""
|
|
110
|
+
pass
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
def on_predict_batch_start(predictor):
|
|
114
|
+
"""Called at the start of each prediction batch."""
|
|
115
|
+
pass
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
def on_predict_batch_end(predictor):
|
|
119
|
+
"""Called at the end of each prediction batch."""
|
|
120
|
+
pass
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
def on_predict_postprocess_end(predictor):
|
|
124
|
+
"""Called after the post-processing of the prediction ends."""
|
|
125
|
+
pass
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
def on_predict_end(predictor):
|
|
129
|
+
"""Called when the prediction ends."""
|
|
130
|
+
pass
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
# Exporter callbacks ---------------------------------------------------------------------------------------------------
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
def on_export_start(exporter):
|
|
137
|
+
"""Called when the model export starts."""
|
|
138
|
+
pass
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def on_export_end(exporter):
|
|
142
|
+
"""Called when the model export ends."""
|
|
143
|
+
pass
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
default_callbacks = {
|
|
147
|
+
# Run in trainer
|
|
148
|
+
'on_pretrain_routine_start': [on_pretrain_routine_start],
|
|
149
|
+
'on_pretrain_routine_end': [on_pretrain_routine_end],
|
|
150
|
+
'on_train_start': [on_train_start],
|
|
151
|
+
'on_train_epoch_start': [on_train_epoch_start],
|
|
152
|
+
'on_train_batch_start': [on_train_batch_start],
|
|
153
|
+
'optimizer_step': [optimizer_step],
|
|
154
|
+
'on_before_zero_grad': [on_before_zero_grad],
|
|
155
|
+
'on_train_batch_end': [on_train_batch_end],
|
|
156
|
+
'on_train_epoch_end': [on_train_epoch_end],
|
|
157
|
+
'on_fit_epoch_end': [on_fit_epoch_end], # fit = train + val
|
|
158
|
+
'on_model_save': [on_model_save],
|
|
159
|
+
'on_train_end': [on_train_end],
|
|
160
|
+
'on_params_update': [on_params_update],
|
|
161
|
+
'teardown': [teardown],
|
|
162
|
+
|
|
163
|
+
# Run in validator
|
|
164
|
+
'on_val_start': [on_val_start],
|
|
165
|
+
'on_val_batch_start': [on_val_batch_start],
|
|
166
|
+
'on_val_batch_end': [on_val_batch_end],
|
|
167
|
+
'on_val_end': [on_val_end],
|
|
168
|
+
|
|
169
|
+
# Run in predictor
|
|
170
|
+
'on_predict_start': [on_predict_start],
|
|
171
|
+
'on_predict_batch_start': [on_predict_batch_start],
|
|
172
|
+
'on_predict_postprocess_end': [on_predict_postprocess_end],
|
|
173
|
+
'on_predict_batch_end': [on_predict_batch_end],
|
|
174
|
+
'on_predict_end': [on_predict_end],
|
|
175
|
+
|
|
176
|
+
# Run in exporter
|
|
177
|
+
'on_export_start': [on_export_start],
|
|
178
|
+
'on_export_end': [on_export_end]}
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
def get_default_callbacks():
|
|
182
|
+
"""
|
|
183
|
+
Return a copy of the default_callbacks dictionary with lists as default values.
|
|
184
|
+
|
|
185
|
+
Returns:
|
|
186
|
+
(defaultdict): A defaultdict with keys from default_callbacks and empty lists as default values.
|
|
187
|
+
"""
|
|
188
|
+
return defaultdict(list, deepcopy(default_callbacks))
|
|
189
|
+
|
|
190
|
+
|
|
191
|
+
def add_integration_callbacks(instance):
|
|
192
|
+
"""
|
|
193
|
+
Add integration callbacks from various sources to the instance's callbacks.
|
|
194
|
+
|
|
195
|
+
Args:
|
|
196
|
+
instance (Trainer, Predictor, Validator, Exporter): An object with a 'callbacks' attribute that is a dictionary
|
|
197
|
+
of callback lists.
|
|
198
|
+
"""
|
|
199
|
+
from .clearml import callbacks as clearml_cb
|
|
200
|
+
from .comet import callbacks as comet_cb
|
|
201
|
+
from .dvc import callbacks as dvc_cb
|
|
202
|
+
from .hub import callbacks as hub_cb
|
|
203
|
+
from .mlflow import callbacks as mlflow_cb
|
|
204
|
+
from .neptune import callbacks as neptune_cb
|
|
205
|
+
from .raytune import callbacks as tune_cb
|
|
206
|
+
from .tensorboard import callbacks as tensorboard_cb
|
|
207
|
+
from .wb import callbacks as wb_cb
|
|
208
|
+
|
|
209
|
+
for x in clearml_cb, comet_cb, hub_cb, mlflow_cb, neptune_cb, tune_cb, tensorboard_cb, wb_cb, dvc_cb:
|
|
210
|
+
for k, v in x.items():
|
|
211
|
+
if v not in instance.callbacks[k]: # prevent duplicate callbacks addition
|
|
212
|
+
instance.callbacks[k].append(v) # callback[name].append(func)
|