{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"To see how tfrecords are generated, check https://www.kaggle.com/code/hoyso48/aslfr-create-tfr\n\nNOTES\n\n1. (Colab only) You should update GCS_PATH by running code in https://www.kaggle.com/hoyso48/aslfr-get-gcs-path/edit yourself as it expires after several weeks.\n2. If you want to use GPU, set device = 'GPU' in get_strategy. You should set policy = 'float16' in CFG and dtype='float16' in get_model if you want fp16 training with GPU.\n3. It is recommended to use the following setting instead: CFG.epoch = 200, CFG.awp_lr = 0.1. It will give you a model with decent performance. set CFG.epoch = 60 and CFG.awp = False if you want the quick results. But if you want to reproduce the exact solution, see 4,5 carefully.\n4. Colab TPU runtime has recently been reduced to 3-4 hours. Therefore you cannot reproduce the solution within a single runtime(single model training in TPUv2-8 takes around ~14 hours). training should be done with the resume from the last checkpoint several times: set CFG.resume = 'auto' and run train_folds multiple times. you need to complete training of seed=42,43,44 with fold='all' to fully reproduce the solution.\n5. There is an instability in training(i.e. nan loss or huge accuracy drop occurs sometimes) with epoch=400 and awp_lr=0.2 setting. You can simply restart the training by setting CFG.resume=0. But if you don't want to throw away the training results, you can retry from the best checkpoint: set CFG.resume=best epoch(ex.130) and CFG.resume_ckpt = best ckpt(ex.'model-best.h5')","metadata":{}},{"cell_type":"code","source":"import os\n# if os.path.isdir('/content/drive/MyDrive'):\n#     os.makedirs('/content/drive/MyDrive/aslfr', exist_ok=True)\n#     os.chdir('/content/drive/MyDrive/aslfr')\n# else:\n#     os.makedirs('/content/aslfr', exist_ok=True)\n#     os.chdir('/content/aslfr')\nprint(os.getcwd())","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# copy our file into the working directory (make sure it has .py suffix)\n# code from @shlomoron @irohith\nfrom shutil import copyfile\ncopyfile(src = \"/kaggle/input/ctc-tpu/CTC_TPU.py\", dst = \"/kaggle/working//CTC_TPU.py\")\n\n# import all our functions\nfrom CTC_TPU import classic_ctc_loss","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q tensorflow-addons\n!pip install -q git+https://github.com/hoyso48/tf-utils@main\n!pip install -q Levenshtein\n!pip install -q keras_nlp","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\n\nimport tensorflow as tf\nimport tensorflow_addons as tfa\nimport keras_nlp\nimport tensorflow.keras.mixed_precision as mixed_precision\n\nfrom tf_utils.schedules import OneCycleLR, ListedLR\nfrom tf_utils.callbacks import Snapshot, SWA\nfrom tf_utils.learners import FGM, AWP\n\nimport matplotlib.pyplot as plt\nimport matplotlib as mpl\nfrom tqdm.autonotebook import tqdm\nimport sklearn\n\nimport os\nimport time\nimport pickle\nimport math\nimport random\nimport sys\nimport cv2\nimport gc\nimport glob\nimport datetime\n\nfrom Levenshtein import distance\n\nprint(f'Tensorflow Version: {tf.__version__}')\nprint(f'Python Version: {sys.version}')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Seed all random number generators\ndef seed_everything(seed=42):\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n\ndef get_strategy(device='TPU'):\n    if \"TPU\" in device:\n        tpu = 'local' if device=='TPU-VM' else None\n        print(\"connecting to TPU...\")\n        tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect(tpu=tpu)\n        strategy = tf.distribute.TPUStrategy(tpu)\n        IS_TPU = True\n\n    if device == \"GPU\"  or device==\"CPU\":\n        ngpu = len(tf.config.experimental.list_physical_devices('GPU'))\n        if ngpu>1:\n            print(\"Using multi GPU\")\n            strategy = tf.distribute.MirroredStrategy()\n        elif ngpu==1:\n            print(\"Using single GPU\")\n            strategy = tf.distribute.get_strategy()\n        else:\n            print(\"Using CPU\")\n            strategy = tf.distribute.get_strategy()\n        IS_TPU = False\n\n    if device == \"GPU\":\n        print(\"Num GPUs Available: \", ngpu)\n\n    AUTO     = tf.data.experimental.AUTOTUNE\n    REPLICAS = strategy.num_replicas_in_sync\n    print(f'REPLICAS: {REPLICAS}')\n\n    return strategy, REPLICAS, IS_TPU\n\nSTRATEGY, N_REPLICAS, IS_TPU = get_strategy('TPU-VM')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#NOTE: you should run KaggleDatasets.get_gcs_path(dataset_name) in the kaggle notebook to update gcs_path as they expires after several weeks..\n#notebook: https://www.kaggle.com/hoyso48/aslfr-get-gcs-path/edit\n\nGCS_PATH = {\n            'aslfr':'/kaggle/input/asl-fingerspelling',\n            'aslfr-5fold':'/kaggle/input/aslfr-5fold',\n            }\n\nTRAIN_FILENAMES = tf.io.gfile.glob(GCS_PATH['aslfr-5fold']+'/*.tfrecords')\nCOMPETITION_PATH = GCS_PATH['aslfr']\n\nprint(len(TRAIN_FILENAMES))\n# !gsutil cp {COMPETITION_PATH}/train.csv .\n# !gsutil cp {COMPETITION_PATH}/character_to_prediction_index.json .","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\nwith open('/kaggle/input/asl-fingerspelling/character_to_prediction_index.json') as json_file:\n    CHAR_TO_NUM = json.load(json_file)\nNUM_TO_CHAR = dict([(y+1,x) for x,y in CHAR_TO_NUM.items()] )\nNUM_TO_CHAR[60] = 'S'\nNUM_TO_CHAR[61] = 'E'\nNUM_TO_CHAR[0] = 'P'\n\n# LABEL_DICT","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TABLE = tf.lookup.StaticHashTable(\n    initializer=tf.lookup.KeyValueTensorInitializer(\n        keys=list(NUM_TO_CHAR.values()),\n        values=list(NUM_TO_CHAR.keys()),\n    ),\n    default_value=tf.constant(-1),\n    name=\"class_weight\"\n)\n\ndef preprocess_phrase(phrase, table=TABLE):\n    phrase = tf.strings.join(['S', phrase, 'E']) #'S'+ phrase + 'E'\n    phrase = tf.strings.bytes_split(phrase)\n    phrase = table.lookup(phrase)\n    return phrase","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train DataFrame\ntrain_df = pd.read_csv('/kaggle/input/asl-fingerspelling/train.csv')\ndisplay(train_df.head())\ndisplay(train_df.info())","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import re\ndef count_data_items(filenames):\n    n = [int(re.compile(r\"-([0-9]*)\\.\").search(filename.split('/')[-1]).group(1)) for filename in filenames]\n    return np.sum(n)\nprint(count_data_items(TRAIN_FILENAMES), len(train_df))\nassert count_data_items(TRAIN_FILENAMES) == len(train_df)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#for the lip_lr function. LEFT[i] is matching with RIGHT[i](i.e LEFT[i](x) == -RIGHT[i](x)).\n#computed from https://github.com/google/mediapipe/blob/master/mediapipe/modules/face_geometry/data/canonical_face_model.obj\n\nLEFT = [\n         248, 249, 250, 251, 252, 253, 254, 255, 256, 257, 258, 259, 260, 261, 262, 263, 264,\n         265, 266, 267, 268, 269, 270, 271, 272, 273, 274, 275, 276, 277, 278, 279, 280, 281,\n         282, 283, 284, 285, 286, 287, 288, 289, 290, 291, 292, 293, 294, 295, 296, 297, 298,\n         299, 300, 301, 302, 303, 304, 305, 306, 307, 308, 309, 310, 311, 312, 313, 314, 315,\n         316, 317, 318, 319, 320, 321, 322, 323, 324, 325, 326, 327, 328, 329, 330, 331, 332,\n         333, 334, 335, 336, 337, 338, 339, 340, 341, 342, 343, 344, 345, 346, 347, 348, 349,\n         350, 351, 352, 353, 354, 355, 356, 357, 358, 359, 360, 361, 362, 363, 364, 365, 366,\n         367, 368, 369, 370, 371, 372, 373, 374, 375, 376, 377, 378, 379, 380, 381, 382, 383,\n         384, 385, 386, 387, 388, 389, 390, 391, 392, 393, 394, 395, 396, 397, 398, 399, 400,\n         401, 402, 403, 404, 405, 406, 407, 408, 409, 410, 411, 412, 413, 414, 415, 416, 417,\n         418, 419, 420, 421, 422, 423, 424, 425, 426, 427, 428, 429, 430, 431, 432, 433, 434,\n         435, 436, 437, 438, 439, 440, 441, 442, 443, 444, 445, 446, 447, 448, 449, 450, 451,\n         452, 453, 454, 455, 456, 457, 458, 459, 460, 461, 462, 463, 464, 465, 466, 467,  #LFACE\n         468, 469, 470, 471, 472, 473, 474, 475, 476, 477, 478, 479, 480, 481, 482, 483, 484, 485, 486, 487, 488, #LHAND\n         493, 494, 495, 497, 499, 501, 503, 505, 507, 509, 511, 513, #LPOSE\n         515, 517, 519, 521, #LLEG\n         ]\n\nRIGHT = [\n         3, 7, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34, 35, 36, 37, 38,\n         39, 40, 41, 42, 43, 44, 45, 46, 47, 48, 49, 50, 51, 52, 53, 54, 55, 56, 57, 58, 59,\n         60, 61, 62, 63, 64, 65, 66, 67, 68, 69, 70, 71, 72, 73, 74, 75, 76, 77, 78, 79, 80,\n         81, 82, 83, 84, 85, 86, 87, 88, 89, 90, 91, 92, 93, 95, 96, 97, 98, 99, 100, 101, 102,\n         103, 104, 105, 106, 107, 108, 109, 110, 111, 112, 113, 114, 115, 116, 117, 118, 119, 120,\n         121, 122, 123, 124, 125, 126, 127, 128, 129, 130, 131, 132, 133, 134, 135, 136, 137, 138,\n         139, 140, 141, 142, 143, 144, 145, 146, 147, 148, 149, 150, 153, 154, 155, 156, 157, 158,\n         159, 160, 161, 162, 163, 165, 166, 167, 169, 170, 171, 172, 173, 174, 176, 177, 178, 179,\n         180, 181, 182, 183, 184, 185, 186, 187, 188, 189, 190, 191, 192, 193, 194, 196, 198, 201,\n         202, 203, 204, 205, 206, 207, 208, 209, 210, 211, 212, 213, 214, 215, 216, 217, 218, 219,\n         220, 221, 222, 223, 224, 225, 226, 227, 228, 229, 230, 231, 232, 233, 234, 235, 236, 237,\n         238, 239, 240, 241, 242, 243, 244, 245, 246, 247, #RFACE\n        522, 523, 524, 525, 526, 527, 528, 529, 530, 531, 532, 533, 534, 535, 536, 537, 538, 539, 540, 541, 542, #RHAND\n        490, 491, 492, 496, 498, 500, 502, 504, 506, 508, 510, 512, #RPOSE\n        514, 516, 518, 520, #RLEG\n        ]\n\nCENTRE = [\n          0, 1, 2, 4, 5, 6, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 94, 151, 152, 164, 168, 175, 195, 197, 199, 200, #FACE\n          489, #POSE\n          ]\n\nprint(len(LEFT+RIGHT+CENTRE))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROWS_PER_FRAME = 543\nMAX_LEN = 384\nCROP_LEN = MAX_LEN\nNUM_CLASSES  = len(NUM_TO_CHAR.values()) #62\nPAD = -100.\n\nLHAND = np.arange(468, 489).tolist()\nRHAND = np.arange(522, 543).tolist()\nPOINT_LANDMARKS = list(range(543))\n\nNUM_NODES = len(POINT_LANDMARKS)\nCHANNELS = 3*NUM_NODES\n\nprint(NUM_NODES)\nprint(CHANNELS)\n\ndef interp1d_(x, target_len, method='random'):\n    length = tf.shape(x)[1]\n    target_len = tf.maximum(1,target_len)\n    if method == 'random':\n        if tf.random.uniform(()) < 0.33:\n            x = tf.image.resize(x, (target_len,tf.shape(x)[1]),'bilinear')\n        else:\n            if tf.random.uniform(()) < 0.5:\n                x = tf.image.resize(x, (target_len,tf.shape(x)[1]),'bicubic')\n            else:\n                x = tf.image.resize(x, (target_len,tf.shape(x)[1]),'nearest')\n    else:\n        x = tf.image.resize(x, (target_len,tf.shape(x)[1]),method)\n    return x\n\ndef tf_nan_mean(x, axis=0, keepdims=False):\n    return tf.reduce_sum(tf.where(tf.math.is_nan(x), tf.zeros_like(x), x), axis=axis, keepdims=keepdims) / tf.reduce_sum(tf.where(tf.math.is_nan(x), tf.zeros_like(x), tf.ones_like(x)), axis=axis, keepdims=keepdims)\n\ndef tf_nan_std(x, center=None, axis=0, keepdims=False):\n    if center is None:\n        center = tf_nan_mean(x, axis=axis,  keepdims=True)\n    d = x - center\n    return tf.math.sqrt(tf_nan_mean(d * d, axis=axis, keepdims=keepdims))\n\ndef is_left_handed(x):\n    lhand = tf.gather(x, LHAND, axis=1)\n    rhand = tf.gather(x, RHAND, axis=1)\n    lhand_nans = tf.reduce_sum(tf.cast(tf.math.is_nan(lhand), tf.int32))\n    rhand_nans = tf.reduce_sum(tf.cast(tf.math.is_nan(rhand), tf.int32))\n    return lhand_nans < rhand_nans\n\ndef flip_lr(x, left=LEFT, right=RIGHT):\n    x,y,z = tf.unstack(x, axis=-1)\n    x = 1-x\n    new_x = tf.stack([x,y,z], -1)\n    new_x = tf.transpose(new_x, [1,0,2])\n    l_x = tf.gather(new_x, left, axis=0)\n    r_x = tf.gather(new_x, right, axis=0)\n    new_x = tf.tensor_scatter_nd_update(new_x, tf.constant(left)[...,None], r_x)\n    new_x = tf.tensor_scatter_nd_update(new_x, tf.constant(right)[...,None], l_x)\n    new_x = tf.transpose(new_x, [1,0,2])\n    return new_x\n\nclass Preprocess(tf.keras.layers.Layer):\n    def __init__(self, max_len=MAX_LEN, point_landmarks=POINT_LANDMARKS, **kwargs):\n        super().__init__(**kwargs)\n        self.max_len = max_len\n        self.point_landmarks = point_landmarks\n\n    def call(self, inputs):\n        # if tf.rank(inputs) == 3:\n        #     x = inputs[None,...]\n        # else:\n        #     x = inputs\n        x = inputs\n        x = filter_nans_tf(x)\n        x = tf.cond(is_left_handed(x), lambda:flip_lr(x), lambda:x)\n        x = x[None,...]\n\n        if self.max_len is not None:\n            x = x[:,:self.max_len]\n        length = tf.shape(x)[1]\n\n        mean = tf_nan_mean(tf.gather(x, self.point_landmarks, axis=2), axis=[1,2], keepdims=True)\n        mean = tf.where(tf.math.is_nan(mean), tf.constant([0.5,0.5,0.],x.dtype), mean)\n        x = tf.gather(x, self.point_landmarks, axis=2) #N,T,P,C\n        std = tf_nan_std(x, center=mean, axis=[1,2], keepdims=True)\n\n        x = (x - mean)/std\n\n        x = tf.concat([\n            tf.reshape(x, (-1,length,3*len(self.point_landmarks))),\n            # tf.reshape(dx, (-1,length,3*len(self.point_landmarks))),\n        ], axis = -1)\n\n        x = tf.where(tf.math.is_nan(x),tf.constant(0.,x.dtype),x)\n\n        return x","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def decode_tfrec(record_bytes):\n    features = tf.io.parse_single_example(record_bytes, {\n        'coordinates': tf.io.FixedLenFeature([], tf.string),\n        'phrase_encoded': tf.io.VarLenFeature(dtype=tf.int64),\n        'phrase': tf.io.FixedLenFeature([], tf.string),\n    })\n    out = {}\n    out['coordinates']  = tf.transpose(tf.reshape(tf.io.decode_raw(features['coordinates'], tf.float32), (-1,3,ROWS_PER_FRAME)), (0,2,1))\n    out['phrase'] = features['phrase']\n    return out\n\ndef filter_nans_tf(x, ref_point=POINT_LANDMARKS):\n    mask = tf.math.logical_not(tf.reduce_all(tf.math.is_nan(tf.gather(x,ref_point,axis=1)), axis=[-2,-1]))\n    x = tf.boolean_mask(x, mask, axis=0)\n    return x\n\ndef preprocess(x, augment=False, max_len=MAX_LEN):\n    coord = x['coordinates']\n    if augment:\n        coord = augment_fn(coord, max_len=max_len)\n    coord = tf.ensure_shape(coord, (None,ROWS_PER_FRAME,3))\n\n    inp = tf.cast(Preprocess(max_len=max_len)(coord)[0],tf.float32)\n    tar = preprocess_phrase(x['phrase'])\n\n    return inp, tar\n\ndef augment_phrase(phrase):\n    phrase = keras_nlp.layers.MaskedLMMaskGenerator(NUM_CLASSES-2,\n                                              mask_selection_rate=0.2,\n                                              mask_token_id=0,\n                                              mask_token_rate=0,\n                                              random_token_rate=1,\n                                              unselectable_token_ids=[0,60,61])(phrase)['token_ids']\n    return phrase\n\ndef is_empty(*args):\n    return tf.shape(args[0])[0] > 1\n\ndef resample(x, rate=(0.8,1.2)):\n    rate = tf.random.uniform((), rate[0], rate[1])\n    length = tf.shape(x)[0]\n    new_size = tf.cast(rate*tf.cast(length,tf.float32), tf.int32)\n    new_x = interp1d_(x, new_size)\n    return new_x\n\ndef spatial_random_affine(xyz,\n    scale  = (0.8,1.2),\n    shear = (-0.15,0.15),\n    shift  = (-0.1,0.1),\n    degree = (-30,30),\n):\n    center = tf.constant([0.5,0.5])\n    if scale is not None:\n        scale = tf.random.uniform((),*scale)\n        xyz = scale*xyz\n\n    if shear is not None:\n        xy = xyz[...,:2]\n        z = xyz[...,2:]\n        shear_x = shear_y = tf.random.uniform((),*shear)\n        if tf.random.uniform(()) < 0.5:\n            shear_x = 0.\n        else:\n            shear_y = 0.\n        shear_mat = tf.identity([\n            [1.,shear_x],\n            [shear_y,1.]\n        ])\n        xy = xy @ shear_mat\n        center = center + [shear_y, shear_x]\n        xyz = tf.concat([xy,z], axis=-1)\n\n    if degree is not None:\n        xy = xyz[...,:2]\n        z = xyz[...,2:]\n        xy -= center\n        degree = tf.random.uniform((),*degree)\n        radian = degree/180*np.pi\n        c = tf.math.cos(radian)\n        s = tf.math.sin(radian)\n        rotate_mat = tf.identity([\n            [c,s],\n            [-s, c],\n        ])\n        xy = xy @ rotate_mat\n        xy = xy + center\n        xyz = tf.concat([xy,z], axis=-1)\n\n    if shift is not None:\n        shift = tf.random.uniform((),*shift)\n        xyz = xyz + shift\n\n    return xyz\n\ndef temporal_crop(x, length=MAX_LEN):\n    l = tf.shape(x)[0]\n    offset = tf.random.uniform((), 0, tf.clip_by_value(l-length,1,length), dtype=tf.int32)\n    x = x[offset:offset+length]\n    return x\n\ndef temporal_mask(x, size=(0.2,0.4), mask_value=float('nan')):\n    l = tf.shape(x)[0]\n    mask_size = tf.random.uniform((), *size)\n    mask_size = tf.cast(tf.cast(l, tf.float32) * mask_size, tf.int32)\n    mask_offset = tf.random.uniform((), 0, tf.clip_by_value(l-mask_size,1,l), dtype=tf.int32)\n    x = tf.tensor_scatter_nd_update(x,tf.range(mask_offset, mask_offset+mask_size)[...,None],tf.fill([mask_size,543,3],mask_value))\n    return x\n\ndef spatial_mask(x, size=(0.2,0.4), mask_value=float('nan')):\n    mask_offset_y = tf.random.uniform(())\n    mask_offset_x = tf.random.uniform(())\n    mask_size = tf.random.uniform((), *size)\n    mask_x = (mask_offset_x<x[...,0]) & (x[...,0] < mask_offset_x + mask_size)\n    mask_y = (mask_offset_y<x[...,1]) & (x[...,1] < mask_offset_y + mask_size)\n    mask = mask_x & mask_y\n    x = tf.where(mask[...,None], mask_value, x)\n    return x\n\ndef augment_fn(x, always=False, max_len=None):\n    if tf.random.uniform(())<0.8 or always:\n        x = resample(x, (0.5,1.5))\n    # if tf.random.uniform(())<0.5 or always:\n    #     x = flip_lr(x)\n    # if max_len is not None:\n    #     x = temporal_crop(x, max_len)\n    if tf.random.uniform(())<0.75 or always:\n        x = spatial_random_affine(x)\n    # if tf.random.uniform(())<0.5 or always:\n    #     x = temporal_mask(x)\n    if tf.random.uniform(())<0.5 or always:\n        x = spatial_mask(x)\n    return x\n\ndef get_tfrec_dataset(tfrecords, batch_size=64, max_len=128, target_len=64, teacher_forcing=True, drop_remainder=False, train=False, augment=False, shuffle=False, repeat=False):\n    # Initialize dataset with TFRecords\n    ds = tf.data.TFRecordDataset(tfrecords, num_parallel_reads=tf.data.AUTOTUNE, compression_type='GZIP')\n    ds = ds.map(decode_tfrec, tf.data.AUTOTUNE)\n    ds = ds.map(lambda x: preprocess(x, augment=augment, max_len=max_len), tf.data.AUTOTUNE)\n\n    if train:\n        ds = ds.filter(is_empty)\n\n    if teacher_forcing:\n        ds = ds.map(lambda x,y:((x,y[:-1]),(y[1:],y[1:-1])), tf.data.AUTOTUNE)\n        if augment:\n            ds = ds.map(lambda x,y:((x[0],augment_phrase(x[1])),y), tf.data.AUTOTUNE)\n\n    if repeat:\n        ds = ds.repeat()\n\n    if shuffle:\n        ds = ds.shuffle(shuffle)\n        options = tf.data.Options()\n        options.experimental_deterministic = (False)\n        ds = ds.with_options(options)\n\n    if batch_size:\n        if teacher_forcing:\n            ds = ds.padded_batch(batch_size, padding_values=((PAD,0),(0,0)), padded_shapes=(([max_len,CHANNELS],[target_len,]),([target_len,],[target_len,])), drop_remainder=drop_remainder)\n        else:\n            ds = ds.padded_batch(batch_size, padding_values=(PAD,0), padded_shapes=([max_len,CHANNELS],[target_len,]), drop_remainder=drop_remainder)\n\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n\n    return ds\n\nds = get_tfrec_dataset(TRAIN_FILENAMES, train=True, augment=True, batch_size=1024, shuffle=1024)\nfor x in ds:\n    temp_train = x\n    break","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tensorflow.python.ops.gen_dataset_ops import filter_dataset_eager_fallback\nclass ECA(tf.keras.layers.Layer):\n    def __init__(self, kernel_size=5, **kwargs):\n        super().__init__(**kwargs)\n        self.supports_masking = True\n        self.kernel_size = kernel_size\n        self.conv = tf.keras.layers.Conv1D(1, kernel_size=kernel_size, strides=1, padding=\"same\", use_bias=False)\n\n    def call(self, inputs, mask=None):\n        nn = tf.keras.layers.GlobalAveragePooling1D()(inputs, mask=mask)\n        nn = tf.expand_dims(nn, -1)\n        nn = self.conv(nn)\n        nn = tf.squeeze(nn, -1)\n        nn = tf.nn.sigmoid(nn)\n        nn = nn[:,None,:]\n        return inputs * nn\n\nclass MaskingDWConv1D(tf.keras.layers.Layer):\n    '''\n    masked DW1Dconv with strides>1, padding=same.\n    NOTE: padded(masked) frames should always be at the beginning or end of the input sequence.\n    '''\n    def __init__(self, kernel_size, strides=1,\n        dilation_rate=1,\n        padding='same',\n        use_bias=False,\n        kernel_initializer='glorot_uniform',**kwargs):\n        super().__init__(**kwargs)\n        assert padding == 'same' or padding == 'causal'\n        self.strides = strides\n        self.kernel_size = kernel_size\n        self.dilation_rate = dilation_rate\n        self.use_bias = use_bias\n        self.padding = padding\n        self.conv = tf.keras.layers.DepthwiseConv1D(\n                            kernel_size,\n                            strides=strides,\n                            dilation_rate=dilation_rate,\n                            padding=padding,\n                            use_bias=use_bias,\n                            kernel_initializer=kernel_initializer)\n        self.supports_masking = True\n\n    def compute_mask(self, inputs, mask=None):\n      if mask is not None:\n        if self.strides > 1:\n          mask = mask[:,::self.strides]\n      return mask\n\n    def call(self, inputs, mask=None):\n        x = inputs\n        if mask is not None:\n            x = tf.where(mask[...,None], x, tf.constant(0., dtype=x.dtype))\n        x = self.conv(x)\n        return x\n\ndef Conv1DBlock(channel_size,\n          kernel_size,\n          dilation_rate=1,\n          strides=1,\n          drop_rate=0.0,\n          expand_ratio=2,\n          activation='swish',\n          name=None):\n    '''\n    efficient conv1d block, @hoyso48\n    '''\n    if name is None:\n        name = str(tf.keras.backend.get_uid(\"conv1dblock\"))\n    # Expansion phase\n    def apply(inputs):\n        channels_in = tf.keras.backend.int_shape(inputs)[-1]\n        channels_expand = channels_in * expand_ratio\n\n        skip = inputs\n\n        x = tf.keras.layers.BatchNormalization(momentum=0.95, name=name + 'pre_bn')(inputs)\n\n        x = tf.keras.layers.Dense(\n            channels_expand,\n            use_bias=True,\n            activation=activation,\n            name=name + '_expand_conv')(x)\n\n        # Depthwise Convolution\n        x = MaskingDWConv1D(kernel_size,\n            dilation_rate=dilation_rate,\n            strides=strides,\n            use_bias=False,\n            name=name + '_dwconv')(x)\n\n        x = tf.keras.layers.BatchNormalization(momentum=0.95, name=name + 'conv_bn')(x)\n\n        x = ECA()(x)\n\n        x = tf.keras.layers.Dense(\n            channel_size,\n            use_bias=True,\n            name=name + '_project_conv')(x)\n\n        if drop_rate > 0:\n            x = tf.keras.layers.Dropout(drop_rate, noise_shape=(None,1,1), name=name + '_drop')(x)\n\n        if (channels_in == channel_size) and (strides == 1):\n            x = tf.keras.layers.add([x, skip], name=name + '_add')\n        return x\n\n    return apply","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PosEmbedding(tf.keras.layers.Layer):\n    def __init__(self, dim=64, max_len=64, **kwargs):\n        super().__init__(**kwargs)\n        self.pos_emb = tf.keras.layers.Embedding(input_dim=max_len, output_dim=dim)\n        self.supports_masking = True\n\n    def call(self, x, positions=None):\n        if positions is None:\n            maxlen = tf.shape(x)[1]\n            positions = tf.range(start=0, limit=maxlen, delta=1)\n        positions = self.pos_emb(positions)\n        return x + positions\n\nclass MultiHeadAttention(tf.keras.layers.Layer):\n    def __init__(self, dim=256, num_heads=4, dropout=0, **kwargs):\n        super().__init__(**kwargs)\n        self.dim = dim\n        self.scale = self.dim ** -0.5\n        self.num_heads = num_heads\n        self.q = tf.keras.layers.Dense(dim, use_bias=False)\n        self.k = tf.keras.layers.Dense(dim, use_bias=False)\n        self.v = tf.keras.layers.Dense(dim, use_bias=False)\n        self.drop1 = tf.keras.layers.Dropout(dropout)\n        self.proj = tf.keras.layers.Dense(dim, use_bias=False)\n        self.supports_masking = True\n\n    def get_causal_mask(self, q, k):\n        q_len = tf.shape(q)[1]\n        k_len = tf.shape(k)[1]\n        i = tf.range(q_len)[:, None]\n        j = tf.range(k_len)\n        mask = i >= j\n        mask = tf.reshape(mask, (q_len, k_len))\n        return mask\n\n    def merge_input_state(self, input, state, layer):\n        if input is not None and state is not None:\n            return tf.keras.layers.Concatenate(axis=1)([state, layer(input)])\n        elif input is not None and state is None:\n            return layer(input)\n        elif input is None and state is not None:\n            return state\n        else:\n            raise ValueError\n\n    def call(self, q, k=None, v=None, key_state=None, value_state=None, return_states=False, use_causal_mask=False):\n        q = self.q(q)\n        k = self.merge_input_state(k, key_state, self.k)\n        v = self.merge_input_state(v, value_state, self.v)\n        mask = getattr(k, '_keras_mask', None) # we only consider mask from the 'key' here.\n        if mask is not None:\n            mask = mask[:,None,None,:]\n        if use_causal_mask:\n            if mask is not None:\n                mask = tf.logical_and(mask, self.get_causal_mask(q,k)[None,None,:,:])\n            else:\n                mask = self.get_causal_mask(q,k)[None,None,:,:]\n        q_ = tf.keras.layers.Permute((2, 1, 3))(tf.keras.layers.Reshape((-1, self.num_heads, self.dim // self.num_heads))(q))\n        k_ = tf.keras.layers.Permute((2, 1, 3))(tf.keras.layers.Reshape((-1, self.num_heads, self.dim // self.num_heads))(k))\n        v_ = tf.keras.layers.Permute((2, 1, 3))(tf.keras.layers.Reshape((-1, self.num_heads, self.dim // self.num_heads))(v))\n        attn = tf.matmul(q_, k_, transpose_b=True) * self.scale\n\n        attn = tf.keras.layers.Softmax(axis=-1)(attn, mask=mask)\n        attn = self.drop1(attn)\n\n        x = attn @ v_\n        x = tf.keras.layers.Reshape((-1, self.dim))(tf.keras.layers.Permute((2, 1, 3))(x))\n        x = self.proj(x)\n        if return_states:\n            return x, k, v\n        else:\n            return x\n\ndef TransformerDecoderBlock(dim=256, num_heads=4, expand=4, attn_dropout=0.2, drop_rate=0., activation='swish', name=None):\n    if name is None:\n        name = str(tf.keras.backend.get_uid(\"transformerdecoderblock\"))\n    def apply(q,k,v):\n        x = q\n        x = tf.keras.layers.BatchNormalization(momentum=0.95, name=name + '_bn1')(x)\n        x = MultiHeadAttention(dim=dim,num_heads=num_heads,dropout=attn_dropout, name=name + '_self_attn')(x,x,x,use_causal_mask=True)\n        x = tf.keras.layers.Dropout(drop_rate, noise_shape=(None,1,1), name=name + '_drop1')(x)\n        x = tf.keras.layers.Add(name=name + '_add1')([q, x])\n        attn_out1 = x\n\n        x = tf.keras.layers.BatchNormalization(momentum=0.95, name=name + '_bn2')(x)\n        x = MultiHeadAttention(dim=dim,num_heads=num_heads,dropout=attn_dropout, name=name + '_cross_attn')(x,k,v)\n        x = tf.keras.layers.Dropout(drop_rate, noise_shape=(None,1,1), name=name + '_drop2')(x)\n        x = tf.keras.layers.Add(name=name + '_add2')([attn_out1, x])\n        attn_out2 = x\n\n        x = tf.keras.layers.BatchNormalization(momentum=0.95, name=name + '_bn3')(x)\n        x = tf.keras.layers.Dense(dim*expand, use_bias=False, activation=activation, name=name + '_fc1')(x)\n        x = tf.keras.layers.Dense(dim, use_bias=False, name=name + '_fc2')(x)\n        x = tf.keras.layers.Dropout(drop_rate, noise_shape=(None,1,1), name=name + '_drop3')(x)\n        x = tf.keras.layers.Add(name=name + '_add3')([attn_out2, x])\n        return x\n    return apply\n\n\nclass MultiHeadSelfAttention(tf.keras.layers.Layer):\n    def __init__(self, dim=256, num_heads=4, dropout=0, **kwargs):\n        super().__init__(**kwargs)\n        self.dim = dim\n        self.scale = self.dim ** -0.5\n        self.num_heads = num_heads\n        self.qkv = tf.keras.layers.Dense(3 * dim, use_bias=False)\n        self.drop1 = tf.keras.layers.Dropout(dropout)\n        self.proj = tf.keras.layers.Dense(dim, use_bias=False)\n        self.supports_masking = True\n\n    def call(self, inputs, mask=None):\n        qkv = self.qkv(inputs)\n        qkv = tf.keras.layers.Permute((2, 1, 3))(tf.keras.layers.Reshape((-1, self.num_heads, self.dim * 3 // self.num_heads))(qkv))\n        q, k, v = tf.split(qkv, [self.dim // self.num_heads] * 3, axis=-1)\n\n        attn = tf.matmul(q, k, transpose_b=True) * self.scale\n\n        if mask is not None:\n            mask = mask[:, None, None, :]\n\n        attn = tf.keras.layers.Softmax(axis=-1)(attn, mask=mask)\n        attn = self.drop1(attn)\n\n        x = attn @ v\n        x = tf.keras.layers.Reshape((-1, self.dim))(tf.keras.layers.Permute((2, 1, 3))(x))\n        x = self.proj(x)\n        return x\n\ndef TransformerBlock(dim=256, num_heads=4, expand=4, attn_dropout=0.2, drop_rate=0.2, activation='swish', name=None):\n    if name is None:\n        name = str(tf.keras.backend.get_uid(\"transformerblock\"))\n    def apply(inputs):\n        x = inputs\n        x = tf.keras.layers.BatchNormalization(momentum=0.95, name=name + '_bn1')(x)\n        x = MultiHeadSelfAttention(dim=dim,num_heads=num_heads,dropout=attn_dropout, name=name + '_mhsa')(x)\n        x = tf.keras.layers.Dropout(drop_rate, noise_shape=(None,1,1), name=name + '_drop1')(x)\n        x = tf.keras.layers.Add(name=name + '_add1')([inputs, x])\n        attn_out = x\n\n        x = tf.keras.layers.BatchNormalization(momentum=0.95, name=name + 'bn2')(x)\n        x = tf.keras.layers.Dense(dim*expand, use_bias=False, activation=activation, name=name + '_fc1')(x)\n        x = tf.keras.layers.Dense(dim, use_bias=False, name=name + '_fc2')(x)\n        x = tf.keras.layers.Dropout(drop_rate, noise_shape=(None,1,1), name=name + '_drop2')(x)\n        x = tf.keras.layers.Add(name=name + '_add2')([attn_out, x])\n        return x\n    return apply","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CTCLoss(tf.keras.losses.Loss):\n    def __init__(self, blank_index=0, input_padding_value=0., target_padding_value=0, **kwargs):\n        super().__init__(**kwargs)\n        self.blank_index = blank_index\n        self.input_padding_value = input_padding_value\n        self.target_padding_value = target_padding_value\n\n    def call(self, y_true, y_pred):\n        y_true = tf.cast(y_true, tf.int32)\n        y_pred = tf.cast(y_pred, tf.float32)\n        batch_len = tf.cast(tf.shape(y_true)[0], dtype=tf.int32)\n        label_length = y_true != tf.cast(self.target_padding_value, tf.int32)\n        label_length = tf.reduce_sum(tf.cast(label_length, tf.int32), axis=1, keepdims=False) #(B,)\n        mask = getattr(y_pred, '_keras_mask', None)\n        if mask is not None:\n            input_length = tf.reduce_sum(tf.cast(mask, tf.int32), axis=-1)\n        else:\n            input_length = tf.cast(tf.shape(y_pred)[1], dtype=tf.int32)\n            input_length = input_length * tf.ones(shape=(batch_len,), dtype=tf.int32)\n\n#         loss = tf.nn.ctc_loss(y_true, y_pred, label_length=label_length, logit_length=input_length, blank_index=0, logits_time_major=False)\n        loss = classic_ctc_loss(y_true, y_pred, label_length=label_length, logit_length=input_length, blank_index=0) #only for the kaggle TPU\n\n        loss = tf.reduce_mean(loss)\n\n        return loss\n\nclass MaskedSCCE(tf.keras.losses.Loss):\n    def __init__(self, num_classes=NUM_CLASSES, from_logits=True, label_smoothing=0.25, **kwargs):\n        super().__init__(**kwargs)\n        self.num_classes = num_classes\n        self.label_smoothing=label_smoothing\n        self.from_logits = from_logits\n\n    def call(self, y_true, y_pred):\n        mask = y_true!=0\n        N = tf.shape(y_true)[0]\n        y_pred = tf.cast(y_pred, tf.float32)\n        y_true = tf.one_hot(y_true, self.num_classes, axis=-1, dtype=tf.float32)\n        loss = tf.keras.losses.categorical_crossentropy(y_true, y_pred, from_logits=self.from_logits, label_smoothing=self.label_smoothing)\n        loss = tf.where(mask, loss, tf.constant(0, dtype=tf.float32))\n        loss = tf.reduce_sum(loss)\n        loss = loss / tf.cast(N, tf.float32)\n        return loss\n\nclass Accuracy(tf.keras.metrics.Metric):\n    def __init__(self, **kwargs):\n        super(Accuracy, self).__init__(name=f'acc', **kwargs)\n        self.acc = tf.keras.metrics.SparseCategoricalAccuracy()\n\n    def update_state(self, y_true, y_pred, sample_weight=None):\n        y_true = tf.reshape(y_true, [-1])\n        y_pred = tf.reshape(y_pred, [-1, tf.shape(y_pred)[-1]])\n        mask = y_true != 0\n        y_true = tf.boolean_mask(y_true, mask)\n        y_pred = tf.boolean_mask(y_pred, mask)\n        self.acc.update_state(y_true, y_pred)\n\n    def result(self):\n        return self.acc.result()\n\n    def reset_state(self):\n        self.acc.reset_state()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model(max_len=128, target_len=64, dim=192, dtype='float32'):\n    ################# ENCODER #################\n    inp1 = tf.keras.Input((max_len,CHANNELS),dtype=dtype)\n    x = tf.keras.layers.Masking(mask_value=PAD,input_shape=(max_len,CHANNELS))(inp1)\n    ksize = 17\n    drop_rate = 0.2\n    x = tf.keras.layers.Dense(dim,use_bias=False,name='stem_conv')(x)\n    x = Conv1DBlock(dim,ksize,expand_ratio=4,drop_rate=drop_rate)(x)\n    x = Conv1DBlock(dim,ksize,expand_ratio=4,drop_rate=drop_rate)(x)\n    x = Conv1DBlock(dim,ksize,expand_ratio=4,drop_rate=drop_rate)(x)\n    x = TransformerBlock(dim,expand=2,num_heads=4,drop_rate=drop_rate,attn_dropout=0.2)(x)\n    x = Conv1DBlock(dim,ksize,expand_ratio=4,drop_rate=drop_rate)(x)\n    x = Conv1DBlock(dim,ksize,expand_ratio=4,drop_rate=drop_rate)(x)\n    x = Conv1DBlock(dim,ksize,expand_ratio=4,drop_rate=drop_rate)(x)\n    x = TransformerBlock(dim,expand=2,num_heads=4,drop_rate=drop_rate,attn_dropout=0.2)(x)\n    x = Conv1DBlock(dim,ksize,expand_ratio=4,drop_rate=0,strides=2)(x) #drop_rate=0 since we don't want to drop the whole output here\n    x = Conv1DBlock(dim,ksize,expand_ratio=4,drop_rate=drop_rate)(x)\n    x = Conv1DBlock(dim,ksize,expand_ratio=4,drop_rate=drop_rate)(x)\n    x = Conv1DBlock(dim,ksize,expand_ratio=4,drop_rate=drop_rate)(x)\n    x = TransformerBlock(dim,expand=2,num_heads=4,drop_rate=drop_rate,attn_dropout=0.2)(x)\n    x = Conv1DBlock(dim,ksize,expand_ratio=4,drop_rate=drop_rate)(x)\n    x = Conv1DBlock(dim,ksize,expand_ratio=4,drop_rate=drop_rate)(x)\n    x = Conv1DBlock(dim,ksize,expand_ratio=4,drop_rate=drop_rate)(x)\n    x = TransformerBlock(dim,expand=2,num_heads=4,drop_rate=drop_rate,attn_dropout=0.2)(x)\n    x = tf.keras.layers.BatchNormalization(momentum=0.95)(x)\n\n    encoder = tf.keras.Model(inp1,x,name='encoder')\n\n    ################# CTC DECDODER #################\n    inp3 = tf.keras.Input((x.shape[1],dim),name='ctc_decoder_inp2',dtype=dtype)\n    x = inp3\n    x = tf.keras.layers.RNN(tf.keras.layers.GRUCell(dim), return_sequences=True)(x)\n    x = tf.keras.layers.Dense(dim*2)(x)\n    x = tf.keras.layers.Dropout(0.5)(x)\n    x = tf.keras.layers.Dense(NUM_CLASSES,name='ctc_classifier')(x) #include sos, eos token\n    ctc_decoder = tf.keras.Model(inp3,x,name='ctc_decoder')\n\n    ################# ATT DECODER #################\n    inp2 = tf.keras.Input((None,),name='att_decoder_inp1',dtype='int32')\n    inp3 = tf.keras.Input((x.shape[1],dim),name='att_decoder_inp2',dtype=dtype)\n\n    x = inp3\n    y = tf.keras.layers.Masking(mask_value=0,input_shape=(None,),name='att_decoder_input_masking')(inp2)\n    y = tf.keras.layers.Embedding(NUM_CLASSES,dim,mask_zero=True,name='att_decoder_token_emb')(y) #include sos token\n    y = PosEmbedding(dim,max_len=target_len,name='att_decoder_pos_emb')(y)\n    y = TransformerDecoderBlock(dim,expand=2,num_heads=4,attn_dropout=0.2,name='att_decoder_block1')(y,x,x)\n    y = tf.keras.layers.Dropout(0.5)(y)\n    y = tf.keras.layers.Dense(NUM_CLASSES,name='att_decoder_classifier')(y)\n\n    decoder = tf.keras.Model([inp2,inp3],y,name='att_decoder')\n\n    ################### MODEL #####################\n    inp1 = tf.keras.Input((max_len,CHANNELS),dtype=dtype)\n    inp2 = tf.keras.Input((None,),dtype='int32')\n\n    x = inp1\n    enc_out = encoder(x)\n    y = inp2\n    dec_out = decoder([y, enc_out])\n    ctc_out = ctc_decoder(enc_out)\n    model = tf.keras.Model([inp1,inp2], [dec_out,ctc_out])\n\n    return model\n\nmodel = get_model()\ny = model(temp_train[0], training=True)\nmodel.summary()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#check supports_masking\nfor x in model.layers:\n    if not x.supports_masking:\n        print(x.supports_masking, x.name)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GreedyDecoder(tf.keras.layers.Layer):\n    def __init__(self, model, max_output_length=64, sos_token_idx=60, eos_token_idx=61, pad_token_idx=0, **kwargs):\n        super().__init__(**kwargs)\n        self.model = model\n        self.encoder = self.model.get_layer('encoder')\n        self.decoder = self.model.get_layer('att_decoder')\n        self.inference_module = self.model.get_layer('att_decoder')\n        self.max_output_length = max_output_length\n        self.sos_token_idx = sos_token_idx\n        self.eos_token_idx = eos_token_idx\n        self.pad_token_idx = pad_token_idx\n\n    def call(self, batch_x):\n        encoder_out = self.encoder(batch_x)\n\n        time = tf.constant(0, dtype=tf.int32)\n        predictions = tf.ones((tf.shape(batch_x)[0],1), dtype=tf.int32) * self.sos_token_idx\n        pad = tf.ones((tf.shape(batch_x)[0],), dtype=tf.int32) * self.pad_token_idx\n        init = True\n\n        def condition(_time, _predictions):\n            return tf.logical_and(_time < self.max_output_length, tf.logical_not(tf.reduce_all(tf.reduce_any(_predictions==self.eos_token_idx, axis=1))))\n\n        def body(_time, _predictions):\n            out = self.inference_module([_predictions, encoder_out])\n            pred_curr = tf.where(tf.reduce_any(_predictions==self.eos_token_idx, axis=1), [self.pad_token_idx], tf.argmax(out[:,-1], axis=-1, output_type=tf.int32))\n            _predictions = tf.concat([_predictions, pred_curr[...,None]], axis=1)\n            return _time+1, _predictions\n\n        _, predictions = tf.while_loop(condition, body, loop_vars=[time, predictions])\n        return predictions[:,1:]\n\nclass KerasCTCDecoder(tf.keras.layers.Layer):\n    def __init__(self, model, greedy=True, beam_width=100, from_logits=True, **kwargs):\n        super().__init__(**kwargs)\n        self.model = model\n        self.greedy = greedy\n        self.beam_width = beam_width\n        self.from_logits = from_logits\n        self.encoder = self.model.get_layer('encoder')\n        self.ctc_decoder = self.model.get_layer('ctc_decoder')\n\n    def call(self, batch_x):\n        encoder_out = self.encoder(batch_x)\n        input_length = tf.reduce_sum(tf.cast(encoder_out._keras_mask, tf.int32), axis=1)\n        predictions = self.ctc_decoder(encoder_out)\n        if not self.greedy and self.from_logits:\n            predictions = tf.nn.softmax(predictions, axis=-1)\n        predictions = tf.keras.backend.ctc_decode(tf.cast(predictions, tf.float32), input_length=input_length, greedy=self.greedy, beam_width=self.beam_width)[0][0]\n        return predictions","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_predictions(recognizer, ds):\n    results = []\n    for batch in tqdm(ds):\n        result = recognizer(batch[0][0])\n        results.append(num_to_char(result.numpy()))\n    results = np.array([item for sublist in results for item in sublist])\n    return results\n\ndef num_to_char(list_of_nums, n2c_dict=NUM_TO_CHAR):\n    def n_to_c(x):\n        return [n2c_dict[a] for a in x if a!=-1]\n    char_list = [''.join(n_to_c(x)).replace('P','').replace('S','').replace('E','') for x in list_of_nums]\n    return np.array(char_list, dtype='str')\n\ndef extract_labels(ds):\n    labels = [num_to_char(x[1][0].numpy()) for x in ds]\n    labels = np.array([item for sublist in labels for item in sublist])\n    return labels\n\nfrom Levenshtein import distance\ndef competition_metric(true, pred):\n    #true: list of strings, ground truths\n    #pred: list of strings, predictions\n    D = sum([distance(x,y) for x,y in zip(true, pred)])\n    N = len(''.join(true))\n    return max((N-D)/N, 0.), D/len(true)\n\ndef display(labels, preds):\n    for target,prediction in zip(labels, preds):\n        print(f\"Target    : {target}\")\n        print(f\"Prediction: {prediction}\")\n        print(\"-\" * 100)\n    return\n\ndef evaluate(model, ds, labels=None, display_index='random', num_display=5):\n    if labels is None:\n        labels = extract_labels(ds)\n    preds = make_predictions(model, ds)\n    score, mean_dist = competition_metric(labels, preds)\n    num_display = min(len(labels), num_display)\n    if display_index=='random':\n        if num_display:\n            idxs = np.random.choice(range(len(labels)),num_display,replace=False)\n            display(labels[idxs], preds[idxs])\n    elif display_index=='init':\n        if num_display:\n            display(labels[:num_display], preds[:num_display])\n    elif isinstance(display_index, list):\n        if display_index:\n            display(labels[display_index], preds[display_index])\n    else:\n        pass\n    print(f'Score: {score:0.4f}')\n    print(f'mean_dist: {mean_dist:0.4f}')\n    # return labels, preds, score\n    del preds, score\n    return\n\nclass Eval(tf.keras.callbacks.Callback):\n    def __init__(self,recognizer,ds,labels=None,eval_epochs=[],display_index='random',num_display=5):\n        super().__init__()\n        self.recognizer = recognizer\n        self.ds = ds\n        self.labels = labels\n        self.eval_epochs = eval_epochs\n        self.display_index = display_index\n        self.num_display = num_display\n\n    def on_epoch_end(self, epoch, logs=None):\n        if epoch in self.eval_epochs and self.ds is not None: # your custom condition\n            evaluate(self.recognizer, self.ds, self.labels, self.display_index, self.num_display)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AWP(tf.keras.Model):\n    def __init__(self, *args, lr=0.1, eps=1e-6, start_step=0, exclude=[], **kwargs):\n        super().__init__(*args, **kwargs)\n        self.lr = lr\n        self.eps = eps\n        self.start_step = start_step\n        self.exclude = exclude\n\n    def compute_perturbation(self, param, param_gradient):\n        grad = tf.zeros_like(param) + param_gradient\n        #delta = tf.math.divide_no_nan(self.lr * grad * tf.norm(param), tf.norm(grad) + self.eps) #original implemenation from the paper\n        delta = tf.math.divide_no_nan(self.lr * grad, tf.norm(grad) + self.eps)\n        return delta\n\n    def train_step_awp(self, data):\n        # Unpack the data. Its structure depends on your model and\n        # on what you pass to `fit()`.\n        x, y = data\n\n        with tf.GradientTape() as tape:\n            y_pred = self(x, training=True)\n            loss = self.compiled_loss(y, y_pred, regularization_losses=self.losses)\n        params = self.trainable_variables\n        params_gradients = tape.gradient(loss, self.trainable_variables)\n\n        for i in range(len(params_gradients)):\n            if not any(s in params[i].name for s in self.exclude):\n                delta = self.compute_perturbation(params[i], params_gradients[i])\n                self.trainable_variables[i].assign_add(delta)\n\n        with tf.GradientTape() as tape2:\n            y_pred = self(x, training=True)\n            new_loss = self.compiled_loss(y, y_pred, regularization_losses=self.losses)\n            if hasattr(self.optimizer, 'get_scaled_loss'):\n                new_loss = self.optimizer.get_scaled_loss(new_loss)\n\n        gradients = tape2.gradient(new_loss, self.trainable_variables)\n        if hasattr(self.optimizer, 'get_unscaled_gradients'):\n            gradients =  self.optimizer.get_unscaled_gradients(gradients)\n\n        for i in range(len(params_gradients)):\n            if not any(s in params[i].name for s in self.exclude):\n                delta = self.compute_perturbation(params[i], params_gradients[i])\n                self.trainable_variables[i].assign_sub(delta)\n\n        #if nan is detected, skip update\n        # nan_detected = tf.reduce_any([tf.reduce_any(tf.math.is_nan(g)) for g in gradients])\n        # _ = tf.cond(nan_detected, lambda:tf.constant(False),lambda:self.optimizer.apply_gradients(zip(gradients, self.trainable_variables)))\n\n        self.optimizer.apply_gradients(zip(gradients, self.trainable_variables))\n        self.compiled_metrics.update_state(y, y_pred)\n        return {m.name: m.result() for m in self.metrics}\n\n    def train_step(self, data):\n        return tf.cond(self._train_counter < self.start_step, lambda:super(AWP,self).train_step(data), lambda:self.train_step_awp(data))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CSVLoggerV2(tf.keras.callbacks.CSVLogger):\n    def __init__(self, filename, separator=\",\", resume=0):\n        self.resume = resume\n        super().__init__(filename=filename, separator=separator, append=bool(resume))\n\n    def on_epoch_end(self, epoch, logs=None):\n        super(CSVLoggerV2,self).on_epoch_end(epoch=epoch+self.resume+1, logs=logs)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_fold(CFG, fold, train_files, valid_files=None, strategy=STRATEGY, summary=True):\n    seed_everything(CFG.seed)\n    tf.keras.backend.clear_session()\n    gc.collect()\n    # tf.config.optimizer.set_jit(True)\n\n    policy = mixed_precision.Policy(CFG.policy)\n    mixed_precision.set_global_policy(policy)\n\n    if CFG.resume == 'auto':\n        if os.path.isfile(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-logs.csv'):\n            resume = pd.read_csv(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-logs.csv')['epoch'].values[-1]\n            resume = 0 if resume == CFG.epoch else resume #restart if training is already fininshed\n        else:\n            resume = 0\n    else:\n        resume = CFG.resume\n\n    if fold != 'all':\n        train_ds = get_tfrec_dataset(train_files, batch_size=CFG.train_batch_size, max_len=CFG.max_len, drop_remainder=True, train=True, augment=True, repeat=True, shuffle=4096)\n        valid_ds = get_tfrec_dataset(valid_files, batch_size=CFG.valid_batch_size, max_len=CFG.max_len, drop_remainder=True, repeat=False, shuffle=False)\n    else:\n        train_ds = get_tfrec_dataset(train_files, batch_size=CFG.train_batch_size, max_len=CFG.max_len, drop_remainder=True, train=True, augment=True, repeat=True, shuffle=4096)\n        valid_ds = None\n        valid_files = []\n\n    num_train = count_data_items(train_files)\n    num_valid = count_data_items(valid_files)\n    steps_per_epoch = num_train//CFG.train_batch_size\n    with strategy.scope():\n        model = get_model(max_len=CFG.max_len, dim=CFG.dim, dtype='bfloat16') #dtype should be matched with CFG.policy\n\n        schedule = OneCycleLR(CFG.lr, CFG.epoch, warmup_epochs=CFG.epoch*CFG.warmup, steps_per_epoch=steps_per_epoch, resume_epoch=resume, decay_epochs=CFG.epoch, lr_min=CFG.lr_min, decay_type=CFG.decay_type, warmup_type=CFG.warmup_type)\n        decay_schedule = OneCycleLR(CFG.lr*CFG.weight_decay, CFG.epoch, warmup_epochs=CFG.epoch*CFG.warmup, steps_per_epoch=steps_per_epoch, resume_epoch=resume, decay_epochs=CFG.epoch, lr_min=CFG.lr_min*CFG.weight_decay, decay_type=CFG.decay_type, warmup_type=CFG.warmup_type)\n\n        awp_start_epoch = max(CFG.awp_start_epoch - resume, 0)\n        awp_step = int(awp_start_epoch * steps_per_epoch)\n        if CFG.fgm:\n            model = FGM(model.input, model.output, lr=CFG.awp_lr, eps=0., start_step=awp_step)\n        elif CFG.awp:\n            model = AWP(model.input, model.output, lr=CFG.awp_lr, eps=0., start_step=awp_step, exclude=['bias','gamma','beta','rnn'])\n\n        # opt = tfa.optimizers.RectifiedAdam(learning_rate=schedule, weight_decay=decay_schedule, sma_threshold=4)\n        # opt = tfa.optimizers.Lookahead(opt,sync_period=5)\n        opt = tfa.optimizers.AdamW(learning_rate=schedule, weight_decay=decay_schedule)\n\n        model.compile(\n            optimizer=opt,\n            loss=[MaskedSCCE(label_smoothing=0.25), CTCLoss()],\n            loss_weights=[0.75,0.25],\n            metrics=[\n                [\n                Accuracy(),\n                ],\n                [],\n            ],\n        )\n\n    if summary:\n        print()\n        model.summary()\n        print()\n        print(train_ds, valid_ds)\n        print()\n        schedule.plot()\n        print()\n        init=False\n    print(f'---------fold{fold}---------')\n    print(f'train:{num_train} valid:{num_valid}')\n    print()\n\n    if resume:\n        print(f'resume from epoch{resume}')\n        if CFG.resume_ckpt:\n            print(f'load weights from {CFG.resume_ckpt}')\n            model.load_weights(CFG.resume_ckpt)\n        else:\n            print(f'load weights from {CFG.output_dir}/{CFG.comment}-fold{fold}-last.h5')\n            model.load_weights(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-last.h5')\n        # if train_ds is not None:\n        #     model.evaluate(train_ds.take(steps_per_epoch))\n        # if valid_ds is not None:\n        #     model.evaluate(valid_ds)\n\n    logger = CSVLoggerV2(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-logs.csv', resume=resume)\n\n    mode = 'min'\n    if fold != 'all':\n        monitor = 'val_loss'\n    else:\n        monitor = 'loss'\n    if resume:\n        prev_best = pd.read_csv(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-logs.csv')[monitor].agg(mode)\n    else:\n        prev_best = None\n    sv_loss = tf.keras.callbacks.ModelCheckpoint(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-best.h5', monitor=monitor, verbose=0, save_best_only=True,\n                  save_weights_only=True, mode='min', save_freq='epoch', initial_value_threshold=prev_best)\n\n    snap = Snapshot(f'{CFG.output_dir}/{CFG.comment}-fold{fold}', snapshot_epochs=[])\n    # swa = SWA(f'{CFG.output_dir}/{CFG.comment}-fold{fold}', CFG.swa_epochs, strategy=strategy, train_ds=train_ds, valid_ds=valid_ds)\n\n    callbacks = []\n    if CFG.save_output:\n        callbacks.append(logger)\n        callbacks.append(snap)\n        # callbacks.append(swa)\n        callbacks.append(sv_loss)\n\n    history = model.fit(\n        train_ds,\n        epochs=CFG.epoch-resume,\n        steps_per_epoch=steps_per_epoch,\n        callbacks=callbacks,\n        validation_data=valid_ds,\n        verbose=CFG.verbose,\n    )\n\n    if fold != 'all':\n        ds = get_tfrec_dataset(valid_files, batch_size=CFG.valid_batch_size, max_len=CFG.max_len, drop_remainder=False, repeat=False, shuffle=False)\n        labels = extract_labels(ds)\n        model.load_weights(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-last.h5')\n        print('ATTENTION EVAL')\n        evaluate(GreedyDecoder(model),ds,labels)\n        print()\n        print('CTC EVAL')\n        evaluate(KerasCTCDecoder(model, greedy=True),ds,labels)\n\n    return model, history\n\ndef train_folds(CFG, folds, strategy=STRATEGY, summary=True):\n    for fold in folds:\n        if fold != 'all':\n            all_files = TRAIN_FILENAMES\n            train_files = [x for x in all_files if f'fold{fold}' not in x]\n            valid_files = [x for x in all_files if f'fold{fold}' in x]\n        else:\n            train_files = TRAIN_FILENAMES\n            valid_files = None\n\n        train_fold(CFG, fold, train_files, valid_files, strategy=strategy, summary=summary)\n    return","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    n_splits = 5\n    save_output = True\n    output_dir = '.'\n\n    seed = 42\n    verbose = 'auto' #0) silent 1) progress bar 2) one line per epoch\n\n    dim = 192\n    max_len = 768\n\n    policy = 'mixed_bfloat16' #'float32') fp32, 'mixed_float16') GPU+fp16, 'mixed_bfloat16') TPU+fp16\n    replicas = N_REPLICAS\n    lr = 5e-4 * replicas\n    weight_decay = 0.01\n    lr_min = 1e-6\n    epoch = 60 #400\n    warmup = 0.1\n    warmup_type = 'linear'\n    decay_type = 'cosine'\n    train_batch_size = 16 * replicas\n    valid_batch_size = 64 * replicas\n\n    fgm = False\n    awp = False #True\n    awp_lr = 0.2\n    awp_start_epoch = 0.1 * epoch\n\n    resume = 0\n    resume_ckpt = ''\n    comment =  f'aslfr-fp16-192d-17l-ctcattjoint-seed{seed}'","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CFG.resume = 'auto'\n# CFG.resume = 134\n# CFG.resume_ckpt = f'{CFG.comment}-fold0-best.h5'\ntrain_folds(CFG, [0])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# CFG.seed = 42\n# CFG.comment = f'aslfr-fp16-192d-17l-ctcattjoint-seed{CFG.seed}'\n# train_folds(CFG, ['all'], summary=False)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# CFG.seed = 43\n# CFG.comment = f'aslfr-fp16-192d-17l-ctcattjoint-seed{CFG.seed}'\n# train_folds(CFG, ['all'], summary=False)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# CFG.seed = 44\n# CFG.comment = f'aslfr-fp16-192d-17l-ctcattjoint-seed{CFG.seed}'\n# train_folds(CFG, ['all'], summary=False)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}],"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"}}