{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Optimizing dataloading\n\nOne of the biggest bottlenecks when working on AI problems is the speed of training loops. There are 2 factors to consider: (1) how fast can we load and prepare data into the GPU (2) how fast can the model process that data on the GPU. Since these two steps can be done in parallel, it's possible to have 100% GPU utilization during training so long as (1) is faster than (2). In this notebook, I will explore how to optimize (1).\n\nSpoilers:\n- fetching 5 sec crops instead of full files speeds up loading by ~2x\n- increasing # of workers to 4 on this machine is another ~2x speedup\n- loading pre-decoded data is another ~3-10x speedup (10x on interactive sessions, 3x on the \"Latest Container Image\")\n\nOther things that don't matter much:\n- increasing batch sizes (slightly better)\n- preloading ogg into memory first (neutral) (Thanks to [lhanhsin](https://www.kaggle.com/lhanhsin) for [this idea and the code example](https://www.kaggle.com/competitions/birdclef-2023/discussion/397086#2203570)!)\n- increasing the prefetch factor (barely better to neutral)\n\nBatch preparation is often slowed down by disk reads, decompressing, and decoding. So why not skip all that and preload the entire decoded dataset into CPU RAM or (even better) GPU RAM? The BirdCLEF 2013 dataset is 5 GB of compressed OGG files that becomes 80 GB of raw audio data when decoded. Kaggle P100 machines have 16 GB of GPU RAM, 13 GB of CPU RAM, and 73 GB of disk space. So we can't fit the uncompressed dataset into memory or the disk!\n\nAn important aspect of this dataset is that it contains audio recordings that are a few seconds to 45 minutes long. Since batches of data are usually fixed length, a common strategy is to take crops from these recordings (usually 5-30 seconds)—which suggests that we might save time by just fetching crops directly instead of the entire recordings when loading data. It turns out that doing that speeds up dataloading by 2x. I see this consistently when reading encoded or decoded data from disk, or encoded data from memory.\n\nOf course, another simple solution would be to throw away part of the dataset via downscaling or filtering until it all fits into memory uncompressed. Some examples of this: taking the first 5 seconds of each recording, converting all recordings to low resolution spectrograms, using only recordings from Kenya, downsampling from 32 KHz to 16 Khz, ignoring the largest files, etc. Those techniques have tradeoffs and can be tested, but even aggressive downsampling strategies could still hit size limits if 10x more data is added from other sources, leaving us with the same problem.\n\nHowever, we can use a filtered version of the dataset to determine whether reading from disk or decoding is slower. Using only recordings that occurred in Kenya (\\~10% of the dataset), I first decoded the audio into numpy arrays (~8 GB). Then I found that it was 3-10x faster to read the numpy arrays from disk than to read encoded data from memory.\n\nEnough talk, let's get to the experiments!","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport torchaudio\nimport torch\nimport random\nimport time\nfrom tqdm import tqdm\nfrom pathlib import Path\nimport io\nimport os\nimport numpy as np","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-04-01T03:08:56.806800Z","iopub.execute_input":"2023-04-01T03:08:56.807297Z","iopub.status.idle":"2023-04-01T03:08:59.024483Z","shell.execute_reply.started":"2023-04-01T03:08:56.807257Z","shell.execute_reply":"2023-04-01T03:08:59.022941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torchaudio.get_audio_backend()","metadata":{"execution":{"iopub.status.busy":"2023-04-01T03:08:59.026833Z","iopub.execute_input":"2023-04-01T03:08:59.027541Z","iopub.status.idle":"2023-04-01T03:08:59.037529Z","shell.execute_reply.started":"2023-04-01T03:08:59.027499Z","shell.execute_reply":"2023-04-01T03:08:59.036069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/birdclef-2023-num-frames/train_metadata_with_kenya_and_num_frames.csv')\ndf_kenya = df[df.is_in_kenya].reset_index()\n\nTRAIN_DIR = '/kaggle/input/birdclef-2023/train_audio/'\nDECODE_DIR = './decoded_train_audio/'\nWINDOW_SIZE = 32_000 * 5 # 5 seconds of frames","metadata":{"execution":{"iopub.status.busy":"2023-04-01T03:08:59.039076Z","iopub.execute_input":"2023-04-01T03:08:59.039492Z","iopub.status.idle":"2023-04-01T03:08:59.241016Z","shell.execute_reply.started":"2023-04-01T03:08:59.039455Z","shell.execute_reply":"2023-04-01T03:08:59.239282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Experiments","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\n\nclass AudioDataset(Dataset):\n    def __init__(self, df, fetch_crop=False, decode_to_device=None):\n        self.df = df\n        self.fetch_crop = fetch_crop\n        self.decode_to_device = decode_to_device\n\n        if self.decode_to_device == 'cpu':\n            self.audio_cache = {}\n            for filename in tqdm(self.df.filename, desc='Caching audio files'):\n                with open(TRAIN_DIR + filename, 'rb') as f:\n                    self.audio_cache[filename] = f.read()\n        elif self.decode_to_device == 'disk':\n            os.makedirs(DECODE_DIR, exist_ok=True)\n            for filename in tqdm(self.df.filename, desc='Creating directories for decoded audio'):\n                path = os.path.join(DECODE_DIR, filename[:-4]) \n                os.makedirs(path, exist_ok=True)\n            for filename in tqdm(self.df.filename, desc='Decoding audio'):\n                audio = torchaudio.load(TRAIN_DIR + filename)[0]\n                np.save(DECODE_DIR + filename + '.npy', audio)\n\n    def __len__(self):\n        return len(self.df)\n\n    def _load_audio(self, filename, frame_offset=0, num_frames=-1):\n        if self.decode_to_device == 'cpu':\n            with io.BytesIO(self.audio_cache[filename]) as fh:\n                return torchaudio.load(fh, frame_offset, num_frames)[0]\n        elif self.decode_to_device == 'disk':\n            mmap_data = np.load(DECODE_DIR + filename + '.npy', mmap_mode='r')\n            frame_offset_end = None if num_frames == -1 else (frame_offset + num_frames)\n            return torch.tensor(mmap_data[:, frame_offset:frame_offset_end])\n        else:\n            return torchaudio.load(TRAIN_DIR + filename, frame_offset, num_frames)[0]\n\n    def __getitem__(self, i):\n        filename = self.df.filename[i]\n        if self.fetch_crop:\n            num_frames = self.df.num_frames[i]\n            frame_offset = random.randint(0, max(0, num_frames - WINDOW_SIZE))\n            # load only the crop from the file\n            crop = self._load_audio(filename, frame_offset, WINDOW_SIZE)[0]\n        else:\n            # load the entire file, then crop\n            audio = self._load_audio(filename)[0]\n            frame_offset = random.randint(0, max(0, len(audio) - WINDOW_SIZE))\n            crop = audio[frame_offset:frame_offset+WINDOW_SIZE]\n\n        if len(crop) < WINDOW_SIZE:\n            crop = torch.concat([crop, torch.zeros(WINDOW_SIZE - len(crop))])\n\n        return crop\n\n\nfull_dataset = AudioDataset(df, fetch_crop=False)\ncrop_dataset = AudioDataset(df, fetch_crop=True)\ncrop_preloaded_dataset = AudioDataset(df, fetch_crop=True, decode_to_device='cpu')\n\nkeyna_dataset = AudioDataset(df_kenya, fetch_crop=False)\nkeyna_crop_dataset = AudioDataset(df_kenya, fetch_crop=True)\nkeyna_decoded_dataset = AudioDataset(df_kenya, fetch_crop=False, decode_to_device='disk')\nkeyna_decoded_crop_dataset = AudioDataset(df_kenya, fetch_crop=True)\nkeyna_decoded_crop_dataset.decode_to_device = 'disk' # rely on data decoded by keyna_decoded_dataset","metadata":{"execution":{"iopub.status.busy":"2023-04-01T03:09:01.122645Z","iopub.execute_input":"2023-04-01T03:09:01.123140Z","iopub.status.idle":"2023-04-01T03:09:03.621417Z","shell.execute_reply.started":"2023-04-01T03:09:01.123098Z","shell.execute_reply":"2023-04-01T03:09:03.619667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"full_dataloaders = [\n    DataLoader(full_dataset, batch_size=32),\n    DataLoader(full_dataset, batch_size=32, num_workers=2),\n    DataLoader(full_dataset, batch_size=32, num_workers=4),\n    \n    DataLoader(crop_dataset, batch_size=32),\n    DataLoader(crop_dataset, batch_size=32, num_workers=4),\n\n    DataLoader(crop_dataset, batch_size=8, num_workers=4),\n    DataLoader(crop_dataset, batch_size=16, num_workers=4),\n    DataLoader(crop_dataset, batch_size=64, num_workers=4),\n    DataLoader(crop_dataset, batch_size=128, num_workers=4),\n\n    DataLoader(crop_dataset, batch_size=32, num_workers=4, prefetch_factor=3),\n    \n    DataLoader(crop_preloaded_dataset, batch_size=32, num_workers=4),\n]\n\nkeyna_dataloaders = [\n    DataLoader(keyna_dataset, batch_size=32, num_workers=4),\n    DataLoader(keyna_crop_dataset, batch_size=32, num_workers=4),\n    DataLoader(keyna_decoded_dataset, batch_size=32, num_workers=4),\n    DataLoader(keyna_decoded_dataset, batch_size=32, num_workers=4),\n    DataLoader(keyna_decoded_crop_dataset, batch_size=32, num_workers=4),\n]","metadata":{"execution":{"iopub.status.busy":"2023-04-01T03:09:03.624391Z","iopub.execute_input":"2023-04-01T03:09:03.625286Z","iopub.status.idle":"2023-04-01T03:09:03.636319Z","shell.execute_reply.started":"2023-04-01T03:09:03.625225Z","shell.execute_reply":"2023-04-01T03:09:03.634734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_dataloaders(dataloaders):\n    for dataloader in dataloaders:\n        desc = \\\n            'fetching crops:' + \\\n            str(dataloader.dataset.fetch_crop) + \\\n            '; decode device:' + \\\n            str(dataloader.dataset.decode_to_device) + \\\n            '; batch size:' + \\\n            str(dataloader.batch_size) + \\\n            '; workers:'  + \\\n            str(dataloader.num_workers) + \\\n            '; prefetch factor:' + \\\n            str(dataloader.prefetch_factor)\n        start = time.time()\n        for batch in tqdm(dataloader, desc):\n            batch[0].sum()\n        print(\"total time:\", time.time() - start)\n        print()\n\nprint('=== full dataset dataloaders ===')\ntest_dataloaders(full_dataloaders)\n\nprint('=== just keyna dataset dataloaders ===')\ntest_dataloaders(keyna_dataloaders)","metadata":{"execution":{"iopub.status.busy":"2023-04-01T03:09:03.638289Z","iopub.execute_input":"2023-04-01T03:09:03.639375Z","iopub.status.idle":"2023-04-01T03:09:35.214598Z","shell.execute_reply.started":"2023-04-01T03:09:03.639332Z","shell.execute_reply":"2023-04-01T03:09:35.212969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}