{"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":"code","source":"import os\nimport pandas as pd\nimport librosa\nimport numpy as np\nimport nltk\nimport IPython.display as ipd\nimport torchaudio\nimport librosa.display\nimport matplotlib.pyplot as plt\nimport random\nimport string\nimport torch\nimport torch.nn as nn\nfrom tqdm import tqdm\nfrom torch.utils.data import DataLoader\n\nfrom nltk.tokenize import word_tokenize\nfrom nltk.corpus import stopwords\nfrom nltk.stem import PorterStemmer\nfrom collections import Counter\nfrom sklearn.feature_extraction.text import CountVectorizer\nfrom sklearn.decomposition import LatentDirichletAllocation\nfrom gensim import models","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-08-26T20:05:04.068378Z","iopub.execute_input":"2023-08-26T20:05:04.068946Z","iopub.status.idle":"2023-08-26T20:05:04.080486Z","shell.execute_reply.started":"2023-08-26T20:05:04.068907Z","shell.execute_reply":"2023-08-26T20:05:04.078443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import (\n    Wav2Vec2ForCTC,\n    Wav2Vec2Processor,\n    Wav2Vec2CTCTokenizer,\n    Wav2Vec2FeatureExtractor\n) ","metadata":{"execution":{"iopub.status.busy":"2023-08-26T19:56:11.239353Z","iopub.execute_input":"2023-08-26T19:56:11.240480Z","iopub.status.idle":"2023-08-26T19:56:11.252860Z","shell.execute_reply.started":"2023-08-26T19:56:11.240441Z","shell.execute_reply":"2023-08-26T19:56:11.251760Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# WE ARE LOADING THE DATA NOW AND PRINTING A HEADER","metadata":{}},{"cell_type":"code","source":"# Define the paths to the data directories\nbasedir = '/kaggle/input/bengaliai-speech'\ntraindata = f\"{basedir}/train_mp3s/\"  \ntestdata = f\"{basedir}/test_mp3s/\" \ntraincsv = f\"{basedir}/train.csv\" \ndomains = f\"{basedir}/examples/\" \n\ntrain_valid_df = pd.read_csv(traincsv)\ntrain_valid_df['path'] = train_valid_df['id'].apply(lambda x: os.path.join(Config['audio_dir'], x+'.mp3'))\ntraindf = train_valid_df[train_valid_df['split'] == 'train'].sample(frac=.005).reset_index(drop=True)\nvaliddf = train_valid_df[train_valid_df['split'] == 'valid'].sample(frac=.005).reset_index(drop=True)\n# Load the train.csv file through pandas\n\n","metadata":{"execution":{"iopub.status.busy":"2023-08-26T19:56:11.254475Z","iopub.execute_input":"2023-08-26T19:56:11.255461Z","iopub.status.idle":"2023-08-26T19:56:18.602660Z","shell.execute_reply.started":"2023-08-26T19:56:11.255426Z","shell.execute_reply":"2023-08-26T19:56:18.601111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Previewing the first 10 rows of the test DataFrame\ndisplay(traindf.head(10))","metadata":{"execution":{"iopub.status.busy":"2023-08-26T19:56:18.605901Z","iopub.execute_input":"2023-08-26T19:56:18.606310Z","iopub.status.idle":"2023-08-26T19:56:18.620458Z","shell.execute_reply.started":"2023-08-26T19:56:18.606275Z","shell.execute_reply":"2023-08-26T19:56:18.619050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Previewing the first 10 rows of the validation DataFrame\ndisplay(validdf.head(10))","metadata":{"execution":{"iopub.status.busy":"2023-08-26T19:56:18.622629Z","iopub.execute_input":"2023-08-26T19:56:18.623170Z","iopub.status.idle":"2023-08-26T19:56:18.641219Z","shell.execute_reply.started":"2023-08-26T19:56:18.623124Z","shell.execute_reply":"2023-08-26T19:56:18.639940Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Config = {\n    'audio_dir': '/kaggle/input/bengaliai-speech/train_mp3s',\n    'lr': 3e-4,\n    'wd': 1e-5,\n    'T_0': 10,\n    'T_mult': 2,\n    'eta_min': 1e-6,\n    'nb_epochs': 5,\n    'train_bs': 16,\n    'valid_bs': 16,\n    'sampling_rate': 16000,\n}","metadata":{"execution":{"iopub.status.busy":"2023-08-26T19:56:18.642724Z","iopub.execute_input":"2023-08-26T19:56:18.643154Z","iopub.status.idle":"2023-08-26T19:56:18.651677Z","shell.execute_reply.started":"2023-08-26T19:56:18.643118Z","shell.execute_reply":"2023-08-26T19:56:18.650363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# domains for test set 17 named domains\nDOMAINS = !ls /kaggle/input/bengaliai-speech/examples/\nDOMAINS","metadata":{"execution":{"iopub.status.busy":"2023-08-26T19:56:18.653536Z","iopub.execute_input":"2023-08-26T19:56:18.654348Z","iopub.status.idle":"2023-08-26T19:56:18.677231Z","shell.execute_reply.started":"2023-08-26T19:56:18.654301Z","shell.execute_reply":"2023-08-26T19:56:18.675623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# Load audio files and corresponding transcriptions\naudiodata = []  # List to store audio data\ntranscriptions = []  # List to store corresponding transcriptions\n\nfor idx, row in traindf.head().iterrows() :\n    audiofilepath = os.path.join(traindata, f\"{row['id']}.mp3\")\nfor idx, row in validdf.head().iterrows() :\n    audiofilepath = os.path.join(traindata, f\"{row['id']}.mp3\")\n    \n    # Load the audio file using librosa\n    audio, sr = librosa.load(audiofilepath, sr=Config['sampling_rate'])\n\n    # Append audio data and transcription to lists\n    audiodata.append(audio)\n    transcriptions.append(row['sentence'])\n    \naudiodata = np.array(audiodata,dtype = 'object')\ntranscriptions = np.array(transcriptions,dtype = 'object')","metadata":{"execution":{"iopub.status.busy":"2023-08-26T20:09:04.880816Z","iopub.execute_input":"2023-08-26T20:09:04.881290Z","iopub.status.idle":"2023-08-26T20:09:04.936770Z","shell.execute_reply.started":"2023-08-26T20:09:04.881251Z","shell.execute_reply":"2023-08-26T20:09:04.935412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Audio data shape:\", audiodata.shape)\nprint(\"Transcriptions shape:\", transcriptions.shape)\n","metadata":{"execution":{"iopub.status.busy":"2023-08-26T20:09:08.605571Z","iopub.execute_input":"2023-08-26T20:09:08.606051Z","iopub.status.idle":"2023-08-26T20:09:08.613377Z","shell.execute_reply.started":"2023-08-26T20:09:08.606016Z","shell.execute_reply":"2023-08-26T20:09:08.611792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get the total number of audio files in the training and test directories\ntrainaudiofiles = os.listdir(traindata)\ntestaudiofiles = os.listdir(testdata)\n\n# Get the total duration of audio data in the training set (in seconds)\ntraintotalduration = 0\nSAMPLESTAKEN = Config['sampling_rate']\nfor idx, row in traindf.head(SAMPLESTAKEN).iterrows():\n    audioinfo = torchaudio.info(audiofilepath)\n    duration = audioinfo.num_frames / audioinfo.sample_rate\n    traintotalduration += duration\nfor idx, row in validdf.head(SAMPLESTAKEN).iterrows():\n    audioinfo = torchaudio.info(audiofilepath)\n    duration = audioinfo.num_frames / audioinfo.sample_rate\n    traintotalduration += duration\ntestduration = 0\nfor testfile in testaudiofiles:\n    audiofilepath = os.path.join(testdata, str(testfile))\n    audioinfo = torchaudio.info(audiofilepath)\n    duration = audioinfo.num_frames / audioinfo.sample_rate\n    testduration += duration\n    \n# Get the total number of samples in the training data\ntotalsamples = traindf.shape[0]\n\n# Print the data summary\nprint(\"Data Summary:\")\nprint(f\"Total number of audio files in the training directory: {len(trainaudiofiles)}\")\nprint(f\"Total number of audio files in the test directory: {len(testaudiofiles)}\")\nprint(f\"Total duration of audio data in the training set (in seconds): {traintotalduration:.2f}\")\nprint(f\"Average duration of audio file in training set (in seconds): {traintotalduration/SAMPLESTAKEN:.2f}\")\nprint(f\"Total duration of audio data in the test set (in seconds): {testduration:.2f}\")\nprint(f\"Total number of samples in the training data: {totalsamples}\")\n","metadata":{"execution":{"iopub.status.busy":"2023-08-26T20:18:09.879340Z","iopub.execute_input":"2023-08-26T20:18:09.879853Z","iopub.status.idle":"2023-08-26T20:18:50.482241Z","shell.execute_reply.started":"2023-08-26T20:18:09.879814Z","shell.execute_reply":"2023-08-26T20:18:50.480922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install mutagen --quiet\nimport mutagen\nfrom mutagen.mp3 import MP3\n","metadata":{"execution":{"iopub.status.busy":"2023-08-26T19:57:33.108516Z","iopub.execute_input":"2023-08-26T19:57:33.108890Z","iopub.status.idle":"2023-08-26T19:57:47.822428Z","shell.execute_reply.started":"2023-08-26T19:57:33.108856Z","shell.execute_reply":"2023-08-26T19:57:47.821080Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check some random indices \nrandom_indices = [0, 10, 20, 30, 40]\n\nfor ind in random_indices:\n    audiofilepath = os.path.join(traindata, f\"{row['id']}.mp3\")\n    row = traindf.iloc[ind]\n    audio = MP3(audiofilepath)\n    \n    print(audio.info.length)\n    \n\n    # Load the audio file using librosa\n    audio, sr = librosa.load(audiofilepath, sr=Config['sampling_rate'])\n    # Print the transcription and play the audio\n    ipd.display(ipd.Audio(audio, rate=sr))\n    print(\"Transcription:\", row['sentence'])\n    ","metadata":{"execution":{"iopub.status.busy":"2023-08-26T20:34:27.018640Z","iopub.execute_input":"2023-08-26T20:34:27.019154Z","iopub.status.idle":"2023-08-26T20:34:27.114355Z","shell.execute_reply.started":"2023-08-26T20:34:27.019117Z","shell.execute_reply":"2023-08-26T20:34:27.112997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for idx in random_indices:\n    row = traindf.iloc[idx]\n    audiofilepath = os.path.join(traindata, f\"{row['id']}.mp3\")\n\n    # Load the audio file using librosa\n    audio, sr = librosa.load(audiofilepath, sr=Config['sampling_rate'])\n\n    # Plot the waveform\n    plt.figure(figsize=(10, 4))\n    librosa.display.waveshow(audio, sr=sr)\n    plt.title(f\"Waveform - Audio File ID: {row['id']}\")\n    plt.xlabel(\"Time (s)\")\n    plt.ylabel(\"Amplitude\")\n    plt.tight_layout()\n    plt.show()\n\n    # Plot the log Mel spectrogram\n    plt.figure(figsize=(10, 4))\n    melspec = librosa.feature.melspectrogram(y=audio, sr=sr)\n    melspecdb = librosa.power_to_db(melspec, ref=np.max)\n    librosa.display.specshow(melspecdb, sr=sr, x_axis='time', y_axis='mel')\n    plt.colorbar(format='%+2.0f dB')\n    plt.title(f\"Log Mel Spectrogram - Audio File ID: {row['id']}\")\n    plt.xlabel(\"Time (s)\")\n    plt.ylabel(\"Mel Frequency\")\n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-08-26T20:43:39.685840Z","iopub.execute_input":"2023-08-26T20:43:39.686468Z","iopub.status.idle":"2023-08-26T20:43:45.929282Z","shell.execute_reply.started":"2023-08-26T20:43:39.686405Z","shell.execute_reply":"2023-08-26T20:43:45.927861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for idx, row in traindf.head(SAMPLESTAKEN).iterrows():\n    audiofilepath = os.path.join(traindata, f\"{row['id']}.mp3\")\n    # Load the audio file using librosa\n    audio, sr = librosa.load(audiofilepath, sr=Config['sampling_rate'])\n    # Compute the MFCCs\n    mfccs = librosa.feature.mfcc(y=audio, sr=sr, n_mfcc=13)\n\n    # Optionally, convert MFCCs to dB scale\n    mfccs_db_train = librosa.power_to_db(mfccs, ref=np.max)\nfor idx, row in validdf.head(SAMPLESTAKEN).iterrows():\n    row = validdf.iloc[idx]\n    audiofilepath = os.path.join(traindata, f\"{row['id']}.mp3\")\n\n    # Load the audio file using librosa\n    audio, sr = librosa.load(audiofilepath, sr=Config['sampling_rate'])\n    # Compute the MFCCs\n    mfccs = librosa.feature.mfcc(y=audio, sr=sr, n_mfcc=13)\n\n    # Optionally, convert MFCCs to dB scale\n    mfccs_db_valid = librosa.power_to_db(mfccs, ref=np.max)\n    ","metadata":{"execution":{"iopub.status.busy":"2023-08-26T20:47:37.045369Z","iopub.execute_input":"2023-08-26T20:47:37.045854Z","iopub.status.idle":"2023-08-26T20:50:36.696292Z","shell.execute_reply.started":"2023-08-26T20:47:37.045819Z","shell.execute_reply":"2023-08-26T20:50:36.694214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mfccs_db_train","metadata":{"execution":{"iopub.status.busy":"2023-08-26T20:51:06.504850Z","iopub.execute_input":"2023-08-26T20:51:06.505259Z","iopub.status.idle":"2023-08-26T20:51:06.515164Z","shell.execute_reply.started":"2023-08-26T20:51:06.505218Z","shell.execute_reply":"2023-08-26T20:51:06.513534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mfccs_db_valid","metadata":{"execution":{"iopub.status.busy":"2023-08-26T20:51:11.354485Z","iopub.execute_input":"2023-08-26T20:51:11.355621Z","iopub.status.idle":"2023-08-26T20:51:11.363871Z","shell.execute_reply.started":"2023-08-26T20:51:11.355577Z","shell.execute_reply":"2023-08-26T20:51:11.362598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Select the first 5 transcriptions\ntranscriptions = traindf['sentence'][:5].tolist()\n\n# Convert transcriptions to lowercase\ntranscriptionslower = [transcription.lower() for transcription in transcriptions]\n\n# Remove punctuation\ntranslator = str.maketrans(\"\", \"\", string.punctuation)\ntranscriptionsnopunct = [transcription.translate(translator) for transcription in transcriptionslower]\n\n# Tokenization\nnltk.download('punkt')  \ntranscriptionstokens = [word_tokenize(transcription) for transcription in transcriptionsnopunct]\n\n\nnltk.download('stopwords')  \nstopwords = set(stopwords.words('bengali'))\ntranscriptionsnostopwords = [\n    [word for word in tokens if word not in stopwords]\n    for tokens in transcriptionstokens\n]\n\nnltk.download('wordnet')  \nstemmer = PorterStemmer()\ntranscriptionsstemmed = [\n    [stemmer.stem(word) for word in tokens]\n    for tokens in transcriptionsnostopwords\n]\n\n# Print the preprocessed transcriptions\nfor i, transcription in enumerate(transcriptionsstemmed):\n    print(f\"Preprocessed transcription {i+1}: {' '.join(transcription)}\")\n","metadata":{"execution":{"iopub.status.busy":"2023-08-26T19:57:56.126166Z","iopub.execute_input":"2023-08-26T19:57:56.126540Z","iopub.status.idle":"2023-08-26T19:57:56.411934Z","shell.execute_reply.started":"2023-08-26T19:57:56.126508Z","shell.execute_reply":"2023-08-26T19:57:56.410514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Build vocabulary\nvocabulary = set()\nfor transcriptiontokens in transcriptionstokens:\n    vocabulary.update(transcriptiontokens)\n\n# print(\"Vocabulary:\")\n# print(vocabulary)\nprint(f\"Vocabulary Size: {len(vocabulary)}\")","metadata":{"execution":{"iopub.status.busy":"2023-08-26T19:57:56.413582Z","iopub.execute_input":"2023-08-26T19:57:56.413947Z","iopub.status.idle":"2023-08-26T19:57:56.421045Z","shell.execute_reply.started":"2023-08-26T19:57:56.413916Z","shell.execute_reply":"2023-08-26T19:57:56.419743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Compute descriptive statistics\nsentencelengths = [len(tokens) for tokens in transcriptionstokens]\nminlength = min(sentencelengths)\nmaxlength = max(sentencelengths)\nmeanlength = sum(sentencelengths) / len(sentencelengths)\nmedianlength = sorted(sentencelengths)[len(sentencelengths) // 2]\n\n# Plot the distribution of sentence lengths\nplt.figure(figsize=(10, 6))\nplt.hist(sentencelengths, bins=50, color='skyblue', edgecolor='black')\nplt.axvline(meanlength, color='red', linestyle='dashed', linewidth=2, label='Mean')\nplt.axvline(medianlength, color='green', linestyle='dashed', linewidth=2, label='Median')\nplt.xlabel('Sentence Length')\nplt.ylabel('Frequency')\nplt.title('Distribution of Sentence Lengths')\nplt.legend()\nplt.show()\n\n# Print the descriptive statistics\nprint(\"Descriptive Statistics:\")\nprint(f\"Minimum Sentence Length: {minlength}\")\nprint(f\"Maximum Sentence Length: {maxlength}\")\nprint(f\"Mean Sentence Length: {meanlength:.2f}\")\nprint(f\"Median Sentence Length: {medianlength}\")","metadata":{"execution":{"iopub.status.busy":"2023-08-26T19:57:56.422698Z","iopub.execute_input":"2023-08-26T19:57:56.423243Z","iopub.status.idle":"2023-08-26T19:57:56.921627Z","shell.execute_reply.started":"2023-08-26T19:57:56.423210Z","shell.execute_reply":"2023-08-26T19:57:56.920773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Select the first 1000 sentences \nsentences = traindf['sentence'][:10000].tolist()\n\n# Tokenization using NLTK\nnltk.download('punkt')  # Download the Punkt tokenizer\nsentences_tokens = [word_tokenize(sentence) for sentence in sentences]\n\n\n# Remove Bengali stopwords\nsentencesnostopwords = [\n    [word for word in tokens if word not in stopwords]\n    for tokens in sentences_tokens\n]\n\n# Convert tokenized sentences back to strings\nsentencesprocessed = [' '.join(tokens) for tokens in sentencesnostopwords]\n\n# Create a CountVectorizer to convert text data to a bag-of-words representation\nvectorizer = CountVectorizer(max_features=1000)\nX = vectorizer.fit_transform(sentencesprocessed)\n\n# Perform LDA topic modeling\nn_topics = 10  # Number of topics to discover\nlda_model = LatentDirichletAllocation(n_components=n_topics, random_state=42)\nlda_model.fit(X)\n\n# Get the top words for each topic\nfeature_names = vectorizer.get_feature_names_out()\ntop_words_per_topic = []\nfor topic_idx, topic in enumerate(lda_model.components_):\n    top_words = [feature_names[i] for i in topic.argsort()[:-4:-1]]\n    top_words_per_topic.append(top_words)\n\n# Print the top words for each topic\nfor i, top_words in enumerate(top_words_per_topic):\n    print(f\"Topic {i + 1}: {' '.join(top_words)}\")\n    ","metadata":{"execution":{"iopub.status.busy":"2023-08-26T19:57:56.923054Z","iopub.execute_input":"2023-08-26T19:57:56.923949Z","iopub.status.idle":"2023-08-26T19:58:06.427980Z","shell.execute_reply.started":"2023-08-26T19:57:56.923911Z","shell.execute_reply":"2023-08-26T19:58:06.426522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass JUAudioModel(nn.Module):\n    def __init__(self, input_size, hidden_size, num_classes):\n        super(JUAudioModel, self).__init__()\n        \n        self.lstm1 = nn.LSTM(input_size, hidden_size, batch_first=True)\n        self.lstm2 = nn.LSTM(hidden_size, hidden_size, batch_first=True)\n        self.lstm3 = nn.LSTM(hidden_size, hidden_size, batch_first=True)\n        \n        self.flatten = nn.Flatten()\n        \n        self.fc1 = nn.Linear(hidden_size, 256)\n        self.fc2 = nn.Linear(256, 128)\n        self.fc3 = nn.Linear(128, num_classes)\n        \n        self.dropout = nn.Dropout(0.3)\n        \n    def forward(self, x):\n        x, _ = self.lstm1(x)\n        x = self.dropout(x)\n        x, _ = self.lstm2(x)\n        x = self.dropout(x)\n        x, _ = self.lstm3(x)\n        x = self.dropout(x)\n        \n        x = self.flatten(x)\n        \n        x = self.fc1(x)\n        x = self.dropout(x)\n        x = self.fc2(x)\n        x = self.dropout(x)\n        x = self.fc3(x)\n        \n        return x\ninput_size = 1 \nhidden_size = 128 \nnum_classes = len(vocabulary) \n\nmodel = JUAudioModel(input_size, hidden_size, num_classes)","metadata":{"execution":{"iopub.status.busy":"2023-08-26T20:01:14.895475Z","iopub.execute_input":"2023-08-26T20:01:14.895914Z","iopub.status.idle":"2023-08-26T20:01:14.916631Z","shell.execute_reply.started":"2023-08-26T20:01:14.895880Z","shell.execute_reply":"2023-08-26T20:01:14.915098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_one_epoch(model, train_loader, optimizer, device='cuda:0'):\n    model.train()\n    pbar = tqdm(train_loader, total=len(train_loader))\n    avg_loss = 0\n    for data in pbar:\n        data = {k: v.to(device) for k, v in data.items()}\n        loss = model(**data).loss\n        loss_itm = loss.item()\n        \n        avg_loss += loss_itm\n        pbar.set_description(f\"loss: {loss_itm:.4f}\")\n        \n        optimizer.zero_grad(set_to_none=True)\n        loss.backward()\n        optimizer.step()\n        \n    return avg_loss / len(train_loader)\n\n@torch.no_grad()\ndef valid_one_epoch(model, valid_loader, device='cuda:0'):\n    pbar = tqdm(valid_loader, total=len(valid_loader))\n    avg_loss = 0\n    for data in pbar:\n        data = {k: v.to(device) for k, v in data.items()}\n        loss = model(**data).loss\n        loss_itm = loss.item()\n        \n        avg_loss += loss_itm\n        pbar.set_description(f\"val_loss: {loss_itm:.4f}\")\n\n    return avg_loss / len(valid_loader)","metadata":{"execution":{"iopub.status.busy":"2023-08-26T19:58:06.547819Z","iopub.status.idle":"2023-08-26T19:58:06.548258Z","shell.execute_reply.started":"2023-08-26T19:58:06.548038Z","shell.execute_reply":"2023-08-26T19:58:06.548058Z"},"trusted":true},"execution_count":null,"outputs":[]}]}