{"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":"gpu","dataSources":[{"sourceId":73047,"databundleVersionId":8149390,"sourceType":"competition"},{"sourceId":8172942,"sourceType":"datasetVersion","datasetId":4837321}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import gc\nimport ctypes\nimport torch\n\ndef clean_memory():\n    gc.collect()\n    ctypes.CDLL(\"libc.so.6\").malloc_trim(0)\n    torch.cuda.empty_cache()\nclean_memory()","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:47:35.890939Z","iopub.execute_input":"2024-04-22T04:47:35.891880Z","iopub.status.idle":"2024-04-22T04:47:35.974601Z","shell.execute_reply.started":"2024-04-22T04:47:35.891849Z","shell.execute_reply":"2024-04-22T04:47:35.973458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install jiwer","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:47:37.259084Z","iopub.execute_input":"2024-04-22T04:47:37.259956Z","iopub.status.idle":"2024-04-22T04:47:51.811618Z","shell.execute_reply.started":"2024-04-22T04:47:37.259923Z","shell.execute_reply":"2024-04-22T04:47:51.810642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\nimport os\nimport math\nimport time\nimport numpy as np\nimport torch\n\nimport pandas as pd\nimport pickle\n\ndef get_logger(filename: str):\n    from logging import getLogger, INFO, StreamHandler, FileHandler, Formatter\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=f\"{filename}.log\")\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\ndef seed_everything(seed: int):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef timeSince(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\ndef shuffle_df_chunk_by_chunk(df, chunk_size=1000):\n    new_df = df[:len(df)//chunk_size*chunk_size].copy()\n    num_sets = len(new_df) // chunk_size\n    indices = list(range(num_sets))\n    np.random.shuffle(indices)\n\n    shuffled_dfs = []\n    for i in indices:\n        start_idx = i * chunk_size\n        end_idx = (i + 1) * chunk_size\n        shuffled_dfs.append(new_df.iloc[start_idx:end_idx])\n        \n    shuffled_df = pd.concat(shuffled_dfs)\n    shuffled_df.reset_index(drop=True, inplace=True)\n    return shuffled_df\n\ndef pickle_file(path, contents):\n    with open(path, 'wb') as f:\n        pickle.dump(contents, f)\n\ndef load_pickle(path):\n    with open(path, 'rb') as f:\n        contents = pickle.load(f)\n    return contents","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:47:55.178580Z","iopub.execute_input":"2024-04-22T04:47:55.178955Z","iopub.status.idle":"2024-04-22T04:47:55.511399Z","shell.execute_reply.started":"2024-04-22T04:47:55.178921Z","shell.execute_reply":"2024-04-22T04:47:55.510542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport sys\nfrom pathlib import Path\nimport pickle\nimport time\nimport random\nimport math\nfrom functools import partial\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nfrom torch.optim import AdamW\nfrom torch.utils.data import DataLoader, Dataset\n\nfrom transformers import get_cosine_schedule_with_warmup\n\nimport librosa\n\nfrom datasets import load_metric\n\nsys.path.append('../')\n# from utils import *\n\nfrom transformers import (\n    Wav2Vec2ForCTC,\n    Wav2Vec2Processor,\n    Wav2Vec2FeatureExtractor,\n)\n\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:47:56.967249Z","iopub.execute_input":"2024-04-22T04:47:56.967849Z","iopub.status.idle":"2024-04-22T04:48:00.823498Z","shell.execute_reply.started":"2024-04-22T04:47:56.967818Z","shell.execute_reply":"2024-04-22T04:48:00.822466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:48:02.548297Z","iopub.execute_input":"2024-04-22T04:48:02.549227Z","iopub.status.idle":"2024-04-22T04:48:02.553217Z","shell.execute_reply.started":"2024-04-22T04:48:02.549190Z","shell.execute_reply":"2024-04-22T04:48:02.552202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"##### device = 'cuda' if torch.cuda.is_available() else 'cpu'\n<!-- os.environ['CUDA_VISIBLE_DEVICES'] = '0'\n\n\nclass config:\n    stage = '1'\n#     audio_dir = \"../data/bengaliai-speech/train_mp3s\"\n    audio_dir = \"/kaggle/input/ben10/ben10/16_kHz_train_audio\"\n    model = \"ai4bharat/indicwav2vec_v1_bengali\"\n    language_model = \"arijitx/wav2vec2-xls-r-300m-bengali\"\n    seed = 42\n    train_bs = 1    #6\n    valid_bs = 2    #12\n    lr = 2e-4\n    weight_decay = 1e-5\n    n_folds = 40\n    epochs = 1 #10   #10\n    apex = True\n    \n    print_freq = 1000\n    num_workers = 16\n    sampling_rate = 16000\n\ntry:\n    os.mkdir(f'stage{config.stage}')\nexcept:\n    pass\n\nOUT_PATH = f'stage{config.stage}/'\n\n\nLOGGER = get_logger(OUT_PATH+'train')\nseed_everything(config.seed) -->","metadata":{}},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\nos.environ['CUDA_VISIBLE_DEVICES'] = '0'\n\nclass config:\n    stage = '1'\n#     audio_dir = \"../data/bengaliai-speech/train_mp3s\"\n    audio_dir = \"/kaggle/input/ben10/ben10/16_kHz_train_audio\"\n    model = \"ai4bharat/indicwav2vec_v1_bengali\"\n    language_model = \"arijitx/wav2vec2-xls-r-300m-bengali\"\n    seed = 42\n    train_bs = 1    #6\n    valid_bs = 2    #12\n    lr = 2e-4\n    weight_decay = 1e-5\n    n_folds = 40\n    epochs = 15 #10   #10\n    apex = True\n    \n    print_freq = 1000\n    num_workers = 16\n    sampling_rate = 16000\n\ntry:\n    os.mkdir(f'stage{config.stage}')\nexcept:\n    pass\n\nOUT_PATH = f'stage{config.stage}/'\n\n\nLOGGER = get_logger(OUT_PATH+'train')\nseed_everything(config.seed)","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:48:09.168191Z","iopub.execute_input":"2024-04-22T04:48:09.168810Z","iopub.status.idle":"2024-04-22T04:48:09.204002Z","shell.execute_reply.started":"2024-04-22T04:48:09.168779Z","shell.execute_reply":"2024-04-22T04:48:09.203149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_audio(mp3_path, target_sr=16000):\n    audio, sr = librosa.load(mp3_path, sr=32000)\n    audio_array = librosa.resample(audio, orig_sr=sr, target_sr=target_sr)\n    return audio_array\n\nclass ASRDataset(Dataset):\n    def __init__(self, df, config, processor, is_test=False):\n        self.df = df\n        self.config = config\n        self.is_test = is_test\n        self.processor = processor\n    \n    def __getitem__(self, idx):\n        audio = read_audio(self.df.loc[idx]['path'])\n        audio = self.processor(\n            audio, \n            sampling_rate=self.config.sampling_rate\n        ).input_values[0]\n        \n        if self.is_test:\n            return {'audio': audio, 'label': -1}\n        else:\n            with self.processor.as_target_processor():\n                labels = self.processor(self.df.loc[idx]['sentence']).input_ids\n            return {'audio': audio, 'label': labels}\n        \n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:48:11.224976Z","iopub.execute_input":"2024-04-22T04:48:11.225360Z","iopub.status.idle":"2024-04-22T04:48:11.234768Z","shell.execute_reply.started":"2024-04-22T04:48:11.225330Z","shell.execute_reply":"2024-04-22T04:48:11.233811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def ctc_data_collator(batch, processor):\n    input_features = [{\"input_values\": sample[\"audio\"]} for sample in batch]\n    label_features = [{\"input_ids\": sample[\"label\"]} for sample in batch]\n    batch = processor.pad(\n        input_features,\n        padding=True,\n        return_tensors=\"pt\",\n    )\n    with processor.as_target_processor():\n        labels_batch = processor.pad(\n            label_features,\n            padding=True,\n            return_tensors=\"pt\",\n        )\n        \n    labels = labels_batch[\"input_ids\"].masked_fill(labels_batch.attention_mask.ne(1), -100)\n    batch[\"labels\"] = labels\n    return batch\n\ndef compute_wer(pred, labels, processor, metric):\n    pred_logits = pred.logits.cpu().detach().numpy()\n    pred_ids = np.argmax(pred_logits, axis=-1)\n\n    labels_copy = labels.copy()\n    labels_copy[labels_copy == -100] = processor.tokenizer.pad_token_id\n\n    pred_str = processor.batch_decode(pred_ids)\n    label_str = processor.batch_decode(labels_copy, group_tokens=False)\n\n    wer = metric.compute(predictions=pred_str, references=label_str)\n    return wer","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:48:12.925941Z","iopub.execute_input":"2024-04-22T04:48:12.926645Z","iopub.status.idle":"2024-04-22T04:48:12.936371Z","shell.execute_reply.started":"2024-04-22T04:48:12.926610Z","shell.execute_reply":"2024-04-22T04:48:12.935249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_fn(fold, train_loader, model, processor, optimizer, epoch, device):\n    model.train()\n    scaler = torch.cuda.amp.GradScaler(enabled=config.apex)\n    losses = AverageMeter()\n    start = end = time.time()\n    for step, data in enumerate(train_loader):\n        data = {k: v.to(device) for k, v in data.items()}\n        batch_size = len(data)\n        with torch.cuda.amp.autocast(enabled=config.apex):\n            pred = model(**data)\n            loss = pred.loss\n        losses.update(loss.item(), batch_size)\n        scaler.scale(loss).backward()\n        # torch.nn.utils.clip_grad_norm_(model.parameters(), 1)\n        scaler.step(optimizer)\n        scaler.update()\n        \n        optimizer.zero_grad()\n        end = time.time()\n        if step % config.print_freq == 0 or step == (len(train_loader)-1):\n            print(\"\\n\\n\\n\\n\")\n            LOGGER.info('Epoch: [{0}][{1}/{2}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  'LR: {lr:.8f} '\n                  .format(epoch+1, step, len(train_loader),\n                          remain=timeSince(start, float(step+1)/len(train_loader)),\n                          loss=losses,\n                          lr=optimizer.param_groups[0]['lr'],))\n            print(\"\\n\\n\\n\\n\")\n    return losses.avg, None","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:48:14.436136Z","iopub.execute_input":"2024-04-22T04:48:14.437001Z","iopub.status.idle":"2024-04-22T04:48:14.447291Z","shell.execute_reply.started":"2024-04-22T04:48:14.436968Z","shell.execute_reply":"2024-04-22T04:48:14.446360Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def valid_fn(valid_loader, model, processor, metric, device):\n    losses = AverageMeter()\n    wers = AverageMeter()\n    model.eval()\n    label_list = []\n    pred_list = []\n    start = end = time.time()\n    for step, data in enumerate(valid_loader):\n        data = {k: v.to(device) for k, v in data.items()}\n        batch_size = len(data)\n        with torch.no_grad():\n            with torch.cuda.amp.autocast(enabled=config.apex):\n                pred = model(**data)\n                loss = pred.loss\n                wer = compute_wer(pred, data['labels'].cpu().detach().numpy(), processor, metric)\n        losses.update(loss.item(), batch_size)\n        wers.update(wer, batch_size)\n        end = time.time()\n        if step % config.print_freq == 0 or step == (len(valid_loader)-1):\n            print(\"\\n\\n\\n\\n\")\n            LOGGER.info('EVAL: [{0}/{1}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  'WER: {wers.val:.4f}({wers.avg:.4f}) '\n                  .format(step, len(valid_loader),\n                          loss=losses,\n                          wers=wers,\n                          remain=timeSince(start, float(step+1)/len(valid_loader))))\n            print(\"\\n\\n\\n\\n\")\n\n    return losses.avg, wers.avg","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:48:15.942997Z","iopub.execute_input":"2024-04-22T04:48:15.943401Z","iopub.status.idle":"2024-04-22T04:48:15.953919Z","shell.execute_reply.started":"2024-04-22T04:48:15.943363Z","shell.execute_reply":"2024-04-22T04:48:15.952873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DATASET","metadata":{}},{"cell_type":"code","source":"import pandas as pd\npath = \"/kaggle/input/train-v3-wav2vec-norm-20-april/train_all_v3_norm_wav2vec2v1.csv\"\ndf = pd.read_csv(path)\ndf.head(2)","metadata":{"execution":{"iopub.status.busy":"2024-04-24T11:53:22.021083Z","iopub.execute_input":"2024-04-24T11:53:22.021530Z","iopub.status.idle":"2024-04-24T11:53:24.204803Z","shell.execute_reply.started":"2024-04-24T11:53:22.021495Z","shell.execute_reply":"2024-04-24T11:53:24.203794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bucket = []\nfor idx, row in df.iterrows():\n    bucket.append((row['wer_wav2vec2_v1'], row['audio_length'], row['path'], row['sentence'], row['prediction_sentence_wav2vec2_v1_norm']))\n\ngood = []\nbad = []\nfor t in bucket:\n    if t[1] >= 15.0 and t[0] <= 0.75:\n        good.append(t)\n    else:\n        bad.append(t)\nprint(len(good), len(bad), len(good)/len(bucket), len(bad)/len(bucket))","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:48:40.740795Z","iopub.execute_input":"2024-04-22T04:48:40.741667Z","iopub.status.idle":"2024-04-22T04:48:41.699640Z","shell.execute_reply.started":"2024-04-22T04:48:40.741626Z","shell.execute_reply":"2024-04-22T04:48:41.698681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\nrandom.seed(57)\n\nrandom.shuffle(good)\ngood[:2]","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:48:52.640588Z","iopub.execute_input":"2024-04-22T04:48:52.640920Z","iopub.status.idle":"2024-04-22T04:48:52.653995Z","shell.execute_reply.started":"2024-04-22T04:48:52.640896Z","shell.execute_reply":"2024-04-22T04:48:52.653112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dataset_creation(bucket):\n    path = []\n    sentence = []\n    audio_length = []\n    prediction_whisper = []\n    wer_whisper = []\n#     folder = \"/kaggle/input/ben10/ben10/16_kHz_train_audio/\"\n    folder = \"\"\n    for t in bucket:\n        path.append(folder + t[2])\n        sentence.append(t[3])\n        audio_length.append(t[1])\n        prediction_whisper.append(t[4])\n        \n        wer_whisper.append(t[0])\n    data = {\n        'path': path,\n        'sentence': sentence,\n        'audio_length': audio_length,\n        'prediction_whisper': prediction_whisper,\n        'wer_whisper': wer_whisper\n    }\n    df = pd.DataFrame(data, columns=data.keys())\n    return df\n\ndef split_creation(tuple_list):\n    sz = len(tuple_list)\n    train_size = int(sz*0.9)\n    train = dataset_creation(tuple_list[:train_size])\n    test_size = (sz-train_size)\n    val = dataset_creation(tuple_list[train_size:])\n#     test = dataset_creation(tuple_list[train_size:train_size+test_size])\n    test = val #dataset_creation(tuple_list[train_size+test_size:])\n    return train, test, val\n\ntrain, test, val = split_creation(tuple_list=good)\nprint(\"========================================\")\nprint(\"Train:\", len(train), \"Test:\", len(test), \"Val:\", len(val))\nprint(\"========================================\")","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:48:54.339505Z","iopub.execute_input":"2024-04-22T04:48:54.340192Z","iopub.status.idle":"2024-04-22T04:48:54.365835Z","shell.execute_reply.started":"2024-04-22T04:48:54.340150Z","shell.execute_reply":"2024-04-22T04:48:54.364768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.head(2)","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:49:01.102189Z","iopub.execute_input":"2024-04-22T04:49:01.102874Z","iopub.status.idle":"2024-04-22T04:49:01.114435Z","shell.execute_reply.started":"2024-04-22T04:49:01.102843Z","shell.execute_reply":"2024-04-22T04:49:01.113391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = train.sort_values('audio_length')\ntest = test.sort_values('audio_length')\nval = val.sort_values('audio_length')","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:49:10.023567Z","iopub.execute_input":"2024-04-22T04:49:10.024199Z","iopub.status.idle":"2024-04-22T04:49:10.035535Z","shell.execute_reply.started":"2024-04-22T04:49:10.024146Z","shell.execute_reply":"2024-04-22T04:49:10.034815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:49:12.035530Z","iopub.execute_input":"2024-04-22T04:49:12.036154Z","iopub.status.idle":"2024-04-22T04:49:12.047523Z","shell.execute_reply.started":"2024-04-22T04:49:12.036120Z","shell.execute_reply":"2024-04-22T04:49:12.046601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.to_csv('train_v3_.csv')\ntest.to_csv('test_v2_487.csv')\nval.to_csv('val_v2_487.csv')","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:49:13.252850Z","iopub.execute_input":"2024-04-22T04:49:13.253500Z","iopub.status.idle":"2024-04-22T04:49:13.486387Z","shell.execute_reply.started":"2024-04-22T04:49:13.253459Z","shell.execute_reply":"2024-04-22T04:49:13.485520Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"clean_memory()","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:50:08.617835Z","iopub.execute_input":"2024-04-22T04:50:08.618517Z","iopub.status.idle":"2024-04-22T04:50:08.779828Z","shell.execute_reply.started":"2024-04-22T04:50:08.618480Z","shell.execute_reply":"2024-04-22T04:50:08.778894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_loop(fold):\n    \n    LOGGER.info(f'================== fold: {fold} training ======================')\n    \n#     train_df = pd.read_csv('../data/preprocessed_train_with_audio_length.csv')\n#     train_df = train_df[train_df['audio_length'] < 15].reset_index(drop=True)\n    \n#     train_df['path'] = train_df['id'].apply(lambda x: os.path.join(config.audio_dir, x+'.mp3'))\n#     train_folds = train_df[train_df['fold'] != fold].reset_index(drop=True)\n#     train_folds = train_folds[train_folds.split == 'valid'].reset_index(drop=True)\n\n#     valid_folds = train_df[train_df['fold'] == fold].reset_index(drop=True)\n    \n#     train_folds = train_folds.sort_values('audio_length').reset_index(drop=True)\n#     valid_folds = valid_folds.sort_values('audio_length').reset_index(drop=True)\n\n    train_folds = pd.read_csv('/kaggle/working/train_v3_.csv')   #.head(100)\n    valid_folds = pd.read_csv('/kaggle/working/val_v2_487.csv')    #.head(10)\n    \n    processor = Wav2Vec2Processor.from_pretrained(config.model)\n    model = Wav2Vec2ForCTC.from_pretrained(config.model,\n        ctc_loss_reduction=\"mean\",\n        ignore_mismatched_sizes=True,\n        pad_token_id=processor.tokenizer.pad_token_id,\n        vocab_size=len(processor.tokenizer)\n    )\n    \n    model.config.ctc_zero_infinity = True\n    model.to(device)\n    model.freeze_feature_encoder()\n\n\n    optimizer = AdamW(model.parameters(), lr=config.lr, weight_decay=config.weight_decay)\n    num_cycles = 0.5\n    scheduler = get_cosine_schedule_with_warmup(\n                optimizer, num_warmup_steps=0, num_training_steps=config.epochs, num_cycles=num_cycles\n            )\n    \n    wer_metric = load_metric(\"wer\")\n\n    valid_dataset = ASRDataset(valid_folds, config, processor)\n    partial_func = partial(ctc_data_collator, processor=processor)\n    valid_loader = DataLoader(valid_dataset,\n                             batch_size=config.valid_bs,\n                             shuffle=False,\n                             collate_fn=partial_func,\n                             num_workers=config.num_workers, pin_memory=True, drop_last=False)\n\n    best_score = float('inf')\n\n    for epoch in range(config.epochs):\n        start_time = time.time()\n\n        train_folds = shuffle_df_chunk_by_chunk(train_folds, chunk_size=config.train_bs)\n        train_dataset = ASRDataset(train_folds, config, processor)\n        train_loader = DataLoader(train_dataset,\n                                batch_size=config.train_bs,\n                                shuffle=False,\n                                collate_fn=partial_func,\n                                num_workers=config.num_workers, pin_memory=True, drop_last=True)\n        # train\n        avg_loss, avg_wer = train_fn(fold, train_loader, model, processor, optimizer, epoch, device)\n        scheduler.step()\n        # eval\n        val_loss, valid_wer = valid_fn(valid_loader, model, processor, wer_metric, device)\n\n        elapsed = time.time() - start_time\n\n        if best_score > valid_wer:\n            best_score = valid_wer\n            LOGGER.info(f'Epoch {epoch+1} - Save Best WER: {valid_wer:.4f} Model')\n            torch.save({'model': model.state_dict(),},\n                        OUT_PATH + f\"model_stage{config.stage}.pth\")\n\n    LOGGER.info(f'[Fold{fold}] Best WER: {best_score}')\n    torch.cuda.empty_cache()\n    gc.collect()\n\nfold = 0\nif __name__ == '__main__':\n    train_loop(fold)","metadata":{"execution":{"iopub.status.busy":"2024-04-22T04:50:13.408866Z","iopub.execute_input":"2024-04-22T04:50:13.409566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"DONE\")\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training End","metadata":{}}]}