
이 글에서는 데이터셋 구축부터 OCR 모델 학습까지 다루어 본다.
OCR모델은 Text Detection과 Text Recognition으로 나뉜다.
여기서는 Recognition만 해보고자 한다.
광학 문자 인식(OCR) 이란 사람이 쓰거나 기계로 인쇄한 문자 이미지를 기계가 읽을 수 있는 데이터로 변환하는 과정




import tensorflow as tf
loss_fn = tf.keras.backend.ctc_batch_cost
def CTC_loss(y_true, y_pred):
batch_len = tf.cast(tf.shape(y_true)[0], dtype="int64") # 16
input_length = tf.cast(tf.shape(y_pred)[1], dtype="int64") # 50
label_length = tf.cast(tf.shape(y_true)[1], dtype="int64") # 5
input_length = input_length * tf.ones(shape=(batch_len, 1), dtype="int64")
label_length = label_length * tf.ones(shape=(batch_len, 1), dtype="int64")
loss = loss_fn(y_true, y_pred, input_length, label_length)
return loss
def build_model(input_shape=(MAX_WIDTH, IMG_HEIGHT, 3)):
input_img = layers.Input(shape=input_shape)
arg_dic = {"activation":"swish", "padding":"same", "kernel_initializer":"he_normal"}
x = layers.Conv2D(32, 3, **arg_dic)(input_img)
x = layers.MaxPooling2D()(x)
x = layers.Conv2D(64, 3, **arg_dic)(x)
x = layers.MaxPooling2D()(x)
arg_dic["activation"] = None
x = layers.Conv2D(64, 3, **arg_dic)(x)
x = layers.BatchNormalization()(x)
x = layers.Activation("swish")(x)
x = layers.Dropout(0.15)(x)
x = layers.MaxPooling2D()(x)
x = layers.Reshape(target_shape=(MAX_WIDTH // 8, -1), name="Reshape")(x)
# x = layers.Conv1D(128, 1, **arg_dic)(x)
x = layers.Conv1D(256, 3, **arg_dic)(x)
x = layers.Conv1D(128, 1, activation=None, kernel_initializer="he_normal")(x)
x = layers.BatchNormalization()(x)
x = layers.Activation("swish")(x)
x = layers.Dropout(0.2)(x)
# RNNs
x = layers.Bidirectional(layers.LSTM(128, return_sequences=True, dropout=0.25))(x)
x = layers.Bidirectional(layers.LSTM(64, return_sequences=True, dropout=0.25))(x)
x = layers.Bidirectional(layers.LSTM(32, return_sequences=True, dropout=0.25))(x)
# # Output layer
# x = layers.Dense(len(char_to_num.get_vocabulary()) + 1,
# activation="softmax")(x)
x = layers.Conv1D(len(char_to_num.get_vocabulary()) + 1, 1, activation='softmax')(x)
# Define the model
model = tf.keras.models.Model(inputs=input_img, outputs=x, name="ocr_model_v1")
return model

