{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\nimport numpy as np\nimport pandas as pd\n\nimport cv2\nfrom PIL import Image\nfrom matplotlib import pyplot as plt\n\nfrom tensorflow.keras.datasets import mnist\nfrom tensorflow.keras import layers, Model, models\nfrom tensorflow.keras.utils import to_categorical\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom tensorflow.keras.callbacks import ReduceLROnPlateau\n\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-03-13T20:22:12.845168Z","iopub.execute_input":"2022-03-13T20:22:12.846312Z","iopub.status.idle":"2022-03-13T20:22:17.784114Z","shell.execute_reply.started":"2022-03-13T20:22:12.846135Z","shell.execute_reply":"2022-03-13T20:22:17.782758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv('../input/ultra-mnist/train.csv')","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:17.787010Z","iopub.execute_input":"2022-03-13T20:22:17.787404Z","iopub.status.idle":"2022-03-13T20:22:17.825450Z","shell.execute_reply.started":"2022-03-13T20:22:17.787339Z","shell.execute_reply":"2022-03-13T20:22:17.824454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's look at random image.","metadata":{}},{"cell_type":"code","source":"random_sample = train_df.sample()\nname_image = random_sample.to_numpy()[0][0]","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:17.827314Z","iopub.execute_input":"2022-03-13T20:22:17.827678Z","iopub.status.idle":"2022-03-13T20:22:17.835788Z","shell.execute_reply.started":"2022-03-13T20:22:17.827623Z","shell.execute_reply":"2022-03-13T20:22:17.834583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"random_sample","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:17.837451Z","iopub.execute_input":"2022-03-13T20:22:17.837955Z","iopub.status.idle":"2022-03-13T20:22:17.858902Z","shell.execute_reply.started":"2022-03-13T20:22:17.837899Z","shell.execute_reply":"2022-03-13T20:22:17.857707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"random_image = cv2.imread(f'../input/ultra-mnist/train/{name_image}.jpeg', 0)\nplt.imshow(random_image, cmap='Greys_r');","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:17.862138Z","iopub.execute_input":"2022-03-13T20:22:17.862513Z","iopub.status.idle":"2022-03-13T20:22:19.879911Z","shell.execute_reply.started":"2022-03-13T20:22:17.862468Z","shell.execute_reply":"2022-03-13T20:22:19.878977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"How we can see - at first we need to preprocess background.","metadata":{}},{"cell_type":"markdown","source":"# Background preprocessing","metadata":{}},{"cell_type":"markdown","source":"We can use this function for each image. It is different from fuction that was used for creation Ultra-MNIST-Black dataset. But my function is some better (in my opinion), because sometimes there were mistakes in that dataset. \nIn this function I look at left and top edges each square and compare them with neighbor squares.","metadata":{"execution":{"iopub.status.busy":"2022-03-11T14:03:41.249103Z","iopub.execute_input":"2022-03-11T14:03:41.249447Z","iopub.status.idle":"2022-03-11T14:03:41.255568Z","shell.execute_reply.started":"2022-03-11T14:03:41.249415Z","shell.execute_reply":"2022-03-11T14:03:41.254332Z"}}},{"cell_type":"code","source":"def background_converter(img):\n    img = np.array(img, dtype='int32')\n    for i in range(4):\n        for j in range(4):\n            top, left = False, False\n            img_slice = img[i * 1000:(i + 1) * 1000, j * 1000:(j + 1) * 1000]\n            \n            top_slice = img[i * 1000, j * 1000:(j + 1) * 1000]\n            left_slice = img[i * 1000:(i + 1) * 1000, j * 1000]\n            if i > 0:\n                top_slice_oppos = img[i * 1000 - 1, j * 1000:(j + 1) * 1000]\n            if j > 0:\n                left_slice_oppos = img[i * 1000:(i + 1) * 1000, j * 1000 - 1]\n            \n            if (i == 0 and top_slice.mean() > 250\n                or i > 0 and (top_slice != top_slice_oppos).sum() > 900):\n                top = True\n            if (j == 0 and left_slice.mean() > 250\n                or j > 0 and (left_slice != left_slice_oppos).sum() > 900):\n                left = True\n            if top or left:\n                img[i * 1000:(i + 1) * 1000, j * 1000:(j + 1) * 1000] = np.abs(img_slice - 255) \n    return img.astype('uint8')","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:19.881317Z","iopub.execute_input":"2022-03-13T20:22:19.881571Z","iopub.status.idle":"2022-03-13T20:22:19.895033Z","shell.execute_reply.started":"2022-03-13T20:22:19.881537Z","shell.execute_reply":"2022-03-13T20:22:19.894066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax = plt.subplots(1, 2, figsize=(18, 10))\nax[0].imshow(random_image, cmap='Greys_r')\nrandom_image_processed = background_converter(random_image)\nax[1].imshow(random_image_processed, cmap='Greys_r');","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:19.896165Z","iopub.execute_input":"2022-03-13T20:22:19.896414Z","iopub.status.idle":"2022-03-13T20:22:25.464410Z","shell.execute_reply.started":"2022-03-13T20:22:19.896385Z","shell.execute_reply":"2022-03-13T20:22:25.463506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Detecting numbers","metadata":{}},{"cell_type":"markdown","source":"Use threshold to get rid of noises.","metadata":{}},{"cell_type":"code","source":"def threshold_image(img):\n    return cv2.threshold(img, 200, 255, cv2.THRESH_BINARY)[1]","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:25.465830Z","iopub.execute_input":"2022-03-13T20:22:25.466099Z","iopub.status.idle":"2022-03-13T20:22:25.471757Z","shell.execute_reply.started":"2022-03-13T20:22:25.466065Z","shell.execute_reply":"2022-03-13T20:22:25.470204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"threshold_img = threshold_image(random_image_processed)\nplt.imshow(threshold_img, cmap='Greys_r');","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:25.473323Z","iopub.execute_input":"2022-03-13T20:22:25.473754Z","iopub.status.idle":"2022-03-13T20:22:27.370262Z","shell.execute_reply.started":"2022-03-13T20:22:25.473717Z","shell.execute_reply":"2022-03-13T20:22:27.368960Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Use OpenCV library to detect numbers and find bounding boxes for them.","metadata":{}},{"cell_type":"code","source":"def get_bound_rects(img):\n    contours, _ = cv2.findContours(img, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    bound_rects = []\n\n    for i, c in enumerate(contours):\n            rect = cv2.boundingRect(c)\n            if rect[3] < 9 or rect[2] < 2 or not 0.5 < rect[3] / rect[2] <= 10:\n                continue\n            bound_rects.append(rect)\n    \n    return np.array(bound_rects)","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:27.372033Z","iopub.execute_input":"2022-03-13T20:22:27.372397Z","iopub.status.idle":"2022-03-13T20:22:27.381211Z","shell.execute_reply.started":"2022-03-13T20:22:27.372349Z","shell.execute_reply":"2022-03-13T20:22:27.379883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bounding_boxes = get_bound_rects(threshold_img)\nbounding_boxes","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:27.383284Z","iopub.execute_input":"2022-03-13T20:22:27.383594Z","iopub.status.idle":"2022-03-13T20:22:27.410667Z","shell.execute_reply.started":"2022-03-13T20:22:27.383544Z","shell.execute_reply":"2022-03-13T20:22:27.409535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now we need to get rid of boxes which intersects with each other. We will remove box with lower square. Parameter **k** will define degree of intersection (the larger **k**, the larger area covered by rectangle)","metadata":{}},{"cell_type":"code","source":"def is_point_in_rect(point, rect, k):\n    x0, y0 = point\n    d = k * max(rect[2], rect[3])\n    centerx = rect[0] + rect[2] // 2\n    centery = rect[1] + rect[3] // 2\n    x1, y1, x2, y2 = centerx - d, centery - d, centerx + d, centery + d\n    if x1 <= x0 <= x2 and y1 <= y0 <= y2:\n        return True\n    return False","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-03-13T20:22:27.412577Z","iopub.execute_input":"2022-03-13T20:22:27.412829Z","iopub.status.idle":"2022-03-13T20:22:27.420140Z","shell.execute_reply.started":"2022-03-13T20:22:27.412800Z","shell.execute_reply":"2022-03-13T20:22:27.418787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def check_intersection(rect1, rect2, k):\n    for i in 0, rect1[2]:\n        for j in 0, rect1[3]:\n            point = rect1[0] + i, rect1[1] + j\n            if is_point_in_rect(point, rect2, k):\n                return True\n    for i in 0, rect2[2]:\n        for j in 0, rect2[3]:\n            point = rect2[0] + i, rect2[1] + j\n            if is_point_in_rect(point, rect1, k):\n                return True\n    return False","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-03-13T20:22:27.421604Z","iopub.execute_input":"2022-03-13T20:22:27.422331Z","iopub.status.idle":"2022-03-13T20:22:27.437378Z","shell.execute_reply.started":"2022-03-13T20:22:27.422279Z","shell.execute_reply":"2022-03-13T20:22:27.436344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def remove_intersec_rects(bound_rects, k):\n    intersec = np.ones(bound_rects.shape[0]).astype(bool)\n    for i, rect1 in enumerate(bound_rects):\n        for j, rect2 in enumerate(bound_rects):\n            if i == j:\n                continue\n            if check_intersection(rect1, rect2, k):\n                s1 = rect1[2] * rect1[3]\n                s2 = rect2[2] * rect2[3]\n                if s1 >= s2:\n                    intersec[j] = False\n                else:\n                    intersec[i] = False\n    return bound_rects[intersec]","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2022-03-13T20:22:27.441429Z","iopub.execute_input":"2022-03-13T20:22:27.441910Z","iopub.status.idle":"2022-03-13T20:22:27.451960Z","shell.execute_reply.started":"2022-03-13T20:22:27.441864Z","shell.execute_reply":"2022-03-13T20:22:27.451186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bounding_boxes = remove_intersec_rects(bounding_boxes, 0.6)\nbounding_boxes","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:27.453758Z","iopub.execute_input":"2022-03-13T20:22:27.454367Z","iopub.status.idle":"2022-03-13T20:22:27.471759Z","shell.execute_reply.started":"2022-03-13T20:22:27.454328Z","shell.execute_reply":"2022-03-13T20:22:27.471050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Cropping numbers","metadata":{}},{"cell_type":"markdown","source":"Now we can crop numbers from the image and process them independently.","metadata":{}},{"cell_type":"code","source":"def crop_numbers(img, bound_rects):\n    numbers = []\n    for rect in bound_rects:\n        x, y, w, h = rect\n        number = img[y:y + h, x:x + w]\n        numbers.append(number)\n    return numbers","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:27.473324Z","iopub.execute_input":"2022-03-13T20:22:27.473747Z","iopub.status.idle":"2022-03-13T20:22:27.478765Z","shell.execute_reply.started":"2022-03-13T20:22:27.473713Z","shell.execute_reply":"2022-03-13T20:22:27.478157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"This function will get rid of strange crops, because probably it is not numbers.","metadata":{}},{"cell_type":"code","source":"def check_anomaly(img, bounding_boxes, numbers):\n    k = 0.65\n    while len(numbers) > 5:\n        bound_rects = remove_intersec_rects(bounding_boxes, k)\n        numbers = crop_numbers(img, bound_rects)\n        k += 0.05\n    return numbers","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:27.480215Z","iopub.execute_input":"2022-03-13T20:22:27.480551Z","iopub.status.idle":"2022-03-13T20:22:27.491136Z","shell.execute_reply.started":"2022-03-13T20:22:27.480519Z","shell.execute_reply":"2022-03-13T20:22:27.490399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We recrop images to avoid mistakes of previous functions.","metadata":{}},{"cell_type":"code","source":"def recrop_numbers(numbers):\n    recropped_numbers = []\n    \n    for img in numbers:\n        h, w = img.shape\n        d = (max(h, w) // 20) * 255\n\n        left = 0\n        while img[:, left].sum() < d:\n            left += 1\n\n        if img[:, -1].sum() < d:\n            right = -1\n            while img[:, right - 1].sum() < d:\n                right -= 1\n        else:\n            right = None\n\n        up = 0\n        while img[up, :].sum() < d:\n            up += 1\n\n        if img[-1, :].sum() < d:\n            down = -1\n            while img[:, down - 1].sum() < d:\n                down -= 1\n        else:\n            down = None\n        \n        recropped_numbers.append(img[up:down, left:right])\n    \n    return recropped_numbers","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:27.494052Z","iopub.execute_input":"2022-03-13T20:22:27.494546Z","iopub.status.idle":"2022-03-13T20:22:27.506710Z","shell.execute_reply.started":"2022-03-13T20:22:27.494493Z","shell.execute_reply":"2022-03-13T20:22:27.505673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"numbers = crop_numbers(threshold_img, bounding_boxes)\ncleared_numbers = check_anomaly(threshold_img, bounding_boxes, numbers)\nrecropped_numbers = recrop_numbers(cleared_numbers)\n\nfig, ax = plt.subplots(1, len(recropped_numbers), figsize=(18, 10))\n\nfor i, number in enumerate(cleared_numbers):\n    ax[i].imshow(number, cmap='Greys_r')","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:27.508264Z","iopub.execute_input":"2022-03-13T20:22:27.508512Z","iopub.status.idle":"2022-03-13T20:22:28.436582Z","shell.execute_reply.started":"2022-03-13T20:22:27.508480Z","shell.execute_reply":"2022-03-13T20:22:28.435858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Resizing of numbers","metadata":{}},{"cell_type":"markdown","source":"Now we can notice that all numbers were scaled in an integer number of times. So we shoud try to find out this value.","metadata":{}},{"cell_type":"code","source":"def find_multiplicity(number):\n    h, w = number.shape\n    d = max(h, w)\n    if d <= 20:\n        return 1\n    divs = [i for i in range(d // 20, d // 18 + 1)]\n    for k in divs:\n        if d // k > 20:\n            continue\n        if any([(h + i) % k == 0 for i in (-1, 0, 1)]) and any([(w + i) % k == 0 for i in (-1, 0, 1)]):\n            return k","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:28.437987Z","iopub.execute_input":"2022-03-13T20:22:28.438393Z","iopub.status.idle":"2022-03-13T20:22:28.446755Z","shell.execute_reply.started":"2022-03-13T20:22:28.438353Z","shell.execute_reply":"2022-03-13T20:22:28.445627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def resize_numbers(numbers):\n    resized_numbers = []\n    n = len(numbers)\n    \n    for number in numbers: \n        k = find_multiplicity(number)\n        h, w = number.shape\n        if k is None:\n            k = 20 / max(h, w)\n            new_h = int(h * k)\n            new_w = int(w * k)\n                \n        else:\n            new_h, new_w = h // k, w // k     \n        \n        if new_h <= 10 and new_w <= 10:\n            new_h, new_w = new_h * 2, new_w * 2\n            \n        resized_numbers.append(cv2.resize(number, (new_w, new_h), interpolation=cv2.INTER_AREA))\n        \n    return resized_numbers","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:28.448699Z","iopub.execute_input":"2022-03-13T20:22:28.448983Z","iopub.status.idle":"2022-03-13T20:22:28.461501Z","shell.execute_reply.started":"2022-03-13T20:22:28.448950Z","shell.execute_reply":"2022-03-13T20:22:28.460738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"resized_numbers = resize_numbers(recropped_numbers)\n\nfig, ax = plt.subplots(1, len(resized_numbers), figsize=(18, 10))\n\nfor i, number in enumerate(resized_numbers):\n    ax[i].imshow(number, cmap='Greys_r')","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:28.462632Z","iopub.execute_input":"2022-03-13T20:22:28.463538Z","iopub.status.idle":"2022-03-13T20:22:29.167525Z","shell.execute_reply.started":"2022-03-13T20:22:28.463498Z","shell.execute_reply":"2022-03-13T20:22:29.166570Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now we should resize this images to shape 28x28 (original mnist dataset). We will use zero padding.","metadata":{}},{"cell_type":"code","source":"def zero_padding(numbers):\n    padding_numbers = []\n    \n    for number in numbers:\n        h, w = number.shape\n        up, down = (28 - h) // 2, 28 - h - (28 - h) // 2\n        number = np.vstack([np.zeros((up, w), 'uint8'), number, np.zeros((down, w), 'uint8')])\n        left, right = (28 - w) // 2, 28 - w - (28 - w) // 2\n        number = np.hstack([np.zeros((28, right), 'uint8'), number, np.zeros((28, left), 'uint8')])\n        padding_numbers.append(number)\n    \n    return padding_numbers","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:29.168859Z","iopub.execute_input":"2022-03-13T20:22:29.169109Z","iopub.status.idle":"2022-03-13T20:22:29.177404Z","shell.execute_reply.started":"2022-03-13T20:22:29.169078Z","shell.execute_reply":"2022-03-13T20:22:29.176669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"result_numbers = zero_padding(resized_numbers)\n\nfig, ax = plt.subplots(1, len(result_numbers), figsize=(18, 10))\n\nfor i, number in enumerate(result_numbers):\n    ax[i].imshow(number, cmap='Greys_r')","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:29.179076Z","iopub.execute_input":"2022-03-13T20:22:29.179366Z","iopub.status.idle":"2022-03-13T20:22:29.783087Z","shell.execute_reply.started":"2022-03-13T20:22:29.179333Z","shell.execute_reply":"2022-03-13T20:22:29.782091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Final function for extracting numbers from image.","metadata":{}},{"cell_type":"code","source":"def extract_numbers(image):\n    cleared_background = background_converter(image)\n    threshold_img = threshold_image(cleared_background)\n    \n    bounding_boxes = get_bound_rects(threshold_img)\n    cleared_bounding_boxes = remove_intersec_rects(bounding_boxes, 0.6)\n    \n    numbers = crop_numbers(threshold_img, cleared_bounding_boxes)\n    cleared_numbers = check_anomaly(threshold_img, cleared_bounding_boxes, numbers)\n    recropped_numbers = recrop_numbers(cleared_numbers)\n    resized_numbers = resize_numbers(recropped_numbers)\n    result_numbers = zero_padding(resized_numbers)\n    \n    result_numbers = (np.array(result_numbers).reshape(-1, 28, 28, 1) > 200).astype(float)\n    return result_numbers","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:29.784447Z","iopub.execute_input":"2022-03-13T20:22:29.784705Z","iopub.status.idle":"2022-03-13T20:22:29.793384Z","shell.execute_reply.started":"2022-03-13T20:22:29.784672Z","shell.execute_reply":"2022-03-13T20:22:29.792310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"random_sample = train_df.sample()\nname_image = random_sample.to_numpy()[0][0]\nrandom_image = cv2.imread(f'../input/ultra-mnist/train/{name_image}.jpeg', 0)\nres_numbers = extract_numbers(random_image)\n\nprint(f'Summ: {random_sample.to_numpy()[0][1]}')\nfig, ax = plt.subplots(1, len(res_numbers), figsize=(18, 10))\n\nfor i, number in enumerate(res_numbers):\n    ax[i].imshow(number, cmap='Greys_r')","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:29.794921Z","iopub.execute_input":"2022-03-13T20:22:29.795603Z","iopub.status.idle":"2022-03-13T20:22:30.517823Z","shell.execute_reply.started":"2022-03-13T20:22:29.795559Z","shell.execute_reply":"2022-03-13T20:22:30.516738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocessing training data","metadata":{}},{"cell_type":"markdown","source":"For training we will use MNIST, but as preprocessing we apply threshold to it and then center the numbers.","metadata":{}},{"cell_type":"code","source":"(X_train, Y_train), (X_test, Y_test) = mnist.load_data()","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:30.519281Z","iopub.execute_input":"2022-03-13T20:22:30.519542Z","iopub.status.idle":"2022-03-13T20:22:30.869516Z","shell.execute_reply.started":"2022-03-13T20:22:30.519508Z","shell.execute_reply":"2022-03-13T20:22:30.868589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_train = X_train.reshape((-1, 28, 28, 1))\nX_test = X_test.reshape((-1, 28, 28, 1))\nY_train = to_categorical(Y_train)\nY_test = to_categorical(Y_test)","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:30.870853Z","iopub.execute_input":"2022-03-13T20:22:30.871111Z","iopub.status.idle":"2022-03-13T20:22:30.878744Z","shell.execute_reply.started":"2022-03-13T20:22:30.871078Z","shell.execute_reply":"2022-03-13T20:22:30.878000Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Functions for preprocessing:","metadata":{}},{"cell_type":"code","source":"def crop_number(img):\n    left = 0\n    while img[:, left].sum() == 0:\n        left += 1\n        \n    right = -1\n    while img[:, right - 1].sum() == 0:\n        right -= 1\n        \n    up = 0\n    while img[up, :].sum() == 0:\n        up += 1\n        \n    down = -1\n    while img[down - 1:, :].sum() == 0:\n        down -= 1\n        \n    return img[up:down, left:right]","metadata":{"_kg_hide-input":false,"_kg_hide-output":false,"execution":{"iopub.status.busy":"2022-03-13T20:22:30.880294Z","iopub.execute_input":"2022-03-13T20:22:30.880523Z","iopub.status.idle":"2022-03-13T20:22:30.891943Z","shell.execute_reply.started":"2022-03-13T20:22:30.880495Z","shell.execute_reply":"2022-03-13T20:22:30.890710Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess_img(img):\n    k = np.random.randint(10, 240)\n    threshold_img = (img > k).reshape(28, 28)\n    number = crop_number(threshold_img)\n    result_number = zero_padding([number])\n    return result_number[0].reshape(28, 28, 1)","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-03-13T20:22:30.894435Z","iopub.execute_input":"2022-03-13T20:22:30.895575Z","iopub.status.idle":"2022-03-13T20:22:30.906106Z","shell.execute_reply.started":"2022-03-13T20:22:30.895512Z","shell.execute_reply":"2022-03-13T20:22:30.905317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"generator = ImageDataGenerator(preprocessing_function=preprocess_img)","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:30.907958Z","iopub.execute_input":"2022-03-13T20:22:30.908691Z","iopub.status.idle":"2022-03-13T20:22:30.919577Z","shell.execute_reply.started":"2022-03-13T20:22:30.908640Z","shell.execute_reply":"2022-03-13T20:22:30.918521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_flow = generator.flow(X_train, Y_train, 1000)\ntest_flow = generator.flow(X_test, Y_test, 1000)","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:30.921226Z","iopub.execute_input":"2022-03-13T20:22:30.921631Z","iopub.status.idle":"2022-03-13T20:22:31.024409Z","shell.execute_reply.started":"2022-03-13T20:22:30.921590Z","shell.execute_reply":"2022-03-13T20:22:31.023390Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's look at the result of preprocessing function. These images are more similar to ours.","metadata":{}},{"cell_type":"code","source":"ind = np.random.randint(0, 60001)\nplt.imshow(preprocess_img(generator.random_transform(X_train[ind])), cmap='Greys_r');","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:31.025808Z","iopub.execute_input":"2022-03-13T20:22:31.026085Z","iopub.status.idle":"2022-03-13T20:22:31.183002Z","shell.execute_reply.started":"2022-03-13T20:22:31.026052Z","shell.execute_reply":"2022-03-13T20:22:31.182038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Building and training a model","metadata":{}},{"cell_type":"markdown","source":"For prediction we will use CNN ensemble (99.7% accuracy on original MNIST dataset). ","metadata":{}},{"cell_type":"code","source":"def conv_model(input_shape):\n    \n    X = layers.Input(input_shape)\n\n    Y = X\n    for f in [8, 16, 24, 32, 48, 64, 96, 128, 192, 256, 384, 512]:\n        Y = layers.Conv2D(f, 3, 1, 'valid')(Y)\n        Y = layers.BatchNormalization()(Y)\n        Y = layers.ReLU()(Y)\n\n    Y = layers.Flatten()(Y)\n\n    Y = layers.Dense(256)(Y)\n    Y = layers.BatchNormalization()(Y)\n    Y = layers.ReLU()(Y)\n    \n    Y = layers.Dense(10)(Y)\n    Y = layers.BatchNormalization()(Y)\n    Y = layers.Softmax()(Y)\n    \n    model = Model(inputs=X, outputs=Y)\n    \n    return model","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:31.184521Z","iopub.execute_input":"2022-03-13T20:22:31.184837Z","iopub.status.idle":"2022-03-13T20:22:31.193772Z","shell.execute_reply.started":"2022-03-13T20:22:31.184739Z","shell.execute_reply":"2022-03-13T20:22:31.192856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"reduceLr = ReduceLROnPlateau(monitor='val_accuracy', \n                             patience=3,\n                             verbose=0,\n                             factor=0.8,\n                             min_lr=1e-5)","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:31.195426Z","iopub.execute_input":"2022-03-13T20:22:31.195875Z","iopub.status.idle":"2022-03-13T20:22:31.210845Z","shell.execute_reply.started":"2022-03-13T20:22:31.195844Z","shell.execute_reply":"2022-03-13T20:22:31.209976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#epochs = 40\n#steps = 60\n#models = []\n#for i in range(10):\n#    model = conv_model((28, 28, 1))\n#    model.compile('adam', 'categorical_crossentropy', ['accuracy'])\n#    model.fit(train_flow, \n#              epochs=epochs, \n#              steps_per_epoch=steps, \n#              verbose=0, \n#              callbacks=[reduceLr], \n#              validation_data=test_flow)\n#    models.append(model)","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:31.212108Z","iopub.execute_input":"2022-03-13T20:22:31.212399Z","iopub.status.idle":"2022-03-13T20:22:31.224646Z","shell.execute_reply.started":"2022-03-13T20:22:31.212366Z","shell.execute_reply":"2022-03-13T20:22:31.223635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def model_ensemble(input_shape, models):\n    \n    X = layers.Input(shape=input_shape)\n    \n    Y = []\n    for model in models:\n        model.trainable = False\n        Y.append(model(X))\n    Y = layers.Add()(Y)\n    \n    for units in [1024, 256, 64]:\n        Y = layers.Dense(units)(Y)\n        Y = layers.BatchNormalization()(Y)\n        Y = layers.ReLU()(Y)\n    \n    Y = layers.Dense(10)(Y)\n    Y = layers.BatchNormalization()(Y)\n    Y = layers.Softmax()(Y)\n    \n    model = Model(inputs=X, outputs=Y)\n    \n    return model","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:31.226040Z","iopub.execute_input":"2022-03-13T20:22:31.226329Z","iopub.status.idle":"2022-03-13T20:22:31.237710Z","shell.execute_reply.started":"2022-03-13T20:22:31.226289Z","shell.execute_reply":"2022-03-13T20:22:31.237054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model = model_ensemble((28,28,1), models)\n#model.compile('adam', 'categorical_crossentropy', ['accuracy'])\n#epochs = 2\n#steps = 10\n#model.fit(train_flow, \n#          epochs=epochs, \n#          steps_per_epoch=steps, \n#          verbose=2, \n#          callbacks=[reduceLr], \n#          validation_data=test_flow)","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:31.239303Z","iopub.execute_input":"2022-03-13T20:22:31.239991Z","iopub.status.idle":"2022-03-13T20:22:31.253402Z","shell.execute_reply.started":"2022-03-13T20:22:31.239950Z","shell.execute_reply":"2022-03-13T20:22:31.252283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model.save('ensemble_model')","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:31.255556Z","iopub.execute_input":"2022-03-13T20:22:31.256162Z","iopub.status.idle":"2022-03-13T20:22:31.273578Z","shell.execute_reply.started":"2022-03-13T20:22:31.256108Z","shell.execute_reply":"2022-03-13T20:22:31.272387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"It will take a long time to train a model, so I will pre-trained model from previous version.","metadata":{}},{"cell_type":"code","source":"model = models.load_model('../input/ultramnist-95-solution/ensemble_model/')","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:22:31.275056Z","iopub.execute_input":"2022-03-13T20:22:31.275839Z","iopub.status.idle":"2022-03-13T20:23:17.855600Z","shell.execute_reply.started":"2022-03-13T20:22:31.275789Z","shell.execute_reply":"2022-03-13T20:23:17.854411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"i = 0\nwhile True:\n    random_sample = train_df.sample()\n    name_image = random_sample.to_numpy()[0][0]\n    random_image = cv2.imread(f'../input/ultra-mnist/train/{name_image}.jpeg', 0)\n    res_numbers = extract_numbers(random_image)\n    pred = model.predict(res_numbers).argmax(axis=1)\n    i += 1\n    if random_sample.to_numpy()[0][1] != pred.sum():\n        break\nprint(f'Succesful iterations: {i - 1}\\n')\nprint(f'Real sum: {random_sample.to_numpy()[0][1]}\\nPredicted sum: {pred.sum()}')\nfig, ax = plt.subplots(1, len(res_numbers), figsize=(18, 10))\n\nfor i, number in enumerate(res_numbers):\n    ax[i].imshow(number, cmap='Greys_r')\n    ax[i].set_title(f'Predicted value: {pred[i]}')","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:23:17.860367Z","iopub.execute_input":"2022-03-13T20:23:17.860690Z","iopub.status.idle":"2022-03-13T20:23:25.578811Z","shell.execute_reply.started":"2022-03-13T20:23:17.860654Z","shell.execute_reply":"2022-03-13T20:23:25.576502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Final pipeline","metadata":{}},{"cell_type":"markdown","source":"Function that accept an image and return predicted sum of numbers. (In this function we will delete numbers, which have low probability or too large sum)","metadata":{}},{"cell_type":"code","source":"def predict_sum(image):\n    numbers = extract_numbers(image)\n    \n    prob = model.predict(numbers)\n    pred = np.argmax(prob, axis=1)\n        \n    s = pred.sum()\n    debug_s = s\n    \n    while s > 27 or prob.max(axis=1).min() < 0.12 and len(pred) > 3:\n        min_prob = prob.max(axis=1).argmin()\n        prob = np.delete(prob, min_prob, axis=0)\n        pred = prob.argmax(axis=1)\n        s = pred.sum()\n        \n    if len(pred) == 2:\n        s += np.random.randint(0, 10) \n    if len(pred) == 1:\n        s += np.random.randint(0, 19) \n    if len(pred) == 0:\n        s += np.random.randint(0, 28) \n    \n    return s","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:23:25.580443Z","iopub.execute_input":"2022-03-13T20:23:25.580780Z","iopub.status.idle":"2022-03-13T20:23:25.590878Z","shell.execute_reply.started":"2022-03-13T20:23:25.580744Z","shell.execute_reply":"2022-03-13T20:23:25.589702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's evaluate our model on train dataset.","metadata":{}},{"cell_type":"code","source":"positive = 0\n\nfor i, row in train_df.sample(1000).iterrows():\n    img_name = row['id']\n    res = row['digit_sum']\n    image = cv2.imread(f'../input/ultra-mnist/train/{img_name}.jpeg', 0)\n    \n    s = predict_sum(image)\n    \n    if res == s:\n        positive += 1\n        \nprint(f'Accuracy: {positive / 10}%')","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:23:25.592934Z","iopub.execute_input":"2022-03-13T20:23:25.593326Z","iopub.status.idle":"2022-03-13T20:24:07.902709Z","shell.execute_reply.started":"2022-03-13T20:23:25.593277Z","shell.execute_reply":"2022-03-13T20:24:07.901859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"markdown","source":"Now use our model with test dataset and make submission.","metadata":{}},{"cell_type":"code","source":"submission = pd.read_csv('../input/ultra-mnist/sample_submission.csv')","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:24:07.904616Z","iopub.execute_input":"2022-03-13T20:24:07.905269Z","iopub.status.idle":"2022-03-13T20:24:07.929138Z","shell.execute_reply.started":"2022-03-13T20:24:07.905217Z","shell.execute_reply":"2022-03-13T20:24:07.928052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i, row in submission.iterrows():\n    img_name = row['id']\n    image = cv2.imread(f'../input/ultra-mnist/test/{img_name}.jpeg', 0)\n    \n    s = predict_sum(image)\n    \n    submission.loc[i, 'digit_sum'] = s","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:24:35.002493Z","iopub.execute_input":"2022-03-13T20:24:35.003744Z","iopub.status.idle":"2022-03-13T20:24:39.113706Z","shell.execute_reply.started":"2022-03-13T20:24:35.003686Z","shell.execute_reply":"2022-03-13T20:24:39.112647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-03-13T20:24:39.115854Z","iopub.execute_input":"2022-03-13T20:24:39.116193Z","iopub.status.idle":"2022-03-13T20:24:39.179862Z","shell.execute_reply.started":"2022-03-13T20:24:39.116138Z","shell.execute_reply":"2022-03-13T20:24:39.178881Z"},"trusted":true},"execution_count":null,"outputs":[]}]}