SceneTextDetection
ocr.DetTrainer
Scene Text Detection 모델을 학습시키기 위한 Trainer class 입니다.
Usage
# Trainer 빌드
error, trainer = DetTrainer.build(CONFIG, DATA)
assert error >= 0
# 0번 GPU로 이동
error, result = trainer.to_device(DEVICE)
assert error >= 0
# validation 가능한지 확인
error, validation_enabled = trainer.validation_enabled()
assert error >= 0
# 학습 루프
for step in range(CONFIG["total_iterations"]):
error, result = trainer.train_one_step()
assert error >= 0
if (step + 1) % PRINT_INTERVAL == 0:
print(f"Train {step + 1}:", result)
if validation_enabled and (step + 1) % VALIDATION_INTERVAL == 0:
error, result = trainer.validate()
assert error >= 0
print(f"Validation {step + 1}:", result)
# 학습 후처리 (default inference options)
error, post_train_process_result = trainer.post_train_process()
assert error >= 0
print("post_train_process :", post_train_process_result)
# 체크포인트 저장
error, _ = trainer.save_checkpoint(
checkpoint_path=CHECKPOINT_PATH,
password=PASSWORD,
metadata={"description": "demo"},
)
build(config, data)
classmethod
DetTrainer class의 instance를 생성합니다.
Parameters:
Returns:
| Name | Type | Description |
|---|---|---|
DetTrainer | build가 완료된 ocr.DetTrainer class의 instance를 반환합니다. |
to_device(device=None)
DetTrainer의 device를 변경합니다. (cpu/gpu)
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
device |
Union[torch.device, str, int]
|
변경하고자 하는 device 입니다.
|
None
|
Returns:
| Name | Type | Description |
|---|---|---|
None | None |
train_one_step()
학습을 1스텝 수행하고, 결과를 반환합니다.
Returns:
validation_enabled()
DetTrainer가 validation이 가능한지 bool로 반환합니다.
Returns:
| Name | Type | Description |
|---|---|---|
bool | DetTrainer가 validation이 가능하면 True를, 불가능하다면 False를 반환합니다. |
validate()
post_train_process()
train 후처리를 수행하고, 해당 모델의 default inference options을 리턴합니다.
Returns:
| Name | Type | Description |
|---|---|---|
Dict |
Dict
|
각 parameter에 대한 설명은 DetRecInferencer.get_default_inference_option를 참고해주세요. |
save_checkpoint(checkpoint_path, metadata=None, password=None)
현재 DetTrainer의 상태를 암호화하여 저장합니다.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
checkpoint_path |
str
|
checkpoint를 저장할 path 입니다. |
required |
metadata |
Optional[Dict]
|
checkpoint에 저장할 추가 metadata. Defaults to None. |
None
|
password |
Optional[str]
|
checkpoint 파일에서 중요한 정보를 암호화 하는데 사용되는 password 입니다. None이면 암호화하지 않습니다. Defaults to None. |
None
|
Returns:
| Name | Type | Description |
|---|---|---|
None | None |
load_checkpoint(checkpoint_path, password=None, **options)
저장한 checkpoint로부터 DetTrainer의 상태를 불러오고, metadata를 반환합니다.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
checkpoint_path |
str
|
load할 checkpoint가 저장 되어 있는 path 입니다. |
required |
password |
Optional[str]
|
checkpoint 파일에서 중요한 정보를 복호화 하는데 사용되는 password 입니다. None이면 복호화하지 않습니다. Defaults to None. |
None
|
Returns:
| Name | Type | Description |
|---|---|---|
Dict |
dict
|
save_checkpoint에서 저장했던 metadata 입니다. |
auto-resize 설명
- 이미지 상의 텍스트 박스의 height 값이 특정한 길이가 되도록 하는 resize factor를 자동으로 찾아 주는 기능.
- 텍스트 검출 성능은 텍스트의 절대적인 크기에 영향을 받는데, 어떤 resize 값이 가장 적절한 지를 자동으로 결정해 줌.
- 학습 데이터셋을 한번 쭉 훑어 텍스트 라벨 height의 평균 값을 구하고, 이 값이 지정된 수치가 되도록 하는 resize factor를 계산함.
- (ex) 데이터셋의 mean text height가 12이고 미리 지정된 수치가 36이라면 resize 배율은 3이 됨.
- 옵션은 None 혹은 16-64 의 integer로, 리사이즈 이후의 텍스트 height 평균 값이 얼마가 될 지 지정함. 지정값이 클수록 최종 이미지 사이즈는 커짐. 디폴트는 None.
Preview API에서 설정한 roi config를 그대로 Trainer.build() config의 roi 섹션에 넣어주면 학습 시 해당 roi가 적용되며, 학습에서 설정하는 roi는 validation 및 inference 시에도 동일하게 적용됩니다.