目录

1. 数据集

2. config文件配置

 3. 测试模型


 

1. 数据集

这里以icdar2015字符检测为例https://blog.csdn.net/jizhidexiaoming/article/details/124149164?spm=1001.2014.3001.5501

2. config文件配置

_base_ = [
    '../../_base_/runtime_10e.py',
    # '../../_base_/schedules/schedule_sgd_1200e.py',
    '../../_base_/det_models/dbnet_r50dcnv2_fpnc.py',
    # '../../_base_/det_datasets/icdar2015.py',
    # '../../_base_/det_pipelines/dbnet_pipeline.py'
]

# datasets
dataset_type = 'IcdarDataset'
# data_root = 'data/icdar2015'
data_root = r'E:\data\ocr\text_detection\icdar2015_task1\icdar2015'
work_dir = "./work_dirs/dbnet_r50_0415"
train = dict(
    type=dataset_type,
    ann_file=f'{data_root}/instances_training.json',
    img_prefix=f'{data_root}/imgs',
    select_first_k=-1,  # 默认是-1,debug时,可以设置成1
    pipeline=None)
test = dict(
    type=dataset_type,
    ann_file=f'{data_root}/instances_test.json',
    img_prefix=f'{data_root}/imgs',
    pipeline=None)
train_list = [train]
test_list = [test]

# pipeline
img_norm_cfg_r50dcnv2 = dict(mean=[0, 0, 0], std=[1, 1, 1], to_rgb=True)
train_pipeline_r50dcnv2 = [
    dict(type='LoadImageFromFile', color_type='color_ignore_orientation'),
    dict(
        type='LoadTextAnnotations',
        with_bbox=True,
        with_mask=True,
        poly2mask=False),
    dict(type='ColorJitter', brightness=32.0 / 255, saturation=0.5),
    dict(type='Normalize', **img_norm_cfg_r50dcnv2),
    dict(
        type='ImgAug',
        args=[['Fliplr', 0.5],
              dict(cls='Affine', rotate=[-10, 10]), ['Resize', [0.5, 3.0]]]),
    dict(type='EastRandomCrop', target_size=(640, 640)),
    dict(type='DBNetTargets', shrink_ratio=0.4),
    dict(type='Pad', size_divisor=32),
    dict(
        type='CustomFormatBundle',
        keys=['gt_shrink', 'gt_shrink_mask', 'gt_thr', 'gt_thr_mask'],
        visualize=dict(flag=False, boundary_key='gt_shrink')),
    dict(
        type='Collect',
        keys=['img', 'gt_shrink', 'gt_shrink_mask', 'gt_thr', 'gt_thr_mask'])
]

test_pipeline = [
    dict(type='LoadImageFromFile', color_type='color_ignore_orientation'),
    # dict(type='Resize', img_scale=(640, 640), keep_ratio=True),
    dict(type='Normalize', **img_norm_cfg_r50dcnv2),
    # dict(type='Pad', size_divisor=32),
    dict(type='ImageToTensor', keys=['img']),
    dict(type='Collect', keys=['img']),
    # dict(
    #     type='MultiScaleFlipAug',
    #     img_scale=(4068, 1024),
    #     flip=False,
    #     transforms=[
    #         dict(type='Resize', img_scale=(2944, 736), keep_ratio=True),
    #         dict(type='Normalize', **img_norm_cfg_r50dcnv2),
    #         dict(type='Pad', size_divisor=32),
    #         dict(type='ImageToTensor', keys=['img']),
    #         dict(type='Collect', keys=['img']),
    #     ])
]

# load_from = None
# load_from = r'D:\code\python\mmocr\tools\work_dirs\dbnet_r50_dcnv2\epoch_20.pth'

data = dict(
    samples_per_gpu=8,
    workers_per_gpu=4,
    val_dataloader=dict(samples_per_gpu=1),
    test_dataloader=dict(samples_per_gpu=1),
    train=dict(
        type='UniformConcatDataset',
        datasets=train_list,
        pipeline=train_pipeline_r50dcnv2),
    val=dict(
        type='UniformConcatDataset',
        datasets=test_list,
        pipeline=test_pipeline),
    test=dict(
        type='UniformConcatDataset',
        datasets=test_list,
        pipeline=test_pipeline))

evaluation = dict(interval=100, metric='hmean-iou')


# # optimizer config
optimizer = dict(type='SGD', lr=0.001, momentum=0.9, weight_decay=0.0001)
# optimizer = dict(type='Adam', lr=0.1, weight_decay=0.0001)
optimizer_config = dict(grad_clip=dict(max_norm=35, norm_type=2))

# # learning policy (scheduler) config
lr_config = dict(
    policy='step',
    warmup='linear',
    warmup_iters=2000,
    warmup_ratio=1.0 / 3,
    step=[2500, 4000])

total_epochs = 1200

# save model
checkpoint_config = dict(interval=10)  # 每隔10个epoch保存一次模型

训练到80epoch就差不多了。

 3. 测试模型

mmocr 测试字符检测和识别模型_Mr.Q的博客-CSDN博客字符检测和识别https://blog.csdn.net/jizhidexiaoming/article/details/124273621

更多推荐