콘텐츠로 이동

Make Train Config

SaigeVAD.VadTrainConfigMaker

Bases: TrainConfigMaker

VAD 학습을 위해 필요한 config을 생성하는 VadTrainConfigMaker class입니다.

Usage
from SaigeVAD.engine.saigevad_api import VadTrainConfigMaker

# VadTrainConfigMaker 빌드
_, train_config_maker = VadTrainConfigMaker.build(config=config)

# VadTrainer를 위한 config 생성
_, config_dict = train_config_maker.make_train_config()

build(config) classmethod

학습을 위한 config을 생성하는 VadTrainConfigMaker class의 instance를 생성합니다.

Parameters:

Name Type Description Default
config Dict

학습 config가 담겨 있는 dictionary 입니다.

config = {
    "train_data": {
        "clip_path": "/NFS/database_personal/bts/VAD/voc/221207/vad_training_sample/clips",
        "frame_dir": [
            "test1_00",
            "test1_01",
            "test1_04",
            "test1_06",
        ],
    },
    "config": {
        "cycle_list": [
            "test1_00-00000000",  # {clip name}-{start frame cycle 1}
            "test1_00-00000027",  # {clip name}-{end frame cycle 1}
            "test1_00-00000035",  # {clip name}-{start frame cycle 2}
            "test1_00-00000059",  # {clip name}-{end frame cycle 2}
            "test1_04-00000000",  # {clip name}-{start frame cycle 3}
            "test1_04-00000029",  # {clip name}-{end frame cycle 3}
        ],
        "iteration": 5000,
        "network_model": "shallow", # str, "shallow" or "deep". Defaults to "shallow"
        "num_workers": 8, # int, Defaults to 0.
    },
    "metadata": {
        "common_settings": {
            "roi_xyxy": [1149, 557, 1278, 746],  # (left, top, right, bottom)
            "fixed_roi_xyxy": [1200, 557, 1458, 909], # (left, top, right, bottom): default to None
            "fps": 25,
            "mask_mode": None,  # "binary", "semi_binary
            "cycle_length": 29, # int
        },
    },
}

required

Returns:

Name Type Description
VadTrainConfigMaker

build가 완료된 saigevad_api.VadTrainConfigMaker class의 instance를 반환합니다.


make_train_config()

VadTrainer.build를 위해 필요한 config dict를 생성합니다. Motion detector를 위한 parameter searching도 진행합니다.

Returns:

Name Type Description
Dict Dict

VadTrainer의 학습 config가 담겨 있는 dictionary 입니다.

output = {
    "train_config": {
        "base_config": base_config,
        "train_dataloader": train_dataloader,
        "motion_detector_dataloader": motion_detector_dataloader,
    },
    "reference_images": {
        "template_image_tensors": template_image_tensors,
        "fixed_template_image_tensors": fixed_template_image_tensors,
        "color_count_map": color_count_map,
    },
    "pretrained_model_path": os.path.join(ROOT_DIR, "../pretrained_weight/vad"),
}