{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":8425227,"sourceType":"datasetVersion","datasetId":5016565},{"sourceId":8425253,"sourceType":"datasetVersion","datasetId":5016563},{"sourceId":8425264,"sourceType":"datasetVersion","datasetId":5016567},{"sourceId":8425271,"sourceType":"datasetVersion","datasetId":5016566}],"dockerImageVersionId":30675,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# !pip install tensorflow -U -q\n!pip install tensorflow_probability==0.16.0 -q\n# !pip install git+https://github.com/awsaf49/tensorflow_extra -q\n# !pip install kecam -q -U\n!pip install tensorflow-addons -q\n!pip install tensorflow-io -q\n!pip install git+https://github.com/hoyso48/tf-utils@main -q -U\n!pip install keras==2.15.0 -q","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2024-05-31T16:18:17.163895Z","iopub.execute_input":"2024-05-31T16:18:17.164806Z","iopub.status.idle":"2024-05-31T16:19:41.193284Z","shell.execute_reply.started":"2024-05-31T16:18:17.164750Z","shell.execute_reply":"2024-05-31T16:19:41.192109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import Audio\nfrom sklearn.model_selection import StratifiedGroupKFold, KFold\n\nimport numpy as np\nimport pandas as pd\nimport tensorflow as tf\nimport tensorflow_addons as tfa\nimport matplotlib.pyplot as plt\nimport matplotlib as mpl\nimport tensorflow.keras.mixed_precision as mixed_precision\n# import tensorflow_extra as tfe\nimport tensorflow_io as tfio\n\nfrom tqdm.auto import tqdm\nimport sklearn\n\nfrom tf_utils.schedules import OneCycleLR, ListedLR\nfrom tf_utils.callbacks import Snapshot, SWA\nfrom tf_utils.learners import FGM, AWP\n\nimport os\nimport time\nimport pickle\nimport math\nimport random\nimport sys\n# import cv2\nimport gc\nimport re\nimport glob\nimport datetime\nimport librosa\nprint(f'Tensorflow Version: {tf.__version__}')\nprint(f'Python Version: {sys.version}')\ntqdm.pandas()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-31T16:19:41.195289Z","iopub.execute_input":"2024-05-31T16:19:41.195587Z","iopub.status.idle":"2024-05-31T16:19:55.957498Z","shell.execute_reply.started":"2024-05-31T16:19:41.195561Z","shell.execute_reply":"2024-05-31T16:19:55.956512Z"},"trusted":true},"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    try:\n        tpu = 'local' if device=='TPU-VM' else None\n        print(\"connecting to TPU...\")\n        tpu = tf.distribute.cluster_resolver.TPUClusterResolver()\n        tf.config.experimental_connect_to_cluster(tpu)\n        tf.tpu.experimental.initialize_tpu_system(tpu)\n        strategy = tf.distribute.TPUStrategy(tpu)\n        IS_TPU = True\n    except:\n        device = 'GPU'\n        IS_TPU = None\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\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()","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:19:55.958875Z","iopub.execute_input":"2024-05-31T16:19:55.959478Z","iopub.status.idle":"2024-05-31T16:19:56.198141Z","shell.execute_reply.started":"2024-05-31T16:19:55.959449Z","shell.execute_reply":"2024-05-31T16:19:56.197130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MAX_DURATION = 5\nVAL_MAX_DURATION = 5\n\nSAMPLING_RATE = 32_000\nMAX_LEN = None\nN_MEL = 128\nPAD = -100.\nN_SPLITS = 4\nSEED = 42\nMFCC_FEAT = False\nSPEC_SHAPE = [((N_MEL*2) if MFCC_FEAT else N_MEL),MAX_LEN]\n\nMAX_SEQ_LENGTH = MAX_DURATION * SAMPLING_RATE\nVAL_MAX_SEQ_LENGTH = VAL_MAX_DURATION * SAMPLING_RATE\nTF_REC = True\ns_df  = pd.read_csv(\"/kaggle/input/birdclef-2024/sample_submission.csv\")\ncls2id = {k:idx for idx,k in enumerate(s_df.columns[1:])}\nid2cls = {idx:k for k,idx in cls2id.items()}\nNUM_CLASS = len(cls2id)\nNUM_CLASS","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:19:56.201396Z","iopub.execute_input":"2024-05-31T16:19:56.201759Z","iopub.status.idle":"2024-05-31T16:19:56.265267Z","shell.execute_reply.started":"2024-05-31T16:19:56.201730Z","shell.execute_reply":"2024-05-31T16:19:56.264140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dupes = {\n    (\"asbfly/XC347563.ogg\", \"asbfly/XC341611.ogg\"),\n    (\"grewar3/XC507658.ogg\", \"grewar3/XC184040.ogg\"),\n    ('asbfly/XC724266.ogg', 'asbfly/XC724148.ogg'),\n    ('barswa/XC575749.ogg', 'barswa/XC575747.ogg'),\n    ('bcnher/XC669544.ogg', 'bcnher/XC669542.ogg'),\n    ('bkskit1/XC350251.ogg', 'bkskit1/XC350249.ogg'),\n    ('blhori1/XC417215.ogg', 'blhori1/XC417133.ogg'),\n    ('blhori1/XC743616.ogg', 'blhori1/XC537503.ogg'),\n    ('blrwar1/XC662286.ogg', 'blrwar1/XC662285.ogg'),\n    ('brakit1/XC743675.ogg', 'brakit1/XC537471.ogg'),\n    ('brcful1/XC197746.ogg', 'brcful1/XC157971.ogg'),\n    ('brnshr/XC510751.ogg', 'brnshr/XC510750.ogg'),\n    ('btbeat1/XC665307.ogg', 'btbeat1/XC513403.ogg'),\n    ('btbeat1/XC743618.ogg', 'btbeat1/XC683300.ogg'),\n    ('btbeat1/XC743619.ogg', 'btbeat1/XC683300.ogg'),\n    ('btbeat1/XC743619.ogg', 'btbeat1/XC743618.ogg'),\n    ('categr/XC787914.ogg', 'categr/XC438523.ogg'),\n    ('cohcuc1/XC253418.ogg', 'cohcuc1/XC241127.ogg'),\n    ('cohcuc1/XC423422.ogg', 'cohcuc1/XC423419.ogg'),\n    ('comgre/XC202776.ogg', 'comgre/XC192404.ogg'),\n    ('comgre/XC602468.ogg', 'comgre/XC175341.ogg'),\n    ('comgre/XC64628.ogg', 'comgre/XC58586.ogg'),\n    ('comior1/XC305930.ogg', 'comior1/XC303819.ogg'),\n    ('comkin1/XC207123.ogg', 'comior1/XC207062.ogg'),\n    ('comkin1/XC691421.ogg', 'comkin1/XC690633.ogg'),\n    ('commyn/XC577887.ogg', 'commyn/XC577886.ogg'),\n    ('commyn/XC652903.ogg', 'commyn/XC652901.ogg'),\n    ('compea/XC665320.ogg', 'compea/XC644022.ogg'),\n    ('comsan/XC385909.ogg', 'comsan/XC385908.ogg'),\n    ('comsan/XC643721.ogg', 'comsan/XC642698.ogg'),\n    ('comsan/XC667807.ogg', 'comsan/XC667806.ogg'),\n    ('comtai1/XC126749.ogg', 'comtai1/XC122978.ogg'),\n    ('comtai1/XC305210.ogg', 'comtai1/XC304811.ogg'),\n    ('comtai1/XC542375.ogg', 'comtai1/XC540351.ogg'),\n    ('comtai1/XC542379.ogg', 'comtai1/XC540352.ogg'),\n    ('crfbar1/XC615780.ogg', 'crfbar1/XC615778.ogg'),\n    ('dafbab1/XC188307.ogg', 'dafbab1/XC187059.ogg'),\n    ('dafbab1/XC188308.ogg', 'dafbab1/XC187068.ogg'),\n    ('dafbab1/XC188309.ogg', 'dafbab1/XC187069.ogg'),\n    ('dafbab1/XC197745.ogg', 'dafbab1/XC157972.ogg'),\n    ('eaywag1/XC527600.ogg', 'eaywag1/XC527598.ogg'),\n    ('eucdov/XC355153.ogg', 'eucdov/XC355152.ogg'),\n    ('eucdov/XC360303.ogg', 'eucdov/XC347428.ogg'),\n    ('eucdov/XC365606.ogg', 'eucdov/XC124694.ogg'),\n    ('eucdov/XC371039.ogg', 'eucdov/XC368596.ogg'),\n    ('eucdov/XC747422.ogg', 'eucdov/XC747408.ogg'),\n    ('eucdov/XC789608.ogg', 'eucdov/XC788267.ogg'),\n    ('goflea1/XC163901.ogg', 'bladro1/XC163901.ogg'),\n    ('goflea1/XC208794.ogg', 'bladro1/XC208794.ogg'),\n    ('goflea1/XC208795.ogg', 'bladro1/XC208795.ogg'),\n    ('goflea1/XC209203.ogg', 'bladro1/XC209203.ogg'),\n    ('goflea1/XC209549.ogg', 'bladro1/XC209549.ogg'),\n    ('goflea1/XC209564.ogg', 'bladro1/XC209564.ogg'),\n    ('graher1/XC357552.ogg', 'graher1/XC357551.ogg'),\n    ('graher1/XC590235.ogg', 'graher1/XC590144.ogg'),\n    ('grbeat1/XC304004.ogg', 'grbeat1/XC303999.ogg'),\n    ('grecou1/XC365426.ogg', 'grecou1/XC365425.ogg'),\n    ('greegr/XC247286.ogg', 'categr/XC197438.ogg'),\n    ('grewar3/XC743681.ogg', 'grewar3/XC537475.ogg'),\n    ('grnwar1/XC197744.ogg', 'grnwar1/XC157973.ogg'),\n    ('grtdro1/XC651708.ogg', 'grtdro1/XC613192.ogg'),\n    ('grywag/XC459760.ogg', 'grywag/XC457124.ogg'),\n    ('grywag/XC575903.ogg', 'grywag/XC575901.ogg'),\n    ('grywag/XC650696.ogg', 'grywag/XC592019.ogg'),\n    ('grywag/XC690448.ogg', 'grywag/XC655063.ogg'),\n    ('grywag/XC745653.ogg', 'grywag/XC745650.ogg'),\n    ('grywag/XC812496.ogg', 'grywag/XC812495.ogg'),\n    ('heswoo1/XC357155.ogg', 'heswoo1/XC357149.ogg'),\n    ('heswoo1/XC744698.ogg', 'heswoo1/XC665715.ogg'),\n    ('hoopoe/XC631301.ogg', 'hoopoe/XC365530.ogg'),\n    ('hoopoe/XC631304.ogg', 'hoopoe/XC252584.ogg'),\n    ('houcro1/XC744704.ogg', 'houcro1/XC683047.ogg'),\n    ('houspa/XC326675.ogg', 'houspa/XC326674.ogg'),\n    ('inbrob1/XC744708.ogg', 'inbrob1/XC744706.ogg'),\n    ('insowl1/XC305214.ogg', 'insowl1/XC301142.ogg'),\n    ('junbab2/XC282587.ogg', 'junbab2/XC282586.ogg'),\n    ('labcro1/XC267645.ogg', 'labcro1/XC265731.ogg'),\n    ('labcro1/XC345836.ogg', 'labcro1/XC312582.ogg'),\n    ('labcro1/XC37773.ogg', 'labcro1/XC19736.ogg'),\n    ('labcro1/XC447036.ogg', 'houcro1/XC447036.ogg'),\n    ('labcro1/XC823514.ogg', 'gybpri1/XC823527.ogg'),\n    ('laudov1/XC185511.ogg', 'grewar3/XC185505.ogg'),\n    ('laudov1/XC405375.ogg', 'laudov1/XC405374.ogg'),\n    ('laudov1/XC514027.ogg', 'eucdov/XC514027.ogg'),\n    ('lblwar1/XC197743.ogg', 'lblwar1/XC157974.ogg'),\n    ('lewduc1/XC261506.ogg', 'lewduc1/XC254813.ogg'),\n    ('litegr/XC403621.ogg', 'bcnher/XC403621.ogg'),\n    ('litegr/XC535540.ogg', 'litegr/XC448898.ogg'),\n    ('litegr/XC535552.ogg', 'litegr/XC447850.ogg'),\n    ('litgre1/XC630775.ogg', 'litgre1/XC630560.ogg'),\n    ('litgre1/XC776082.ogg', 'litgre1/XC663244.ogg'),\n    ('litspi1/XC674522.ogg', 'comtai1/XC674522.ogg'),\n    ('litspi1/XC722435.ogg', 'litspi1/XC721636.ogg'),\n    ('litspi1/XC722436.ogg', 'litspi1/XC721637.ogg'),\n    ('litswi1/XC443070.ogg', 'litswi1/XC440301.ogg'),\n    ('lobsun2/XC197742.ogg', 'lobsun2/XC157975.ogg'),\n    ('maghor2/XC197740.ogg', 'maghor2/XC157978.ogg'),\n    ('maghor2/XC786588.ogg', 'maghor2/XC786587.ogg'),\n    ('malpar1/XC197770.ogg', 'malpar1/XC157976.ogg'),\n    ('marsan/XC383290.ogg', 'marsan/XC383288.ogg'),\n    ('marsan/XC733175.ogg', 'marsan/XC716673.ogg'),\n    ('mawthr1/XC455222.ogg', 'mawthr1/XC455211.ogg'),\n    ('orihob2/XC557991.ogg', 'orihob2/XC557293.ogg'),\n    ('piebus1/XC165050.ogg', 'piebus1/XC122395.ogg'),\n    ('piebus1/XC814459.ogg', 'piebus1/XC792272.ogg'),\n    ('placuc3/XC490344.ogg', 'placuc3/XC486683.ogg'),\n    ('placuc3/XC572952.ogg', 'placuc3/XC572950.ogg'),\n    ('plaflo1/XC615781.ogg', 'plaflo1/XC614946.ogg'),\n    ('purher1/XC467373.ogg', 'graher1/XC467373.ogg'),\n    ('purher1/XC827209.ogg', 'purher1/XC827207.ogg'),\n    ('pursun3/XC268375.ogg', 'comtai1/XC241382.ogg'),\n    ('pursun4/XC514853.ogg', 'pursun4/XC514852.ogg'),\n    ('putbab1/XC574864.ogg', 'brcful1/XC574864.ogg'),\n    ('rewbul/XC306398.ogg', 'bkcbul1/XC306398.ogg'),\n    ('rewbul/XC713308.ogg', 'asbfly/XC713467.ogg'),\n    ('rewlap1/XC733007.ogg', 'rewlap1/XC732874.ogg'),\n    ('rorpar/XC199488.ogg', 'rorpar/XC199339.ogg'),\n    ('rorpar/XC402325.ogg', 'comior1/XC402326.ogg'),\n    ('rorpar/XC516404.ogg', 'rorpar/XC516402.ogg'),\n    ('sbeowl1/XC522123.ogg', 'brfowl1/XC522123.ogg'),\n    ('sohmyn1/XC744700.ogg', 'sohmyn1/XC743682.ogg'),\n    ('spepic1/XC804432.ogg', 'spepic1/XC804431.ogg'),\n    ('spodov/XC163930.ogg', 'bladro1/XC163901.ogg'),\n    ('spodov/XC163930.ogg', 'goflea1/XC163901.ogg'),\n    ('spoowl1/XC591485.ogg', 'spoowl1/XC591177.ogg'),\n    ('stbkin1/XC266782.ogg', 'stbkin1/XC266682.ogg'),\n    ('stbkin1/XC360661.ogg', 'stbkin1/XC199815.ogg'),\n    ('stbkin1/XC406140.ogg', 'stbkin1/XC406138.ogg'),\n    ('vefnut1/XC197738.ogg', 'vefnut1/XC157979.ogg'),\n    ('vefnut1/XC293526.ogg', 'vefnut1/XC289785.ogg'),\n    ('wemhar1/XC581045.ogg', 'comsan/XC581045.ogg'),\n    ('wemhar1/XC590355.ogg', 'wemhar1/XC590354.ogg'),\n    ('whbbul2/XC335671.ogg', 'whbbul2/XC335670.ogg'),\n    ('whbsho3/XC856465.ogg', 'whbsho3/XC856463.ogg'),\n    ('whbsho3/XC856468.ogg', 'whbsho3/XC856463.ogg'),\n    ('whbsho3/XC856468.ogg', 'whbsho3/XC856465.ogg'),\n    ('whbwat1/XC840073.ogg', 'whbwat1/XC840071.ogg'),\n    ('whbwoo2/XC239509.ogg', 'rufwoo2/XC239509.ogg'),\n    ('whcbar1/XC659329.ogg', 'insowl1/XC659329.ogg'),\n    ('whiter2/XC265271.ogg', 'whiter2/XC265267.ogg'),\n    ('whtkin2/XC197737.ogg', 'whtkin2/XC157981.ogg'),\n    ('whtkin2/XC430267.ogg', 'whtkin2/XC430256.ogg'),\n    ('whtkin2/XC503389.ogg', 'comior1/XC503389.ogg'),\n    ('whtkin2/XC540094.ogg', 'whtkin2/XC540087.ogg'),\n    ('woosan/XC184466.ogg', 'marsan/XC184466.ogg'),\n    ('woosan/XC545316.ogg', 'woosan/XC476064.ogg'),\n    ('woosan/XC587076.ogg', 'woosan/XC578599.ogg'),\n    ('woosan/XC742927.ogg', 'woosan/XC740798.ogg'),\n    ('woosan/XC825766.ogg', 'grnsan/XC825765.ogg'),\n    ('zitcis1/XC303866.ogg', 'zitcis1/XC302781.ogg'),\n}\nduplicates = [i[1] for i in dupes]","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:19:56.266986Z","iopub.execute_input":"2024-05-31T16:19:56.267406Z","iopub.status.idle":"2024-05-31T16:19:56.294849Z","shell.execute_reply.started":"2024-05-31T16:19:56.267371Z","shell.execute_reply":"2024-05-31T16:19:56.293834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(\"/kaggle/input/birdclef-2024/train_metadata.csv\")\ndf = df[~df.filename.isin(duplicates)].reset_index(drop = True)\ndf['label'] = df.primary_label.apply(lambda x:cls2id[x])\ndf['path'] = '/kaggle/input/birdclef-2024/train_audio/'+df.filename\neb_df = pd.read_csv(\"/kaggle/input/birdclef-2024/eBird_Taxonomy_v2021.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:19:56.296244Z","iopub.execute_input":"2024-05-31T16:19:56.296626Z","iopub.status.idle":"2024-05-31T16:19:56.596722Z","shell.execute_reply.started":"2024-05-31T16:19:56.296593Z","shell.execute_reply":"2024-05-31T16:19:56.595786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_weights =dict( (\n    df.primary_label.replace(cls2id).value_counts() / \n    df.primary_label.replace(cls2id).value_counts().sum()\n)  ** (-0.5))\nsample_weights","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2024-05-31T16:19:56.598171Z","iopub.execute_input":"2024-05-31T16:19:56.598570Z","iopub.status.idle":"2024-05-31T16:19:59.268168Z","shell.execute_reply.started":"2024-05-31T16:19:56.598533Z","shell.execute_reply":"2024-05-31T16:19:59.267117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TF_REC:\n    TRAIN_FILENAMES = glob.glob('/kaggle/input/bclef-*/*.tfrecords')\n    print(len(TRAIN_FILENAMES))\n    def 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)\n    print(count_data_items(TRAIN_FILENAMES), len(df))\n    assert count_data_items(TRAIN_FILENAMES) == len(df)","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:19:59.269619Z","iopub.execute_input":"2024-05-31T16:19:59.270005Z","iopub.status.idle":"2024-05-31T16:19:59.332604Z","shell.execute_reply.started":"2024-05-31T16:19:59.269971Z","shell.execute_reply":"2024-05-31T16:19:59.331590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"audio = tfio.audio.AudioIOTensor('/kaggle/input/birdclef-2024/train_audio/asbfly/XC134896.ogg')\nprint(audio)\naudio_tensor = tf.squeeze(audio.to_tensor(), axis=[-1])\nAudio(audio_tensor.numpy()[:5*32000], rate=audio.rate.numpy())","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:19:59.333660Z","iopub.execute_input":"2024-05-31T16:19:59.333939Z","iopub.status.idle":"2024-05-31T16:20:00.642544Z","shell.execute_reply.started":"2024-05-31T16:19:59.333915Z","shell.execute_reply":"2024-05-31T16:20:00.641608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = '/kaggle/input/birdclef-2024/train_audio/asbfly/XC134896.ogg'\naudio = tf.io.read_file(path)\naudio = tfio.audio.decode_vorbis(audio)","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:20:00.647241Z","iopub.execute_input":"2024-05-31T16:20:00.647900Z","iopub.status.idle":"2024-05-31T16:20:00.851408Z","shell.execute_reply.started":"2024-05-31T16:20:00.647868Z","shell.execute_reply":"2024-05-31T16:20:00.850352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def decode_tfrec(record_bytes):\n    features = tf.io.parse_single_example(record_bytes, {\n        'data': tf.io.FixedLenFeature([], tf.string),\n        'label': tf.io.FixedLenFeature([], tf.int64),\n    })\n    out = {}\n    out['data']  = tf.io.decode_raw(features['data'], tf.float32)\n    out['label'] = features['label']\n    return out\n\ndef decode(path,label):\n    audio = tf.io.read_file(path)\n    audio = tf.squeeze(tfio.audio.decode_vorbis(audio),axis = -1)\n    return dict(data = audio,label = label)","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:20:00.852817Z","iopub.execute_input":"2024-05-31T16:20:00.853184Z","iopub.status.idle":"2024-05-31T16:20:00.860207Z","shell.execute_reply.started":"2024-05-31T16:20:00.853150Z","shell.execute_reply":"2024-05-31T16:20:00.859129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess(out):\n    audio = tf.expand_dims(tf.cast(out['data'],tf.float32),-1)\n    y = tf.one_hot(out['label'],NUM_CLASS)\n    return audio,y\n\ndef limit_seconds_audio(x, y,training = False):\n    audio_len = tf.shape(x)[0]\n    diff_len = abs(audio_len-MAX_SEQ_LENGTH)\n    if training:\n        if audio_len > MAX_SEQ_LENGTH:  # do cropping if audio length is larger\n            idx = tf.random.uniform([], maxval=diff_len, dtype=tf.int32)\n            x = x[idx : (idx + MAX_SEQ_LENGTH)]\n    else:\n        x = x[:VAL_MAX_SEQ_LENGTH]\n    x = tf.where(tf.math.is_nan(x),tf.constant(0.,x.dtype),x)\n    return x, y\n\ndef augment_freq_time_mask(spectrogram,\n                           frequency_masking_para=15,\n                           time_masking_para=15,\n                           frequency_mask_num=1,\n                           time_mask_num=1):\n    time_max = tf.shape(spectrogram)[1]\n    freq_max = tf.shape(spectrogram)[2]\n    # Frequency masking\n    for _ in range(frequency_mask_num):\n        f = tf.random.uniform(shape=(), minval=0, maxval=frequency_masking_para, dtype=tf.dtypes.int32)\n        f0 = tf.random.uniform(shape=(), minval=0, maxval=freq_max - f, dtype=tf.dtypes.int32)\n        value_ones_freq_prev = tf.ones(shape=[1, time_max, f0])\n        value_zeros_freq = tf.zeros(shape=[1, time_max, f])\n        value_ones_freq_next = tf.ones(shape=[1, time_max, freq_max-(f0+f)])\n        freq_mask = tf.concat([value_ones_freq_prev, value_zeros_freq, value_ones_freq_next], axis=2)\n        # mel_spectrogram[:, f0:f0 + f, :] = 0 #can't assign to tensor\n        # mel_spectrogram[:, f0:f0 + f, :] = value_zeros_freq #can't assign to tensor\n        spectrogram = spectrogram*freq_mask\n\n    # Time masking\n    for _ in range(time_mask_num):\n        t = tf.random.uniform(shape=(), minval=0, maxval=time_masking_para, dtype=tf.dtypes.int32)\n        t0 = tf.random.uniform(shape=(), minval=0, maxval=time_max - t, dtype=tf.dtypes.int32)\n        value_zeros_time_prev = tf.ones(shape=[1, t0, freq_max])\n        value_zeros_time = tf.zeros(shape=[1, t, freq_max])\n        value_zeros_time_next = tf.ones(shape=[1, time_max-(t0+t), freq_max])\n        time_mask = tf.concat([value_zeros_time_prev, value_zeros_time, value_zeros_time_next], axis=1)\n        # mel_spectrogram[:, :, t0:t0 + t] = 0 #can't assign to tensor\n        # mel_spectrogram[:, :, t0:t0 + t] = value_zeros_time #can't assign to tensor\n        spectrogram = spectrogram*time_mask\n\n    return spectrogram\n\ndef augment_pitch_and_tempo(spectrogram,\n                            max_tempo=1.2,\n                            max_pitch=1.1,\n                            min_pitch=0.95):\n    original_shape = tf.shape(spectrogram)\n    choosen_pitch = tf.random.uniform(shape=(), minval=min_pitch, maxval=max_pitch)\n    choosen_tempo = tf.random.uniform(shape=(), minval=1, maxval=max_tempo)\n    new_freq_size = tf.cast(tf.cast(original_shape[2], tf.float32)*choosen_pitch, tf.int32)\n    new_time_size = tf.cast(tf.cast(original_shape[1], tf.float32)/(choosen_tempo), tf.int32)\n    spectrogram_aug = tf.image.resize(tf.expand_dims(spectrogram, -1), [new_time_size, new_freq_size])\n    spectrogram_aug = tf.image.crop_to_bounding_box(spectrogram_aug, offset_height=0, offset_width=0, target_height=tf.shape(spectrogram_aug)[1], target_width=tf.minimum(original_shape[2], new_freq_size))\n    spectrogram_aug = tf.cond(choosen_pitch < 1,\n                              lambda: tf.image.pad_to_bounding_box(spectrogram_aug, offset_height=0, offset_width=0,\n                                                                   target_height=tf.shape(spectrogram_aug)[1], target_width=original_shape[2]),\n                              lambda: spectrogram_aug)\n    return spectrogram_aug[:, :, :, 0]\n\ndef augment_speed_up(spectrogram,\n                     speed_std=0.1):\n    original_shape = tf.shape(spectrogram)\n    choosen_speed = tf.math.abs(tf.random.normal(shape=(), stddev=speed_std)) # abs makes sure the augmention will only speed up\n    choosen_speed = 1 + choosen_speed\n    new_freq_size = tf.cast(tf.cast(original_shape[2], tf.float32), tf.int32)\n    new_time_size = tf.cast(tf.cast(original_shape[1], tf.float32)/(choosen_speed), tf.int32)\n    spectrogram_aug = tf.image.resize(tf.expand_dims(spectrogram, -1), [new_time_size, new_freq_size])\n    return spectrogram_aug[:, :, :, 0]\n\ndef augment_dropout(spectrogram,\n                    keep_prob=0.9):\n    return tf.nn.dropout(spectrogram, rate=1-keep_prob)\nclass MelSpectrogram(tf.keras.layers.Layer):\n    \"\"\"\n    Mel Spectrogram Layer to convert audio to mel spectrogram which works with single or batched inputs.\n\n    Args:\n        n_fft (int): Size of the FFT window.\n        hop_length (int): Number of samples between successive STFT columns.\n        win_length (int): Size of the STFT window. If None, defaults to n_fft.\n        window_fn (str): Name of the window function to use.\n        sr (int): Sample rate of the input signal.\n        n_mels (int): Number of mel bins to generate.\n        fmin (float): Minimum frequency of the mel bins.\n        fmax (float): Maximum frequency of the mel bins. If None, defaults to sr / 2.\n        power (float): Exponent for the magnitude spectrogram.\n        power_to_db (bool): Whether to convert the power spectrogram to decibels.\n        top_db (float): Maximum decibel value for the output spectrogram.\n        power_to_db (bool): Whether to convert spectrogram from energy to power.\n        out_channels (int): Number of output channels. If None, no channel is created.\n\n    Call Args:\n        input (tf.Tensor): Audio signal of shape (audio_len,) or (None, audio_len)\n\n    Returns:\n        tf.Tensor: Mel spectrogram of shape (..., n_mels, time, out_channels)\n        or (..., n_mels, time) if out_channels is None.\n\n    \"\"\"\n\n    def __init__(\n        self,\n        n_fft=2048,\n        hop_length=512,\n        win_length=None,\n        window=\"hann_window\",\n        sr=SAMPLING_RATE,\n        n_mels=128,\n        fmin=20.0,\n        fmax=None,\n        power_to_db=True,\n        top_db=80.0,\n        power=2.0,\n        amin=1e-10,\n        ref=1.0,\n        out_channels=None,\n        name=\"mel_spectrogram\",\n        **kwargs,\n    ):\n        super(MelSpectrogram, self).__init__(name=name, **kwargs)\n        self.n_fft = n_fft\n        self.hop_length = hop_length\n        self.win_length = win_length or n_fft\n        self.window = window\n        self.sr = sr\n        self.n_mels = n_mels\n        self.fmin = fmin\n        self.fmax = fmax or int(sr / 2)\n        self.power_to_db = power_to_db\n        self.top_db = top_db\n        self.power = power\n        self.amin = amin\n        self.ref = ref\n        self.out_channels = out_channels\n\n    @tf.function\n    def call(self, input):\n        spec = self.spectrogram(input)  # audio to spectrogram with shape\n        spec = self.melscale(spec)  # spectrogram to mel spectrogram\n        if self.power_to_db:\n            spec = self.dbscale(spec)  # mel spectrogram to decibel mel spectrogram\n        spec = tf.linalg.matrix_transpose(\n            spec\n        )  # (..., time, n_mels) to (..., n_mels, time)\n        if self.out_channels is not None:\n            spec = self.update_channels(spec)\n        return spec\n\n    def spectrogram(self, input):\n        spec = tf.signal.stft(\n            input,\n            frame_length=self.win_length,\n            frame_step=self.hop_length,\n            fft_length=self.n_fft,\n            window_fn=getattr(tf.signal, self.window),\n            pad_end=True,\n        )\n        spec = tf.math.pow(tf.math.abs(spec), self.power)\n        return spec\n\n    def melscale(self, input):\n        nbin = tf.shape(input)[-1]\n        matrix = tf.signal.linear_to_mel_weight_matrix(\n            num_mel_bins=self.n_mels,\n            num_spectrogram_bins=nbin,\n            sample_rate=self.sr,\n            lower_edge_hertz=self.fmin,\n            upper_edge_hertz=self.fmax,\n        )\n        return tf.tensordot(input, matrix, axes=1)\n\n    def dbscale(self, input):\n        log_spec = 10.0 * (\n            tf.math.log(tf.math.maximum(input, self.amin)) / tf.math.log(10.0)\n        )\n        if callable(self.ref):\n            ref_value = self.ref(log_spec)\n        else:\n            ref_value = tf.math.abs(self.ref)\n        log_spec -= (\n            10.0\n            * tf.math.log(tf.math.maximum(ref_value, self.amin))\n            / tf.math.log(10.0)\n        )\n        log_spec = tf.math.maximum(log_spec, tf.math.reduce_max(log_spec) - self.top_db)\n        return log_spec\n\n    def update_channels(self, input):\n        spec = input[..., tf.newaxis]\n        if self.out_channels > 1:\n            multiples = tf.concat(\n                [\n                    tf.ones(tf.rank(spec) - 1, dtype=tf.int32),\n                    tf.constant([self.out_channels], dtype=tf.int32),\n                ],\n                axis=0,\n            )\n            spec = tf.tile(spec, multiples)\n        return spec\n\n    def get_config(self):\n        config = super(MelSpectrogram, self).get_config()\n        config.update(\n            {\n                \"n_fft\": self.n_fft,\n                \"hop_length\": self.hop_length,\n                \"win_length\": self.win_length,\n                \"window\": self.window,\n                \"sr\": self.sr,\n                \"n_mels\": self.n_mels,\n                \"fmin\": self.fmin,\n                \"fmax\": self.fmax,\n                \"power_to_db\": self.power_to_db,\n                \"top_db\": self.top_db,\n                \"power\": self.power,\n                \"amin\": self.amin,\n                \"ref\": self.ref,\n                \"out_channels\": self.out_channels,\n            }\n        )\n        return config\n\n    \ndef audio_to_mel_spectrogram(x_audio, y):\n    x_audio = tf.squeeze(x_audio, axis=-1)\n    x_audio = MelSpectrogram(dtype = tf.float32,out_channels=None)(x_audio)\n\n    return x_audio, y\n\ndef audio_to_mfcc_mel_spectogram(x_audio,y):\n    x_audio = tf.squeeze(x_audio, axis=-1)\n    x_audio = MelSpectrogram(dtype = tf.float32,out_channels=None)(x_audio)\n    # Compute MFCCs from log_mel_spectrograms and take the first 13.\n    mfccs = tf.signal.mfccs_from_log_mel_spectrograms(\n      x_audio)#[..., :13]\n    return mfccs,y\n\ndef spectrogram_data_augmentation(spectrogram, y):\n    spectrogram = tf.expand_dims(spectrogram, axis=0)\n    if  tf.random.uniform([]) > 0.5:\n        spectrogram = augment_freq_time_mask(spectrogram)\n    rand = tf.random.uniform([])\n    if  rand > 0.5:\n        spectrogram = augment_dropout(spectrogram, keep_prob=1-rand)\n    if  tf.random.uniform([]) > 0.5:\n        spectrogram = augment_pitch_and_tempo(spectrogram)\n    spectrogram = tf.squeeze(spectrogram, axis=0)\n    return spectrogram,y\n\n\ndef freq_mask(x_spectrogram, y):\n    x_spectrogram = tfio.audio.freq_mask(x_spectrogram, param=10)\n    return x_spectrogram,y\n\ndef standarize(x,y):\n    x = x - tf.math.reduce_mean(x)\n    x = x / tf.math.reduce_std(x)\n    x = tf.where(tf.math.is_nan(x),tf.constant(0.,x.dtype),x)\n    return x,y\n","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:20:00.861541Z","iopub.execute_input":"2024-05-31T16:20:00.861874Z","iopub.status.idle":"2024-05-31T16:20:00.910421Z","shell.execute_reply.started":"2024-05-31T16:20:00.861849Z","shell.execute_reply":"2024-05-31T16:20:00.909468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def update_channels(input,y):\n    spec = input[..., tf.newaxis]\n    multiples = tf.concat(\n        [\n            tf.ones(tf.rank(spec) - 1, dtype=tf.int32),\n            tf.constant([3], dtype=tf.int32),\n        ],\n        axis=0,\n    )\n    spec = tf.tile(spec, multiples)\n    return spec,y","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:20:00.911684Z","iopub.execute_input":"2024-05-31T16:20:00.912053Z","iopub.status.idle":"2024-05-31T16:20:00.925382Z","shell.execute_reply.started":"2024-05-31T16:20:00.912015Z","shell.execute_reply":"2024-05-31T16:20:00.924392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def concat_feat(audio,y):\n    mfcc = audio_to_mfcc_mel_spectogram(audio,y)[0]\n    spec = audio_to_mel_spectrogram(audio,y)[0]\n    return tf.concat([spec,mfcc],axis = -1),y","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:20:00.926515Z","iopub.execute_input":"2024-05-31T16:20:00.926782Z","iopub.status.idle":"2024-05-31T16:20:00.936129Z","shell.execute_reply.started":"2024-05-31T16:20:00.926752Z","shell.execute_reply":"2024-05-31T16:20:00.935115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_data(tfrecords, batch_size=64,drop_remainder=False,augment = False,shuffle=False, repeat=False):\n    if TF_REC:\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    else:\n        ds = tf.data.Dataset.from_tensor_slices((tfrecords.path,tfrecords.label))\n        ds = ds.map(decode,tf.data.AUTOTUNE)\n    if IS_TPU:\n        ds = ds.cache()  # cache data for speedup\n    ds = ds.map(preprocess, tf.data.AUTOTUNE)\n    ds = ds.map(lambda x,y:(limit_seconds_audio(x,y,training = augment)), num_parallel_calls=tf.data.AUTOTUNE)\n    if not IS_TPU:\n        ds = ds.cache()  # cache data for speedup\n\n    if MFCC_FEAT:\n        ds = ds.map(concat_feat, num_parallel_calls=tf.data.AUTOTUNE)\n        \n    else:\n        ds = ds.map(audio_to_mel_spectrogram, num_parallel_calls=tf.data.AUTOTUNE)\n    \n    if augment:\n        ds = ds.map(spectrogram_data_augmentation, num_parallel_calls=tf.data.AUTOTUNE)\n        \n    ds = ds.map(standarize, num_parallel_calls=tf.data.AUTOTUNE)\n    ds = ds.map(update_channels, num_parallel_calls=tf.data.AUTOTUNE)\n    if repeat:\n        ds = ds.repeat()\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    if batch_size:\n        ds = ds.padded_batch(batch_size, padding_values=PAD, padded_shapes=(SPEC_SHAPE+[3],[NUM_CLASS]), drop_remainder=drop_remainder)\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n    return ds","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:20:00.937362Z","iopub.execute_input":"2024-05-31T16:20:00.937975Z","iopub.status.idle":"2024-05-31T16:20:00.950204Z","shell.execute_reply.started":"2024-05-31T16:20:00.937949Z","shell.execute_reply":"2024-05-31T16:20:00.949315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TF_REC:\n    train_ds = get_data(TRAIN_FILENAMES,drop_remainder=True, shuffle=True,augment = True, repeat=True)\nelse:\n    train_ds = get_data(df,drop_remainder=True, shuffle=True,augment = True, repeat=True)\nprint(train_ds)\nfor idx, x in enumerate(train_ds):\n    tmp_data = x\n    break","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:20:00.951528Z","iopub.execute_input":"2024-05-31T16:20:00.951911Z","iopub.status.idle":"2024-05-31T16:20:06.961003Z","shell.execute_reply.started":"2024-05-31T16:20:00.951849Z","shell.execute_reply":"2024-05-31T16:20:06.960119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _batch_norm(name, params):\n    def _bn_layer(layer_input):\n        return tf.keras.layers.BatchNormalization(\n          name=name,\n          center=params.batchnorm_center,\n          scale=params.batchnorm_scale,\n          epsilon=params.batchnorm_epsilon)(layer_input)\n    return _bn_layer\n\n\ndef _conv(name, kernel, stride, filters, params):\n    def _conv_layer(layer_input):\n        output = tf.keras.layers.Conv2D(name='{}/conv'.format(name),\n                               filters=filters,\n                               kernel_size=kernel,\n                               strides=stride,\n                               padding=params.conv_padding,\n                               use_bias=False,\n                               activation=None)(layer_input)\n        output = _batch_norm('{}/conv/bn'.format(name), params)(output)\n        output = tf.keras.layers.ReLU(name='{}/relu'.format(name))(output)\n        return output\n    return _conv_layer\n\n\ndef _separable_conv(name, kernel, stride, filters, params):\n    def _separable_conv_layer(layer_input):\n        output = tf.keras.layers.DepthwiseConv2D(name='{}/depthwise_conv'.format(name),\n                                        kernel_size=kernel,\n                                        strides=stride,\n                                        depth_multiplier=1,\n                                        padding=params.conv_padding,\n                                        use_bias=False,\n                                        activation=None)(layer_input)\n        output = _batch_norm('{}/depthwise_conv/bn'.format(name), params)(output)\n        output = tf.keras.layers.ReLU(name='{}/depthwise_conv/relu'.format(name))(output)\n        output = tf.keras.layers.Conv2D(name='{}/pointwise_conv'.format(name),\n                               filters=filters,\n                               kernel_size=(1, 1),\n                               strides=1,\n                               padding=params.conv_padding,\n                               use_bias=False,\n                               activation=None)(output)\n        output = _batch_norm('{}/pointwise_conv/bn'.format(name), params)(output)\n        output = tf.keras.layers.ReLU(name='{}/pointwise_conv/relu'.format(name))(output)\n        return output\n    return _separable_conv_layer\n\n\n_YAMNET_LAYER_DEFS = [\n    # (layer_function, kernel, stride, num_filters)\n    (_conv,          [3, 3], 2,   32),\n    (_separable_conv, [3, 3], 1,   64),\n    (_separable_conv, [3, 3], 2,  128),\n    (_separable_conv, [3, 3], 1,  128),\n    (_separable_conv, [3, 3], 2,  256),\n    (_separable_conv, [3, 3], 1,  256),\n    (_separable_conv, [3, 3], 2,  512),\n    (_separable_conv, [3, 3], 1,  512),\n    (_separable_conv, [3, 3], 1,  512),\n    (_separable_conv, [3, 3], 1,  512),\n    (_separable_conv, [3, 3], 1,  512),\n    (_separable_conv, [3, 3], 1,  512),\n    (_separable_conv, [3, 3], 2, 1024),\n    (_separable_conv, [3, 3], 1, 1024)\n]\n\n\nfrom dataclasses import dataclass\n\n# The following hyperparameters (except patch_hop_seconds) were used to train YAMNet,\n# so expect some variability in performance if you change these. The patch hop can\n# be changed arbitrarily: a smaller hop should give you more patches from the same\n# clip and possibly better performance at a larger computational cost.\n@dataclass(frozen=True)  # Instances of this class are immutable.\nclass Params:\n    num_classes: int = NUM_CLASS\n    conv_padding: str = 'same'\n    batchnorm_center: bool = True\n    batchnorm_scale: bool = False\n    batchnorm_epsilon: float = 1e-4\n    classifier_activation: str = 'sigmoid'\n\n    tflite_compatible: bool = False\ndef yamnet(features):\n    \"\"\"Define the core YAMNet mode in Keras.\"\"\"\n#     net = tf.keras.layers.Reshape(\n#       (params.patch_frames, params.patch_bands, 1),\n#       input_shape=(params.patch_frames, params.patch_bands))(features)\n    net = features\n    for (i, (layer_fun, kernel, stride, filters)) in enumerate(_YAMNET_LAYER_DEFS):\n        net = layer_fun('layer{}'.format(i + 1), kernel, stride, filters, params = Params)(net)\n    x = tf.keras.layers.GlobalAveragePooling2D()(net)\n    x = tf.keras.layers.Dropout(0.8)(x)\n    logits = tf.keras.layers.Dense(units=Params.num_classes, use_bias=True)(x)\n    return logits","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:20:06.962549Z","iopub.execute_input":"2024-05-31T16:20:06.962901Z","iopub.status.idle":"2024-05-31T16:20:06.984909Z","shell.execute_reply.started":"2024-05-31T16:20:06.962861Z","shell.execute_reply":"2024-05-31T16:20:06.984007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model():\n    \"\"\"Defines the YAMNet waveform-to-class-scores model.\n\n    Args:\n    params: An instance of Params containing hyperparameters.\n\n    Returns:\n    A model accepting (num_samples,) waveform input and emitting:\n    - predictions: (num_patches, num_classes) matrix of class scores per time frame\n    - embeddings: (num_patches, embedding size) matrix of embeddings per time frame\n    - log_mel_spectrogram: (num_spectrogram_frames, num_mel_bins) spectrogram feature matrix\n    \"\"\"\n    inputs = tf.keras.layers.Input(shape=SPEC_SHAPE+[3], dtype=tf.float32, name='audio_waveform')\n    x = tf.keras.layers.Masking(mask_value=PAD,input_shape=SPEC_SHAPE+[3])(inputs)\n    predictions = yamnet(x)\n    frames_model = tf.keras.Model(\n        name='yamnet_frames', inputs=inputs,\n        outputs=[predictions])\n    return frames_model\n","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:20:06.986233Z","iopub.execute_input":"2024-05-31T16:20:06.986520Z","iopub.status.idle":"2024-05-31T16:20:07.001473Z","shell.execute_reply.started":"2024-05-31T16:20:06.986495Z","shell.execute_reply":"2024-05-31T16:20:07.000605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = get_model()\ny = model(tmp_data[0])\ntf.keras.losses.CategoricalCrossentropy(from_logits=True)(tmp_data[1],y)","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:20:07.002688Z","iopub.execute_input":"2024-05-31T16:20:07.002963Z","iopub.status.idle":"2024-05-31T16:20:09.193868Z","shell.execute_reply.started":"2024-05-31T16:20:07.002942Z","shell.execute_reply":"2024-05-31T16:20:09.192898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AWP(tf.keras.Model):\n    def __init__(self, *args, delta=0.1, eps=1e-4, start_step=0, **kwargs):\n        super().__init__(*args, **kwargs)\n        self.delta = delta\n        self.eps = eps\n        self.start_step = start_step\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        try:\n            x, y, sample_weight = data\n        except:\n            x, y = data\n            sample_weight = None\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,sample_weight=sample_weight)\n        params = self.trainable_variables\n        params_gradients = tape.gradient(loss, self.trainable_variables)\n        for i in range(len(params_gradients)):\n            grad = tf.zeros_like(params[i]) + params_gradients[i]\n            delta = tf.math.divide_no_nan(self.delta * grad , tf.math.sqrt(tf.reduce_sum(grad**2)) + self.eps)\n            self.trainable_variables[i].assign_add(delta)\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,sample_weight=sample_weight)\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        for i in range(len(params_gradients)):\n            grad = tf.zeros_like(params[i]) + params_gradients[i]\n            delta = tf.math.divide_no_nan(self.delta * grad , tf.math.sqrt(tf.reduce_sum(grad**2)) + self.eps)\n            self.trainable_variables[i].assign_sub(delta)\n        self.optimizer.apply_gradients(zip(gradients, self.trainable_variables))\n        # self_loss.update_state(loss)\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":{"iopub.status.busy":"2024-05-31T16:20:09.195218Z","iopub.execute_input":"2024-05-31T16:20:09.195933Z","iopub.status.idle":"2024-05-31T16:20:09.210240Z","shell.execute_reply.started":"2024-05-31T16:20:09.195899Z","shell.execute_reply":"2024-05-31T16:20:09.209431Z"},"trusted":true},"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    if CFG.fp16:\n        try:\n            policy = mixed_precision.Policy('mixed_bfloat16')\n            mixed_precision.set_global_policy(policy)\n        except:\n            policy = mixed_precision.Policy('mixed_float16')\n            mixed_precision.set_global_policy(policy)\n    else:\n        policy = mixed_precision.Policy('float32')\n        mixed_precision.set_global_policy(policy)\n\n    if fold != 'all':\n        train_ds = get_data(train_files, batch_size=CFG.batch_size, drop_remainder=True, augment=True, repeat=True, shuffle=2048)\n        valid_ds = get_data(valid_files, batch_size=CFG.batch_size, drop_remainder=False, augment=False, repeat=False, shuffle=False)\n    else:\n        train_ds = get_data(train_files, batch_size=CFG.batch_size, drop_remainder=False, augment=True, repeat=True, shuffle=2048)\n        valid_ds = None\n        valid_files = []\n    \n    if TF_REC:\n        num_train = count_data_items(train_files)\n        num_valid = count_data_items(valid_files)\n        \n    else:\n        num_train = len(train_files)\n        num_valid = len(valid_files)\n    \n    steps_per_epoch = num_train//CFG.batch_size\n    print('steps_per_epoch',steps_per_epoch)\n    \n    with strategy.scope():\n        dropout_step = CFG.dropout_start_epoch * steps_per_epoch\n        model = get_model()\n\n        schedule = OneCycleLR(CFG.lr, CFG.epoch, warmup_epochs=CFG.epoch*CFG.warmup, steps_per_epoch=steps_per_epoch, resume_epoch=CFG.resume, decay_epochs=CFG.epoch, lr_min=CFG.lr_min, decay_type=CFG.decay_type, warmup_type='linear')\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=CFG.resume, decay_epochs=CFG.epoch, lr_min=CFG.lr_min*CFG.weight_decay, decay_type=CFG.decay_type, warmup_type='linear')\n                \n        awp_step = CFG.awp_start_epoch * steps_per_epoch\n        if CFG.fgm:\n            model = FGM(model.input, model.output, delta=CFG.awp_lambda, eps=0., start_step=awp_step)\n        elif CFG.awp:\n            model = AWP(model.input, model.output, delta=CFG.awp_lambda, eps=0., start_step=awp_step)\n\n        opt = tfa.optimizers.RectifiedAdam(learning_rate=schedule, weight_decay=decay_schedule, sma_threshold=4)#, clipvalue=1)\n        opt = tfa.optimizers.Lookahead(opt,sync_period=5)\n#         opt = tfa.optimizers.AdamW(learning_rate=schedule, weight_decay=decay_schedule,)\n        model.compile(\n            optimizer=opt,\n            loss=[\n                tf.keras.losses.CategoricalCrossentropy(from_logits=True,label_smoothing = 0.25),\n                tf.keras.losses.CategoricalFocalCrossentropy(from_logits=True,label_smoothing = 0.25),\n            ],\n            metrics=[\n                [\n                tf.keras.metrics.AUC(from_logits=True),\n                tf.keras.metrics.CategoricalAccuracy(),\n                ],\n            ],\n            steps_per_execution= None if os.environ['KAGGLE_KERNEL_RUN_TYPE']=='Interactive' else steps_per_epoch,\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 CFG.resume:\n        print(f'resume from epoch{CFG.resume}')\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 = tf.keras.callbacks.CSVLogger(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-logs.csv')\n    sv_loss = tf.keras.callbacks.ModelCheckpoint(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-best.h5', monitor='val_loss', verbose=0, save_best_only=True,\n                save_weights_only=True, mode='min', save_freq='epoch')\n    sv_loss2 = tf.keras.callbacks.ModelCheckpoint(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-best.h5', monitor='loss', verbose=0, save_best_only=True,\n                save_weights_only=True, mode='min', save_freq='epoch')\n    snap = Snapshot(f'{CFG.output_dir}/{CFG.comment}-fold{fold}', CFG.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, valid_steps=-(num_valid//-CFG.batch_size))\n    ton = tf.keras.callbacks.TerminateOnNaN()\n    callbacks = []\n    if CFG.save_output:\n        callbacks.append(logger)\n        callbacks.append(snap)\n        callbacks.append(swa)\n        callbacks.append(ton)\n        if fold != 'all':\n            callbacks.append(sv_loss)\n        else:\n            callbacks.append(sv_loss2)\n\n    history = model.fit(\n        train_ds,\n        epochs=CFG.epoch-CFG.resume,\n        steps_per_epoch=steps_per_epoch,\n        callbacks=callbacks,\n        validation_data=valid_ds,\n        verbose=CFG.verbose,\n        validation_steps=-(num_valid//-CFG.batch_size),\n#         class_weight = sample_weights,\n\n    )\n\n    if CFG.save_output:\n        try:\n            model.load_weights(f'{CFG.output_dir}/{CFG.comment}-fold{fold}-best.h5')\n        except:\n            pass\n    if fold != 'all':\n        print('valid data')\n        cv = model.evaluate(valid_ds,verbose=CFG.verbose,steps=-(num_valid//-CFG.batch_size))\n        print('train data')\n        cv = model.evaluate(train_ds,verbose=CFG.verbose,steps=steps_per_epoch)\n\n    else:\n        cv = model.evaluate(train_ds,verbose=CFG.verbose,steps=steps_per_epoch)\n\n    return model, cv, history\n\ndef main(CFG, folds, strategy=STRATEGY, summary=True):\n    for fold in folds:\n        if fold != 'all':\n            if TF_REC:\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                all_files = train_folds.copy()#TRAIN_FILENAMES\n                train_files = all_files[~all_files.fold.isin([fold])]\n                valid_files =  all_files[all_files.fold.isin([fold])]\n\n        else:\n            train_files = TRAIN_FILENAMES if TF_REC else train_folds.copy()\n            valid_files = None\n\n    return train_fold(CFG, fold, train_files, valid_files, strategy=strategy, summary=summary)","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:20:09.211631Z","iopub.execute_input":"2024-05-31T16:20:09.212174Z","iopub.status.idle":"2024-05-31T16:20:09.243506Z","shell.execute_reply.started":"2024-05-31T16:20:09.212142Z","shell.execute_reply":"2024-05-31T16:20:09.242560Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    n_splits = 4 if TF_REC else 5\n    save_output = True\n    output_dir = '/kaggle/working'\n    \n    seed = 45\n    verbose = 1 #0) silent 1) progress bar 2) one line per epoch\n    \n#     max_len = MAX_LEN\n    replicas = 8\n    lr = 5e-4 * replicas\n    weight_decay = 0.1\n    lr_min = 1e-6\n    epoch = 500 #400\n    warmup = 0\n    batch_size = 32 * replicas\n    snapshot_epochs = []\n    swa_epochs = list(range(epoch//2,epoch+1))\n    \n    fp16 = True\n    fgm = False\n    awp = True\n    awp_lambda = 0.2\n    awp_start_epoch = 0\n    dropout_start_epoch = 0\n    resume = 0\n    decay_type = 'cosine'\n    dim = 192\n    comment = f'BirdCLEF-fp16-YAMNet-{replicas}-seed{seed}'","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:20:09.244639Z","iopub.execute_input":"2024-05-31T16:20:09.244900Z","iopub.status.idle":"2024-05-31T16:20:09.257890Z","shell.execute_reply.started":"2024-05-31T16:20:09.244878Z","shell.execute_reply":"2024-05-31T16:20:09.257105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_folds = df.copy()\ntrain_folds['fold']=-1\nnum_bins = 5\n\nkfold = KFold(n_splits=CFG.n_splits, shuffle=True, random_state=CFG.seed)\nprint(f'{CFG.n_splits}fold training', len(train_folds), 'samples')\nfor fold_idx, (train_idx, valid_idx) in enumerate(kfold.split(train_folds)):\n    train_folds.loc[valid_idx,'fold'] = fold_idx\n    print(f'fold{fold_idx}:', 'train', len(train_idx), 'valid', len(valid_idx))\n\nassert not (train_folds['fold']==-1).sum()\nassert len(np.unique(train_folds['fold']))==CFG.n_splits\ntrain_folds.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-31T16:20:09.259138Z","iopub.execute_input":"2024-05-31T16:20:09.259779Z","iopub.status.idle":"2024-05-31T16:20:09.311824Z","shell.execute_reply.started":"2024-05-31T16:20:09.259748Z","shell.execute_reply":"2024-05-31T16:20:09.310921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model,cv,history = main(CFG, [0])","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2024-05-31T16:20:09.313243Z","iopub.execute_input":"2024-05-31T16:20:09.313906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model,cv,history = main(CFG, [1])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model,cv,history = main(CFG, [2])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model,cv,history = main(CFG, [3])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model,cv,history = main(CFG, ['all'])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# CFG.seed = 42\n# CFG.comment = f'BirdCLEF-fp16-{CFG.dim}-{CFG.replicas}-seed{CFG.seed}'\n# model,cv,history = train_folds(CFG, ['all'], summary=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hist = history.history","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax1 = plt.subplots(figsize = (20,7))\n\nax2 = ax1.twinx()\nline1 = ax1.plot(hist['loss'],'r-o',label = 'loss')\nline2 = ax2.plot(hist['categorical_accuracy'],'-o',label = 'accuracy')\nline3 = ax2.plot(hist['auc'],'-o',label = 'auc')\n\ntry:\n    line4 = ax1.plot(hist['val_loss'], 'g-o',label = 'val_loss')\n    line5 = ax2.plot(hist['val_categorical_accuracy'],'-o',label = 'val_accuracy')\n    line6 = ax2.plot(hist['val_auc'],'-o',label = 'val_auc')\n\nexcept:\n    pass\n\ntry:\n    lns =line1+line4+line2+line5+line3+line6\nexcept:\n    lns =line1+line2+line3\n\nlabs = [l.get_label() for l in lns]\nplt.legend(lns, labs, loc=0)\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}