{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.13"},"papermill":{"default_parameters":{},"duration":24.931552,"end_time":"2026-09-03T21:53:36.080506+00:00","environment_variables":{},"exception":true,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-09-03T21:53:11.148954+00:00","version":"2.7.0"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":7634}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"8350be03","cell_type":"markdown","source":"## Preface\nThis notebooks aims to build a light-weight CNN.\n\nIt uses specgrams of resampled wav files(rate 8000) as inputs.\n\nDue to Kaggle cloud hardware limitations, this script is a 'crippled' version of the original one.\n\nIn order to get LB 0.74, you need to set epoch to 5, set chop_audio(num=1000) and double all Conv layer parameters.\n\nAlthough this script is a slight imrpovement over Alex Ozerin's baseline, I believe by using original wav files(16000 sample rate) one can achieve higher scores.\n\n\n## File Structure\nThis script assumes data are stored in following strcuture:\n\nspeech\n\n├── test            \n\n│   └── audio #test wavfiles\n\n├── train           \n\n│   ├── audio #train wavfiles\n\n└── model #store models\n\n│\n\n└── out #store sub.csv\n\n## Improve This Script\nSince this is only a light-weight CNN, it's performance is limited.\nHere are some ways to improve it's performance.\n1. Use original wav files instead resampled ones.\n2. Create more 'silence' wav files using chop_audio.\n3. Build deeper CNN or use RNN.\n4. Train for longer epochs\n\n## After Words\nIt's still a long way to reach LB 0.88.\n\nIn fact, I doubt CNN would ever reach that high.\n\nFeel free to share your ideas in the comment sections about using CNN to label wav files :)\n\n## Appendix\nThanks __DavidS__ and __Alex Ozerin__ for their great notebooks!","metadata":{"_cell_guid":"d524ad8f-d024-4ba3-8bbc-2435ec3d0dba","_uuid":"7bb30a6f2a18be45a491092a6e2dfcdcce71a2a0","papermill":{"duration":0.003191,"end_time":"2026-09-03T21:53:13.523625+00:00","exception":false,"start_time":"2026-09-03T21:53:13.520434+00:00","status":"completed"},"tags":[]}},{"id":"69d4e45c","cell_type":"code","source":"import os\nimport numpy as np\nfrom scipy.fftpack import fft\nfrom scipy.io import wavfile\nfrom scipy import signal\nfrom glob import glob\nimport re\nimport pandas as pd\nimport gc\nfrom scipy.io import wavfile\n\nfrom keras import optimizers, losses, activations, models\nfrom keras.layers import Convolution2D, Dense, Input, Flatten, Dropout, MaxPooling2D, BatchNormalization\nfrom sklearn.model_selection import train_test_split\nimport keras","metadata":{"_cell_guid":"f7fd8bcb-4451-4d47-bfe8-491c94b3b4eb","_uuid":"712710f20b00f97271136cfeab9937a4c6a2458b","collapsed":true,"execution":{"iopub.execute_input":"2026-09-03T21:53:13.529342Z","iopub.status.busy":"2026-09-03T21:53:13.529053Z","iopub.status.idle":"2026-09-03T21:53:29.894642Z","shell.execute_reply":"2026-09-03T21:53:29.893955Z"},"jupyter":{"outputs_hidden":true},"papermill":{"duration":16.370326,"end_time":"2026-09-03T21:53:29.896359+00:00","exception":false,"start_time":"2026-09-03T21:53:13.526033+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"7bb99f16","cell_type":"markdown","source":"The original sample rate is 16000, and we will resample it to 8000 to reduce data size.","metadata":{"_cell_guid":"fb35a2f1-9301-4693-a9ef-9d180b630f05","_uuid":"4b1ba61998e14e15c822c605dbe5961bfed36014","papermill":{"duration":0.002224,"end_time":"2026-09-03T21:53:29.90113+00:00","exception":false,"start_time":"2026-09-03T21:53:29.898906+00:00","status":"completed"},"tags":[]}},{"id":"85ec7b25","cell_type":"code","source":"import os\nfrom glob import glob\n\nL = 16000\nlegal_labels = 'yes no up down left right on off stop go silence unknown'.split()\n\ndef find_train_audio_dir():\n    # 1. Check known directories\n    candidates = [\n        './train/audio',\n        './train/train/audio',\n        '/kaggle/working/train/audio',\n        '/kaggle/working/train/train/audio',\n        '../input/train/audio',\n        '../input/tensorflow-speech-recognition-challenge/train/audio',\n        '/kaggle/input/tensorflow-speech-recognition-challenge/train/audio',\n    ]\n    for c in candidates:\n        if os.path.isdir(c):\n            wavs = glob(os.path.join(c, '*', '*.wav'))\n            if len(wavs) > 0:\n                return c\n                \n    # 2. Search for directory containing label folders with wavs\n    for search_root in ['.', '/kaggle/working', '../input', '/kaggle/input']:\n        if os.path.exists(search_root):\n            for root, dirs, files in os.walk(search_root):\n                if 'test' in root:\n                    continue\n                if os.path.basename(root) in legal_labels:\n                    parent = os.path.dirname(root)\n                    wavs = glob(os.path.join(parent, '*', '*.wav'))\n                    if len(wavs) > 0:\n                        return parent\n    return None\n\ntrain_data_path = find_train_audio_dir()\n\n# If not extracted yet, extract train.7z\nif train_data_path is None:\n    candidate_7z = [\n        '../input/tensorflow-speech-recognition-challenge/train.7z',\n        '/kaggle/input/tensorflow-speech-recognition-challenge/train.7z',\n        '../input/train.7z',\n        './train.7z',\n        'train.7z'\n    ]\n    candidate_7z += glob('../input/**/train.7z', recursive=True)\n    candidate_7z += glob('/kaggle/input/**/train.7z', recursive=True)\n    \n    archive_path = None\n    for arc in candidate_7z:\n        if os.path.exists(arc):\n            archive_path = arc\n            break\n            \n    if archive_path:\n        print(f'Extracting {archive_path}...')\n        # Extract into current directory (train.7z contains train/audio/ inside it)\n        ret = os.system(f\"7z x '{archive_path}' -o. -y > /dev/null\")\n        if ret != 0:\n            os.system('apt-get update -qq && apt-get install -y -qq p7zip-full')\n            os.system(f\"7z x '{archive_path}' -o. -y > /dev/null\")\n        \n        train_data_path = find_train_audio_dir()\n        if train_data_path is None:\n            train_data_path = './train/audio'\n    else:\n        train_data_path = './train/audio'\n\nprint(f'Active train_data_path: {train_data_path}')\n\ntest_data_path = './test/audio'\nroot_path = r'..'\nout_path = r'.'\nmodel_path = r'.'\n","metadata":{"_cell_guid":"dc66e1df-f1eb-4df4-ba1a-65b9f1675953","_uuid":"4cc586519523b28d1d595716d8709ace9f27ac9c","collapsed":true,"execution":{"iopub.execute_input":"2026-09-03T21:53:29.906965Z","iopub.status.busy":"2026-09-03T21:53:29.906542Z","iopub.status.idle":"2026-09-03T21:53:29.911164Z","shell.execute_reply":"2026-09-03T21:53:29.910404Z"},"jupyter":{"outputs_hidden":true},"papermill":{"duration":0.009268,"end_time":"2026-09-03T21:53:29.912645+00:00","exception":false,"start_time":"2026-09-03T21:53:29.903377+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"5f9026f0","cell_type":"markdown","source":"Here are custom_fft and log_specgram functions written by __DavidS__.","metadata":{"_cell_guid":"e53561e4-1c98-44c0-9245-d87f7957faa5","_uuid":"d9a08781f22e574bb1eb0dc29adeb8dddebc8b51","papermill":{"duration":0.00281,"end_time":"2026-09-03T21:53:29.917815+00:00","exception":false,"start_time":"2026-09-03T21:53:29.915005+00:00","status":"completed"},"tags":[]}},{"id":"c702ea2f","cell_type":"code","source":"def custom_fft(y, fs):\n    T = 1.0 / fs\n    N = y.shape[0]\n    yf = fft(y)\n    xf = np.linspace(0.0, 1.0/(2.0*T), N//2)\n    # FFT is simmetrical, so we take just the first half\n    # FFT is also complex, to we take just the real part (abs)\n    vals = 2.0/N * np.abs(yf[0:N//2])\n    return xf, vals\n\ndef log_specgram(audio, sample_rate, window_size=20,\n                 step_size=10, eps=1e-10):\n    nperseg = int(round(window_size * sample_rate / 1e3))\n    noverlap = int(round(step_size * sample_rate / 1e3))\n    freqs, times, spec = signal.spectrogram(audio,\n                                    fs=sample_rate,\n                                    window='hann',\n                                    nperseg=nperseg,\n                                    noverlap=noverlap,\n                                    detrend=False)\n    return freqs, times, np.log(spec.T.astype(np.float32) + eps)","metadata":{"_cell_guid":"0fd0b579-8b6f-4253-bf3a-7f75115a42d6","_uuid":"e7ea2c277b6459e532721452ec3cd80d585eae1e","collapsed":true,"execution":{"iopub.execute_input":"2026-09-03T21:53:29.923867Z","iopub.status.busy":"2026-09-03T21:53:29.923179Z","iopub.status.idle":"2026-09-03T21:53:29.928708Z","shell.execute_reply":"2026-09-03T21:53:29.92807Z"},"jupyter":{"outputs_hidden":true},"papermill":{"duration":0.009965,"end_time":"2026-09-03T21:53:29.930052+00:00","exception":false,"start_time":"2026-09-03T21:53:29.920087+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"54cd537b","cell_type":"markdown","source":"Following is the utility function to grab all wav files inside train data folder.","metadata":{"_cell_guid":"c54cda36-777e-4129-bac1-af2d1ed2706e","_uuid":"5a04e71fe7e66e1a31835feebdfef4c63920faf8","papermill":{"duration":0.002334,"end_time":"2026-09-03T21:53:29.934828+00:00","exception":false,"start_time":"2026-09-03T21:53:29.932494+00:00","status":"completed"},"tags":[]}},{"id":"0628d9cf","cell_type":"code","source":"def list_wavs_fname(dirpath, ext='wav'):\n    print('Scanning audio files in:', dirpath)\n    if not os.path.exists(dirpath):\n        print(f'Error: Directory {dirpath} does not exist!')\n        return [], [], []\n    \n    # 1. Look directly in subfolders: dirpath/label/*.wav\n    fpaths = glob(os.path.join(dirpath, '*', f'*.{ext}'))\n    \n    # 2. If not found, search recursively (in case of nested folders)\n    if len(fpaths) == 0:\n        fpaths = [p for p in glob(os.path.join(dirpath, '**', f'*.{ext}'), recursive=True) if 'test' not in p]\n        \n    labels = []\n    fnames = []\n    for fpath in fpaths:\n        parent = os.path.basename(os.path.dirname(fpath))\n        fname = os.path.basename(fpath)\n        labels.append(parent)\n        fnames.append(fname)\n        \n    print(f'Total wav files found: {len(fpaths)}')\n    return fpaths, labels, fnames\n","metadata":{"_cell_guid":"956f3150-544d-46ed-b0ec-da1c1fb142b4","_uuid":"964d71a229e9d4560b9118fa1c80804ebf8d6be8","collapsed":true,"execution":{"iopub.execute_input":"2026-09-03T21:53:29.940861Z","iopub.status.busy":"2026-09-03T21:53:29.940143Z","iopub.status.idle":"2026-09-03T21:53:29.944993Z","shell.execute_reply":"2026-09-03T21:53:29.944484Z"},"jupyter":{"outputs_hidden":true},"papermill":{"duration":0.00923,"end_time":"2026-09-03T21:53:29.946394+00:00","exception":false,"start_time":"2026-09-03T21:53:29.937164+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"27ad2d86","cell_type":"markdown","source":"__pad_audio__ will pad audios that are less than 16000(1 second) with 0s to make them all have the same length.\n\n__chop_audio__ will chop audios that are larger than 16000(eg. wav files in background noises folder) to 16000 in length. In addition, it will create several chunks out of one large wav files given the parameter 'num'.\n\n__label_transform__ transform labels into dummies values. It's used in combination with softmax to predict the label.","metadata":{"_cell_guid":"41025a55-8497-43cf-b316-003af7d9d19f","_uuid":"fc18e87793888952e81a867dd95b1dcc455f9932","papermill":{"duration":0.002316,"end_time":"2026-09-03T21:53:29.951167+00:00","exception":false,"start_time":"2026-09-03T21:53:29.948851+00:00","status":"completed"},"tags":[]}},{"id":"4e74e044","cell_type":"code","source":"def pad_audio(samples):\n    if len(samples) >= L: return samples\n    else: return np.pad(samples, pad_width=(L - len(samples), 0), mode='constant', constant_values=(0, 0))\n\ndef chop_audio(samples, L=16000, num=20):\n    for i in range(num):\n        beg = np.random.randint(0, len(samples) - L)\n        yield samples[beg: beg + L]\n\ndef augment_waveform(samples, shift_range=1600, noise_factor=0.005):\n    \"\"\"Applies random time-shift and subtle noise injection to waveform.\"\"\"\n    shift = np.random.randint(-shift_range, shift_range)\n    augmented = np.roll(samples, shift)\n    if shift > 0:\n        augmented[:shift] = 0\n    elif shift < 0:\n        augmented[shift:] = 0\n    if np.random.rand() > 0.5:\n        noise = np.random.randn(len(augmented))\n        augmented = augmented + noise_factor * noise * np.max(np.abs(augmented))\n    return augmented\n\ndef spec_augment(spec, num_time_masks=1, num_freq_masks=1, time_mask_max=10, freq_mask_max=8):\n    \"\"\"Applies SpecAugment: Frequency and Time masking on 2D Spectrogram (99, 81).\"\"\"\n    spec_aug = spec.copy()\n    num_times, num_freqs = spec_aug.shape[:2]\n    for _ in range(num_freq_masks):\n        f = np.random.randint(0, freq_mask_max)\n        f0 = np.random.randint(0, max(1, num_freqs - f))\n        spec_aug[:, f0:f0 + f] = 0\n    for _ in range(num_time_masks):\n        t = np.random.randint(0, time_mask_max)\n        t0 = np.random.randint(0, max(1, num_times - t))\n        spec_aug[t0:t0 + t, :] = 0\n    return spec_aug\n\ndef label_transform(labels):\n    nlabels = []\n    for label in labels:\n        if label == '_background_noise_':\n            nlabels.append('silence')\n        elif label not in legal_labels:\n            nlabels.append('unknown')\n        else:\n            nlabels.append(label)\n    cat_series = pd.Categorical(nlabels, categories=legal_labels)\n    return pd.get_dummies(cat_series)\n","metadata":{"_cell_guid":"200c34a1-851a-4447-9ff7-b4e541f090c6","_uuid":"94e40aef3899acfd3ed85557caa66fee5dd47db2","collapsed":true,"execution":{"iopub.execute_input":"2026-09-03T21:53:29.957465Z","iopub.status.busy":"2026-09-03T21:53:29.957018Z","iopub.status.idle":"2026-09-03T21:53:29.962633Z","shell.execute_reply":"2026-09-03T21:53:29.961988Z"},"jupyter":{"outputs_hidden":true},"papermill":{"duration":0.010082,"end_time":"2026-09-03T21:53:29.964019+00:00","exception":false,"start_time":"2026-09-03T21:53:29.953937+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"77096b8a","cell_type":"markdown","source":"Next, we use functions declared above to generate x_train and y_train.\nlabel_index is the index used by pandas to create dummy values, we need to save it for later use.","metadata":{"_cell_guid":"dae2a45a-f7ab-4e84-bc73-688eda6eca8e","_uuid":"267314ef41c459c8b6ab903d721980fdd62b4106","papermill":{"duration":0.002358,"end_time":"2026-09-03T21:53:29.969096+00:00","exception":false,"start_time":"2026-09-03T21:53:29.966738+00:00","status":"completed"},"tags":[]}},{"id":"f15e6bc3","cell_type":"code","source":"from concurrent.futures import ThreadPoolExecutor\n\nfpaths, labels, fnames = list_wavs_fname(train_data_path)\nassert len(fpaths) > 0, f\"No wav files found in {train_data_path}!\"\n\nnew_sample_rate = 8000\n\ndef process_single_audio(item):\n    fpath, label = item\n    try:\n        sample_rate, samples = wavfile.read(fpath)\n    except Exception:\n        return [], []\n    samples = pad_audio(samples)\n    if len(samples) > 16000:\n        n_samples = list(chop_audio(samples))\n    else:\n        n_samples = [samples]\n        \n    specs, lbls = [], []\n    for s in n_samples:\n        # Original spectrogram\n        if sample_rate == 16000 and new_sample_rate == 8000:\n            resampled = s[::2]\n        else:\n            resampled = signal.resample(s, int(new_sample_rate / sample_rate * s.shape[0]))\n        _, _, specgram = log_specgram(resampled, sample_rate=new_sample_rate)\n        specs.append(specgram)\n        lbls.append(label)\n        \n        # Data Augmentation: add augmented copies for the 10 target commands\n        # This effectively balances target words against dominant 'unknown' samples\n        if label in legal_labels[:10]:\n            aug_s = augment_waveform(s)\n            aug_resampled = aug_s[::2] if sample_rate == 16000 and new_sample_rate == 8000 else signal.resample(aug_s, int(new_sample_rate / sample_rate * aug_s.shape[0]))\n            _, _, aug_spec = log_specgram(aug_resampled, sample_rate=new_sample_rate)\n            specs.append(spec_augment(aug_spec))\n            lbls.append(label)\n            \n    return specs, lbls\n\nprint(\"Extracting spectrograms with data augmentation across CPU threads...\")\nitems = list(zip(fpaths, labels))\n\nx_train = []\ny_train = []\n\nwith ThreadPoolExecutor(max_workers=4) as executor:\n    results = executor.map(process_single_audio, items)\n    for specs, lbls in results:\n        x_train.extend(specs)\n        y_train.extend(lbls)\n\nx_train = np.array(x_train, dtype=np.float32)\nx_train = x_train.reshape(tuple(list(x_train.shape) + [1]))\ny_train = label_transform(y_train)\nlabel_index = y_train.columns.values\ny_train = y_train.values\ny_train = np.array(y_train, dtype=np.float32)\n\ndel fpaths, labels, fnames, items\ngc.collect()\nprint(f\"Ready! x_train shape: {x_train.shape}, y_train shape: {y_train.shape}\")\n","metadata":{"_cell_guid":"4c8d9fdf-ea3e-45fa-b7ef-52542c70b9db","_uuid":"81bc9722dfb036c73721ae44829d429489662e75","collapsed":true,"execution":{"iopub.execute_input":"2026-09-03T21:53:29.974857Z","iopub.status.busy":"2026-09-03T21:53:29.974586Z","iopub.status.idle":"2026-09-03T21:53:30.188316Z","shell.execute_reply":"2026-09-03T21:53:30.187449Z"},"jupyter":{"outputs_hidden":true},"papermill":{"duration":0.218254,"end_time":"2026-09-03T21:53:30.18979+00:00","exception":false,"start_time":"2026-09-03T21:53:29.971536+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"7503b999","cell_type":"markdown","source":"CNN declared below.\nThe specgram created will be of shape (99, 81), but in order to fit into Conv2D layer, we need to reshape it.","metadata":{"_cell_guid":"56921cf3-1269-4b29-876d-abdd31eb150a","_uuid":"a87a77b76c42da61ca0bec395c71bef795a9e928","papermill":{"duration":0.002587,"end_time":"2026-09-03T21:53:30.195169+00:00","exception":false,"start_time":"2026-09-03T21:53:30.192582+00:00","status":"completed"},"tags":[]}},{"id":"5003948c","cell_type":"code","source":"from keras.layers import (\n    Input, Conv2D, BatchNormalization, ReLU, Add, \n    MaxPooling2D, Dropout, GlobalAveragePooling2D, Dense\n)\nfrom keras import models, optimizers, losses\nfrom keras.callbacks import ReduceLROnPlateau, ModelCheckpoint, EarlyStopping\n\ndef resnet_block(x, filters, stride=1):\n    \"\"\"Residual Block with skip connection.\"\"\"\n    res = x\n    out = Conv2D(filters, kernel_size=3, strides=stride, padding='same')(x)\n    out = BatchNormalization()(out)\n    out = ReLU()(out)\n    out = Conv2D(filters, kernel_size=3, strides=1, padding='same')(out)\n    out = BatchNormalization()(out)\n    \n    if stride != 1 or res.shape[-1] != filters:\n        res = Conv2D(filters, kernel_size=1, strides=stride, padding='same')(res)\n        res = BatchNormalization()(res)\n        \n    out = Add()([out, res])\n    out = ReLU()(out)\n    return out\n\ndef build_audio_resnet(input_shape=(99, 81, 1), num_classes=12):\n    inp = Input(shape=input_shape)\n    \n    # Stem\n    x = Conv2D(32, kernel_size=3, strides=1, padding='same')(inp)\n    x = BatchNormalization()(x)\n    x = ReLU()(x)\n    \n    # Residual stages (32 -> 64 -> 128 -> 256)\n    x = resnet_block(x, 32)\n    x = MaxPooling2D(pool_size=2)(x)\n    x = Dropout(0.2)(x)\n    \n    x = resnet_block(x, 64)\n    x = MaxPooling2D(pool_size=2)(x)\n    x = Dropout(0.2)(x)\n    \n    x = resnet_block(x, 128)\n    x = MaxPooling2D(pool_size=2)(x)\n    x = Dropout(0.3)(x)\n    \n    x = resnet_block(x, 256)\n    x = GlobalAveragePooling2D()(x)\n    \n    # Classifier Head\n    x = Dense(128, activation='relu')(x)\n    x = BatchNormalization()(x)\n    x = Dropout(0.4)(x)\n    out = Dense(num_classes, activation='softmax')(x)\n    \n    return models.Model(inputs=inp, outputs=out)\n\nmodel = build_audio_resnet(input_shape=(99, 81, 1), num_classes=12)\n\nopt = optimizers.Adam(learning_rate=0.001)\nloss_fn = losses.CategoricalCrossentropy(label_smoothing=0.05)\nmodel.compile(optimizer=opt, loss=loss_fn, metrics=['accuracy'])\nmodel.summary()\n\n# Training Callbacks\nmodel_save_file = os.path.join(model_path, 'cnn.keras')\ncallbacks = [\n    ReduceLROnPlateau(\n        monitor='val_loss', \n        factor=0.5, \n        patience=2, \n        min_lr=1e-5, \n        verbose=1\n    ),\n    ModelCheckpoint(\n        model_save_file, \n        monitor='val_accuracy', \n        save_best_only=True, \n        verbose=1\n    ),\n    EarlyStopping(\n        monitor='val_accuracy', \n        patience=5, \n        restore_best_weights=True, \n        verbose=1\n    )\n]\n\nBATCH_SIZE = 64\nEPOCHS = 15\n\nprint(f\"Training {len(x_train)} samples with ResNet-CNN, batch_size={BATCH_SIZE}...\")\nx_train, x_valid, y_train, y_valid = train_test_split(x_train, y_train, test_size=0.1, random_state=2017)\nhistory = model.fit(\n    x_train, y_train, \n    batch_size=BATCH_SIZE, \n    validation_data=(x_valid, y_valid), \n    epochs=EPOCHS, \n    shuffle=True, \n    callbacks=callbacks,\n    verbose=2\n)\n\nif os.path.exists(model_save_file):\n    print(f\"Loading best checkpoint from {model_save_file}...\")\n    model = models.load_model(model_save_file)\n","metadata":{"_cell_guid":"b97e8887-b593-4d88-95c8-fc8f1dd5ca72","_uuid":"60af394ad8e91fb868ea32dbb6ac6a725b5935c9","collapsed":true,"execution":{"iopub.execute_input":"2026-09-03T21:53:30.20181Z","iopub.status.busy":"2026-09-03T21:53:30.201045Z","iopub.status.idle":"2026-09-03T21:53:32.732209Z","shell.execute_reply":"2026-09-03T21:53:32.731065Z"},"jupyter":{"outputs_hidden":true},"papermill":{"duration":2.535778,"end_time":"2026-09-03T21:53:32.733481+00:00","exception":true,"start_time":"2026-09-03T21:53:30.197703+00:00","status":"failed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"58eaaa80","cell_type":"markdown","source":"Test data is way too large to fit in RAM, we need to process them one by one.\nGenerator test_data_generator will create batches of test wav files to feed into CNN.","metadata":{"_cell_guid":"060811f5-34ac-4fc2-92ba-c4276606c2a0","_uuid":"2e3fa8d9706f47e69d0b74afcb68280f8a5de706","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"acd2e782","cell_type":"code","source":"def find_test_audio_dir():\n    candidates = [\n        './test/audio',\n        './test/test/audio',\n        '/kaggle/working/test/audio',\n        '../input/test/audio',\n        '../input/tensorflow-speech-recognition-challenge/test/audio',\n        '/kaggle/input/tensorflow-speech-recognition-challenge/test/audio',\n    ]\n    for c in candidates:\n        if os.path.isdir(c):\n            wavs = glob(os.path.join(c, '*.wav'))\n            if len(wavs) > 0:\n                return c\n    return None\n\ntest_data_path = find_test_audio_dir()\n\n# Extract test.7z only when not found\nif test_data_path is None:\n    candidate_test_7z = [\n        '../input/tensorflow-speech-recognition-challenge/test.7z',\n        '/kaggle/input/tensorflow-speech-recognition-challenge/test.7z',\n        '../input/test.7z',\n        './test.7z'\n    ] + glob('../input/**/test.7z', recursive=True) + glob('/kaggle/input/**/test.7z', recursive=True)\n    for arc in candidate_test_7z:\n        if os.path.exists(arc):\n            print(f'Extracting test archive: {arc} (158,538 files, takes ~2-3 mins)...')\n            # Extract without silent mute so progress is clear\n            os.system(f\"7z x '{arc}' -o. -y -bso0\")\n            print('Test extraction finished!')\n            break\n    test_data_path = find_test_audio_dir()\n\nif test_data_path is None:\n    test_data_path = './test/audio'\nprint(f'Active test_data_path: {test_data_path}')\n\ndef extract_spec(samples):\n    \"\"\"Helper to convert audio samples to resampled spectrogram.\"\"\"\n    if len(samples) < 16000:\n        samples = pad_audio(samples)\n    elif len(samples) > 16000:\n        samples = samples[:16000]\n    resampled = samples[::2]\n    _, _, specgram = log_specgram(resampled, sample_rate=8000)\n    return specgram\n\ndef tta_predict_batch(batch_samples, model):\n    \"\"\"\n    TTA: averages predictions over 3 acoustic variations:\n      1) Original audio\n      2) Shift left by 800 samples (-50ms)\n      3) Shift right by 800 samples (+50ms)\n    \"\"\"\n    batch_orig = []\n    batch_left = []\n    batch_right = []\n    \n    for s in batch_samples:\n        batch_orig.append(extract_spec(s))\n        batch_left.append(extract_spec(np.roll(s, -800)))\n        batch_right.append(extract_spec(np.roll(s, 800)))\n        \n    imgs_orig = np.expand_dims(np.array(batch_orig, dtype=np.float32), -1)\n    imgs_left = np.expand_dims(np.array(batch_left, dtype=np.float32), -1)\n    imgs_right = np.expand_dims(np.array(batch_right, dtype=np.float32), -1)\n    \n    p0 = model.predict(imgs_orig, verbose=0)\n    p1 = model.predict(imgs_left, verbose=0)\n    p2 = model.predict(imgs_right, verbose=0)\n    \n    return (p0 + p1 + p2) / 3.0\n","metadata":{"_cell_guid":"7dfe0801-a636-4123-8367-ab2f19c97800","_uuid":"646b6bcfbde7eae53cd8822b8838c575859e51ce","collapsed":true,"jupyter":{"outputs_hidden":true},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null},{"id":"7602aca5","cell_type":"markdown","source":"We use the trained model to predict the test data's labels.\nHowever, since Kaggle doesn't provide test data, the following sections won't be executed here.","metadata":{"_cell_guid":"22992a27-deda-4a35-b34c-4aa87ad173ec","_uuid":"c6d2516a6d5bd3c6a108d7e28565edaa65830958","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"1acd46ec","cell_type":"code","source":"import shutil\n\ndel x_train, y_train\ngc.collect()\n\nprint(\"Starting Test-Time Augmentation (TTA) inference on 158k test files...\")\nfpaths = glob(os.path.join(test_data_path, '*wav'))\nif len(fpaths) == 0:\n    fpaths = glob(os.path.join(test_data_path, '**', '*wav'), recursive=True)\n\nbatch_size = 64\nall_preds = []\nall_fnames = []\ntotal_files = len(fpaths)\n\nfor i in range(0, total_files, batch_size):\n    batch_paths = fpaths[i:i + batch_size]\n    batch_raw = [pad_audio(wavfile.read(p)[1]) for p in batch_paths]\n    \n    avg_probs = tta_predict_batch(batch_raw, model)\n    pred_indices = np.argmax(avg_probs, axis=1)\n    \n    all_preds.extend([label_index[idx] for idx in pred_indices])\n    all_fnames.extend([os.path.basename(p) for p in batch_paths])\n    \n    # Progress logging every 20,000 files\n    if (i + batch_size) % 20000 < batch_size or (i + batch_size) >= total_files:\n        current_count = min(i + batch_size, total_files)\n        print(f\"Inference progress: {current_count:,}/{total_files:,} files ({current_count/total_files*100:.1f}%)\")\n\ndf = pd.DataFrame({'fname': all_fnames, 'label': all_preds})\n\n# Save both submission files\nsub_path = os.path.join(out_path, 'submission.csv')\ndf.to_csv(sub_path, index=False)\ndf.to_csv(os.path.join(out_path, 'sub.csv'), index=False)\nprint(f\"TTA Submission successfully saved to {sub_path}. Total predictions: {len(df):,}\")\n\n# Clean up extracted audio files so Kaggle output remains ~3 MB\nprint(\"Cleaning up temporary extracted audio directories...\")\nfor temp_dir in ['./train', './test']:\n    if os.path.exists(temp_dir):\n        shutil.rmtree(temp_dir, ignore_errors=True)\nprint(\"Done! Clean submission ready for scoring.\")\n","metadata":{"_cell_guid":"7fa8feb8-236e-46c5-8432-014f7e27484d","_uuid":"56194039cac16f5d86a322e67641cbeafda9857d","collapsed":true,"jupyter":{"outputs_hidden":true},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null}]}