{"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":"markdown","source":"# TensorFlow, TPU, UltraMNIST: Simple tricks to facilitate extra large image classification\n\n\"UltraMNIST Classification Challenge\" sets the task to classify extra large grayscale images (4000x4000 pixels) into 28 classes. Each image contains several MNIST digits of different sizes on top of some artificial background. The goal is to predict the sum of all digits in the picture. Large resolution of the images by itself is not that important because we can always apply something like `tf.image.resize`. What makes the task complicated is the presence of very small digits which may just disappear if we resize images too small. In this notebook I will describe several tricks which I used in my solution (6th place). They may be useful in similar tasks where we have very large images and the target signal is represented by some small objects. Please note that some of the following tricks trade speed for accuracy.\n\nMy solution is based on the EfficientNet-B7 backbone and 2-stage training procedure. In the first stage we train the model using 1024x1024 resolution, and then in the second stage we finetune the model using 1536x1536 resolution. Huge HBM capacity of TPUv3 allows the use of a reasonably large global batch size of 16 in both stages. Single fold accuracy for the stage-1 is 0.87, and for the stage-2 is 0.94.\n","metadata":{}},{"cell_type":"code","source":"import IPython\nIPython.display.Image('/kaggle/input/umnist-misc/tiny_eight.png')","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-04-15T11:51:36.368896Z","iopub.execute_input":"2022-04-15T11:51:36.369275Z","iopub.status.idle":"2022-04-15T11:51:36.406851Z","shell.execute_reply.started":"2022-04-15T11:51:36.369162Z","shell.execute_reply":"2022-04-15T11:51:36.406041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Tricks\n\n## Trick 1. Extract all available information from the image files\n\nThis one does not sound like a trick at all because what may be more obvious than just reading and decoding an image file? But in the case of a JPEG this may be not that straightforward. By default decoding is optimized for speed and the DCT (discrete cosine transform) method is set to ‘INTEGER_FAST’. This method is faster but introduces some approximation. In order to get exact pixel values, set this parameter to 'INTEGER_ACCURATE'.\n\n`tf.image.decode_jpeg(image_bytes, dct_method='INTEGER_ACCURATE')`\n\n\n## Trick 2. Use antialias when resizing\n\nAntialias helps to avoid distortion after resizing. It was very important in this competition. When we have small objects in our images, antialias may significantly improve the result.\n\n`tf.image.resize(image, size=[1024, 1024], antialias=True)`\n\n\n## Trick 3. Try different filters for resizing\n\nFilter choice may be task specific. In general Bilinear interpolation is a preferred option and it worked best for me in this competition. But sometimes other high quality filters e.g. Bicubic or Lanczos may give better results.\n\n`tf.image.resize(image, size=[1024, 1024], antialias=True, method='bilinear')`\n\n\n## Trick 4. Use least possible or lossless compression\n\nIf we preprocess images in some way (e.g. resize) we will have to recompres them. In this case while creating TFRecords don’t forget to use the minimal possible compression ratio (JPEG) or even use lossless compression (PNG) in order to preserve the most of the available information. Also we can resize original images online during training. It was my choice in this competition.\n\n\n## Trick 5. Select augmentations carefully\n\nSome augmentations may introduce distortions which are not acceptable for the task at hand. For example, when using standard functions to perform rotation (e.g. `tfa.image.rotate`), we should remember that if angle is not multiple of 90, then a small fraction of image corners will be truncated by default. If the small object is located close to the corner it may be lost. Another example is brightness/contrast/gamma augmentations. If the pixel values of our small object are close to the min/max values, such augmentations may lead to saturation, making pixel values just the same. In this notebook I use only one specific type of augmentation: inverted images (i.e. 255 - image).\n\n\n## Trick 6. Try large pretrained models\n\nIt’s considered a good practice to start from some small model e.g. EfficientNet-B0. But we need to remember that smaller models are often pretrained on (and possibly optimized for) smaller image sizes. When we train a small model on very large images it may just not learn, or the training process may become unstable. In this notebook I use the EfficientNet-B7 model.\n\n\n## Trick 7. Choose large enough resolution, but not too large\n\nAs a first thought we may try to train on the largest resolution which our hardware can afford, hoping to capture more signal from the small objects. But this approach does not necessarily lead to the best performance. First of all, large image resolution forces us to use a small batch size which may lead to an unstable optimization process. Training becomes slower, it becomes much harder to optimize hyperparameters. The better choice is to choose some moderately large resolution which ensures stable training progress, and then when this model reaches some reasonable accuracy, finetune it using larger resolution (see next trick). It requires some experimentation to choose the best starting resolution for a given task. To be specific in this competition I tried 512x512, 1024x1024, and 1536x1536. Eventually 1024x1024 turned out to ensure the best accuracy.\n\n\n## Trick 8. Finetune the model using increasingly larger resolution\n\nThis approach not only allows to significantly improve the accuracy, but also can speed up the whole training process compared with the case if we opt to use larger resolution from the beginning. I started from 1024x1024, and then finetuned using 1536x1536.","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"markdown","source":"# Code\n\n### Import","metadata":{}},{"cell_type":"code","source":"!pip install --user -q efficientnet\n\nimport os\nimport glob\nimport shutil\nimport numpy as np\nimport tensorflow as tf\nimport efficientnet.tfkeras as efn\nfrom kaggle_datasets import KaggleDatasets","metadata":{"execution":{"iopub.status.busy":"2022-04-15T11:51:42.547032Z","iopub.execute_input":"2022-04-15T11:51:42.547751Z","iopub.status.idle":"2022-04-15T11:52:02.363139Z","shell.execute_reply.started":"2022-04-15T11:51:42.547710Z","shell.execute_reply":"2022-04-15T11:52:02.362515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Settings\n\nWe will train and finetune single fold only.  \nWhen early stopping interrupts stage-1 training, uncomment stage-2 settings and restart.  \nTraining loop below allows to train all folds consequently, or any single fold out of five.","metadata":{}},{"cell_type":"code","source":"# Stage-1 settings (1024x1024)\n\nclass args:\n    data_tfrec_dir = KaggleDatasets().get_gcs_path('umnist-tfrec-5f')\n    model_func = efn.EfficientNetB7 # Model function\n    job = 'train'            # Job 'train' or 'cont'. Use 'cont' to continue training\n    n_folds = 5              # Number of folds\n    initial_fold = 0         # Initial fold (0..4)\n    final_fold = 1           # Final fold (1..5)\n    n_channels = 3           # Number of image channels\n    dim = 1024               # Image size in pixels\n    n_examples_train = 22400 # 4/5 of train set (i.e. 4 folds)\n    initial_epoch = 0        # Initial epoch. Set appropriately if continue training\n    n_epochs = 200           # Max number of epochs\n    batch_size = 16          # Batch size\n    lr = 5e-4                # Learning rate\n    \n# Stage-2 settings (1536x1536)\n\n# args.job = 'cont'\n# args.dim = 1536\n# args.initial_epoch = 50 # set your actual value\n# args.lr = 1e-4\n    \nassert args.n_folds <= 10, 'Max 10 folds supported'\nfor a in [x for x in sorted(vars(args)) if '__' not in x]:\n    print('%-20s %s' % (a, vars(args)[a]))","metadata":{"execution":{"iopub.status.busy":"2022-04-15T11:52:05.909059Z","iopub.execute_input":"2022-04-15T11:52:05.909679Z","iopub.status.idle":"2022-04-15T11:52:06.412878Z","shell.execute_reply.started":"2022-04-15T11:52:05.909637Z","shell.execute_reply":"2022-04-15T11:52:06.411947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Definitions","metadata":{}},{"cell_type":"code","source":"def init_tpu():\n    \"\"\"\n    Seamlessly init any accelerator: CPU, GPU, multi-GPU, TPU\n    \"\"\"\n    try:\n        tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect()\n        strategy = tf.distribute.TPUStrategy(tpu)\n        print('Master:      ', tpu.master())\n    except ValueError:\n        strategy = tf.distribute.MirroredStrategy()\n        print('TPU was not found!')\n    print('Num replicas:', strategy.num_replicas_in_sync)\n    return strategy\n\n\ndef init_tfdata(files_glob, deterministic=True, batch_size=32, \n                auto=-1, parse_example=None, aug=None, norm=None, \n                repeat=False, buffer_size=None, cache=False, \n                drop_remainder=False):\n    \"\"\"\n    Init tf.data.TFRecordDataset with a specified \n    parameters and preprocessing\n    \n    files_glob : str\n        Glob wildcard for TFRecord files\n    deterministic : bool\n    batch_size : int\n    auto : int\n    parse_example, aug, norm : callable\n        Processing functions\n    repeat : bool\n        Whether to repeat dataset\n    buffer_size : int or None\n        Shuffle buffer size. None means do not shuffle\n    cache : bool\n        Whether to cache data\n    drop_remainder : bool\n        Whether to drop remainder (incomplete batch)\n    \"\"\"\n    options = tf.data.Options()\n    options.experimental_deterministic = deterministic\n    files = tf.data.Dataset.list_files(files_glob, \n        shuffle=not deterministic).with_options(options)\n    print('N tfrec files:', len(files))\n    ds = tf.data.TFRecordDataset(files, num_parallel_reads=auto)\n    ds = ds.with_options(options)\n    ds = ds.map(parse_example, num_parallel_calls=auto)\n    if aug:\n        ds = ds.map(aug, num_parallel_calls=auto)\n    if norm:\n        ds = ds.map(norm, num_parallel_calls=auto)\n    if repeat:\n        ds = ds.repeat()\n    if buffer_size:\n        ds = ds.shuffle(buffer_size=buffer_size, \n            reshuffle_each_iteration=True)\n    ds = ds.batch(batch_size=batch_size, \n        drop_remainder=drop_remainder)\n    ds = ds.prefetch(auto)\n    if cache:\n        ds = ds.cache()\n    return ds\n\n\nclass KeepLastCKPT(tf.keras.callbacks.Callback):\n    \"\"\"\n    Sort checkpoints by name and remove all except the last.\n    Without this callback keeping only last ckpt would require\n    to use the same name for all ckpts e.g. \"model.h5\"\n    to force each new ckpt to overwrite the previous.\n    Instead this callback allows to use arbitrary name\n    e.g. \"model-epoch-001-loss-0.95.h5\"\n    \"\"\"\n    def __init__(self, wildcard):\n        super(KeepLastCKPT, self).__init__()\n        self.wildcard = wildcard\n    def on_epoch_begin(self, epoch, logs=None):\n        files = sorted(tf.io.gfile.glob(self.wildcard))\n        if len(files):\n            for file in files[:-1]:\n                tf.io.gfile.remove(file)\n            print('Kept ckpt: %s' % files[-1])\n        else:\n            print('No ckpt to keep')\n    def on_train_end(self, logs=None):\n        files = sorted(tf.io.gfile.glob(self.wildcard))\n        if len(files):\n            for file in files[:-1]:\n                tf.io.gfile.remove(file)\n            print('\\nKept ckpt (final): %s' % files[-1])\n        else:\n            print('\\nNo ckpt to keep (final)')\n\n\nfeature_description = {\n    'image':    tf.io.FixedLenFeature([], tf.string),\n    'label':    tf.io.FixedLenFeature([], tf.int64),\n}\n\n\ndef parse_example(example_proto):\n    \"\"\"\n    Parse TFRecord example.\n    Using \"dct_method='INTEGER_ACCURATE'\" to get precise pixel values\n    \"\"\"\n    d = tf.io.parse_single_example(example_proto, feature_description)\n    image = tf.image.decode_jpeg(d['image'], \n                channels=args.n_channels, \n                dct_method='INTEGER_ACCURATE')\n    label = tf.cast(d['label'], tf.int32)\n    return image, label\n\n\ndef aug(image, label):\n    \"\"\"\n    Perform inversion augmentation with 50% proba\n    \"\"\"\n    if tf.random.uniform([], minval=0, maxval=2, dtype=tf.int32) == 1:\n        image = 255 - image\n    return image, label\n\n\ndef norm(image, label):\n    \"\"\"\n    Resize and normalize the image after all transformations.\n    Use of \"antialias=True\" is important for small objects\n    \"\"\"\n    image = tf.image.resize(image, [args.dim, args.dim],    \n                antialias=True, \n                method='bilinear')\n    image = tf.reshape(image, [args.dim, args.dim, args.n_channels])\n    label = tf.cast(label, tf.int32)\n    image = image / 255.0\n    return image, label\n\n\ndef init_model(print_summary=True):\n    \"\"\"\n    Init model and print summary\n    \"\"\"\n    model = tf.keras.Sequential([\n        args.model_func(input_shape=(args.dim, args.dim, args.n_channels), \n                        weights='imagenet', \n                        include_top=False),\n        tf.keras.layers.GlobalAveragePooling2D(),\n        tf.keras.layers.Dense(28, activation='softmax')\n    ], name='model')\n    model.compile(optimizer=tf.keras.optimizers.Adam(args.lr), \n                  loss='sparse_categorical_crossentropy',\n                  metrics=['acc'])\n    if print_summary:\n        model.summary()\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-04-15T11:52:14.883520Z","iopub.execute_input":"2022-04-15T11:52:14.883803Z","iopub.status.idle":"2022-04-15T11:52:14.914730Z","shell.execute_reply.started":"2022-04-15T11:52:14.883762Z","shell.execute_reply":"2022-04-15T11:52:14.913863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training loop","metadata":{}},{"cell_type":"code","source":"for fold_id in range(args.initial_fold, args.final_fold):\n    print('\\n*****')\n    print('Fold:', fold_id)\n    print('*****\\n')\n    #--------------------------------------------------------------------------\n    print('Clear session...')\n    tf.keras.backend.clear_session()\n    #--------------------------------------------------------------------------\n    print('Full batch shape: %d x %d x %d x %d' % (args.batch_size, args.dim, \n                                                   args.dim, args.n_channels))\n    #--------------------------------------------------------------------------\n    print('Init TPU')\n    strategy = init_tpu()\n    #--------------------------------------------------------------------------\n    print('Init globs')\n    all_fold_ids = np.array(range(args.n_folds))\n    train_fold_ids = all_fold_ids[all_fold_ids != fold_id]\n    train_glob = os.path.join(args.data_tfrec_dir, \n        ('fold.[' + '%d'*(args.n_folds-1) + '].tfrecord*') % tuple(train_fold_ids))\n    val_glob   = os.path.join(args.data_tfrec_dir, \n        'fold.[%d].tfrecord*' % fold_id)\n    print('TRAIN GLOB:', train_glob)\n    print('VAL   GLOB:', val_glob)\n    #--------------------------------------------------------------------------\n    print('Init datasets')\n    train_ds = init_tfdata(train_glob, \n                           deterministic=False,  \n                           batch_size=args.batch_size, \n                           auto=-1,\n                           parse_example=parse_example, \n                           aug=aug, \n                           norm=norm,\n                           repeat=True,\n                           buffer_size=256)\n    val_ds = init_tfdata(val_glob, \n                         deterministic=True,  \n                         batch_size=args.batch_size, \n                         auto=-1,\n                         parse_example=parse_example,\n                         norm=norm)\n    #--------------------------------------------------------------------------\n    print('Init model')\n    with strategy.scope():\n        model = init_model()\n    #--------------------------------------------------------------------------\n    if 'cont' in args.job:\n        m = sorted(glob.glob('model-f%d*.h5' % fold_id))[-1]\n        print('Continue (fold %d) from: %s' % (fold_id, m))\n        model.load_weights(m)\n        args.job = args.job.replace('cont', 'train')\n        shutil.copy(m, 'backup-' + m)\n    #--------------------------------------------------------------------------\n    print('Init callbacks')\n    call_ckpt = tf.keras.callbacks.ModelCheckpoint(\n        'model-f%d-e{epoch:03d}-{val_loss:.4f}-{val_%s:.4f}.h5' % (fold_id, 'acc'),\n         monitor='val_acc',\n         save_best_only=True,\n         save_weights_only=True,\n         mode='auto',\n         verbose=1)\n    call_reduce_lr = tf.keras.callbacks.ReduceLROnPlateau(\n        monitor='val_acc', \n        factor=0.5, \n        patience=4, \n        min_delta=1e-4,\n        min_lr=1e-8,\n        verbose=1,\n        mode='auto')\n    call_early_stop = tf.keras.callbacks.EarlyStopping(\n        monitor='val_acc',\n        patience=8,\n        min_delta=1e-4,\n        mode='auto',\n        verbose=1)\n    call_keep_last = KeepLastCKPT(wildcard='model-f%d-e*.h5' % fold_id)\n    #--------------------------------------------------------------------------\n    if 'train' in args.job:\n        print('Fit (fold %d)' % fold_id)\n        h = model.fit(\n            train_ds,\n            steps_per_epoch=args.n_examples_train // args.batch_size,\n            epochs=args.n_epochs,\n            initial_epoch=args.initial_epoch,\n            validation_data=val_ds,\n            callbacks=[call_ckpt,\n                       call_reduce_lr,\n                       call_early_stop,\n                       call_keep_last])\n    # Non-zero initial_epoch value is actve for the curent fold only;\n    # for all other folds it is reset to 0\n    args.initial_epoch = 0\n    #--------------------------------------------------------------------------\n    #--------------------------------------------------------------------------","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# References\n\nhttps://www.tensorflow.org/api_docs/python/tf/image  \nhttps://www.kaggle.com/docs/tpu","metadata":{}}]}