{"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":"# STARTER [Pytorch Version]","metadata":{}},{"cell_type":"markdown","source":"## SOURCE\n - [Automatic speech recognition](https://huggingface.co/docs/transformers/main/tasks/asr)\n - [int8 training for automatic speech recognition](https://huggingface.co/docs/peft/task_guides/int8-asr#int8-training-for-automatic-speech-recognition)\n - [Detailed EDA, Normalizer and WER](https://www.kaggle.com/code/mbmmurad/detailed-eda-normalizer-and-wer)\n - [Github_BanglaSpeech2Text](https://github.com/shhossain/BanglaSpeech2Text)\n - [WANDB_Fine-Tuning Whisper ASR Models](https://wandb.ai/parambharat/whisper_finetuning/reports/Fine-Tuning-Whisper-ASR-Models---VmlldzozMTEzNDE5)\n     - [Colab_whisper_tiny_ta.ipynb](https://colab.research.google.com/drive/1RkboArXsuXIEDTE5OHfJe-0Gn7v3gXI1?usp=sharing)\n - [HF_Fine-Tune Whisper For Multilingual ASR with 🤗 Transformers](https://huggingface.co/blog/fine-tune-whisper)\n     - [HF_COLAB_fine_tune_whisper.ipynb](https://colab.research.google.com/github/sanchit-gandhi/notebooks/blob/main/fine_tune_whisper.ipynb)\n - [BengaliAI-ASR: Baseline Whisper Inference](https://www.kaggle.com/code/emphymachine/bengaliai-asr-baseline-whisper-inference)\n     - [Whisper in Transformers.ipynb](https://colab.research.google.com/drive/16HO7if9iwfpSJzhqlaNOu6iiMhUBLMKE?usp=sharing)","metadata":{}},{"cell_type":"markdown","source":"# Install","metadata":{}},{"cell_type":"code","source":"%%capture\n%pip install wandb pyctcdecode colorama kenlm\n%pip install git+https://github.com/huggingface/peft.git\n%pip install -U git+https://github.com/huggingface/accelerate.git\n%pip install -q transformers datasets evaluate jiwer bitsandbytes","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:22:07.811415Z","iopub.execute_input":"2023-08-06T15:22:07.811776Z","iopub.status.idle":"2023-08-06T15:24:13.202829Z","shell.execute_reply.started":"2023-08-06T15:22:07.811745Z","shell.execute_reply":"2023-08-06T15:24:13.201505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### Another Install \n# %%capture\n# %pip install git+https://github.com/huggingface/peft.git\n# %pip install \"transformers==4.27.2\" \"datasets==2.9.0\" \"accelerate==0.17.1\" \"evaluate==0.4.0\" \"bitsandbytes==0.37.1\" loralib --upgrade --quiet\n# %pip install \"transformers==4.27.2\" \"datasets==2.9.0\" \"accelerate==0.17.1\" \"evaluate==0.4.0\"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import","metadata":{}},{"cell_type":"code","source":"import gc\nimport os\nimport time\nimport copy\nimport random\n\nimport numpy as np\nimport pandas as pd\n\nimport matplotlib.pyplot as plt\n# import seaborn as sns\n\n# Pytorch\nimport torch \nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import DataLoader, Dataset\n\n# Soundfile\n# import mutagen\n# from mutagen.mp3 import MP3\n\nimport librosa\nimport librosa.display\n\n# accelerate\nimport accelerate\nfrom accelerate import Accelerator\n# from accelerate.utils import set_seed \n\n# evaluate : WER\nimport evaluate\n# Utils\nfrom tqdm.auto import tqdm, trange\n\n# huggingface - transformers\nfrom transformers import WhisperForConditionalGeneration, WhisperFeatureExtractor, WhisperTokenizer, WhisperProcessor, AdamW\n# from transformers import TrainerCallback, TrainerState, TrainerControl, TrainingArguments, Trainer,  Wav2Vec2ProcessorWithLM\n# from transformers.trainer_utils import PREFIX_CHECKPOINT_DIR\n\n# HuggingFace peft \nfrom peft import LoraConfig, PeftModel, LoraModel, LoraConfig, get_peft_model, prepare_model_for_int8_training, TaskType\n\n\n# Suppress warnings\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# For colored terminal text\nfrom colorama import Fore, Back, Style\nb_ = Fore.BLUE\ny_ = Fore.YELLOW\nsr_ = Style.RESET_ALL\n\n# display pandas\nfrom IPython.display import display\n# import IPython.display as ipd\n\n# accelerate version\n# accelerate.__version__\n\n# os.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"\n# os.environ[\"CUDA_DEVICE_ORDER\"]=\"PCI_BUS_ID\"  # Arrange GPU devices starting from 0\n# os.environ[\"CUDA_VISIBLE_DEVICES\"]= \"0,1\"  # Set the GPUs 2 and 3 to use","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:13.206523Z","iopub.execute_input":"2023-08-06T15:24:13.206909Z","iopub.status.idle":"2023-08-06T15:24:28.529667Z","shell.execute_reply.started":"2023-08-06T15:24:13.206875Z","shell.execute_reply":"2023-08-06T15:24:28.528501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Wandb Login","metadata":{}},{"cell_type":"code","source":"import wandb\n\ntry:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    api_key = user_secrets.get_secret(\"wandb_api\")\n    wandb.login(key = api_key)\n    anony = None\nexcept:\n    anony = \"must\"\n    print(\"You need your user token\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:28.531338Z","iopub.execute_input":"2023-08-06T15:24:28.532327Z","iopub.status.idle":"2023-08-06T15:24:28.707692Z","shell.execute_reply.started":"2023-08-06T15:24:28.532285Z","shell.execute_reply":"2023-08-06T15:24:28.706673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"############### Model Lists 'Whispers' ###############\n##### Source: https://github.com/shhossain/BanglaSpeech2Text/blob/main/listed_models.json\n### 1) \"https://huggingface.co/Rakib/whisper-tiny-bn\" - \"name\": \"whisper-tiny-bn\",\n### 2) \"https://huggingface.co/shhossain/whisper-base-bn\" - \"name\": \"whisper-base-bn\",\n### 3) \"https://huggingface.co/anuragshas/whisper-large-v2-bn\" - \"name\": \"whisper-large-v2-bn\",\n### 4) \"https://huggingface.co/anuragshas/whisper-small-bn\" -  \"name\": \"whisper-small-bn\",","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = {\"seed\": 2023, \n          \"model_name\": \"shhossain/whisper-base-bn\", #\"openai/whisper-base\",\n          \"competition\": \"Bengali.AI Spech Recognition 2023\",\n          \"language\":\"bengali\",\n          \"task\": \"transcribe\",\n          \n          'n_epochs': 3,\n          'train_batch_size' : 8,\n          'valid_batch_size' : 32, \n          'max_length': 512, \n          \n           # 'grad_clipping' : True,\n          'learning_rate': 3e-4, # 5e-5\n          \"min_lr\": 1e-6,\n          \"T_max\": 500,\n          \"weight_decay\": 1e-6,\n          'n_accumulate': 1,\n          'max_grad_norm': 5,\n          \n           # LoRA\n          'is_lora': True,\n          'lora_r': 8, # 8\n          'lora_alpha': 16, # 32\n          'lora_target_modules': [\"q_proj\", \"v_proj\"],\n          'lora_dropout_p': 0.1,\n          # 'lora_task_type': TaskType.CAUSAL_LM,\n          \n          # Audio\n          \"sample_rate\": 16_000,\n          \"max_time\": 5, \n          \"n_mels\": 224, \n          \"n_fft\": 1024\n          }","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:28.710559Z","iopub.execute_input":"2023-08-06T15:24:28.710941Z","iopub.status.idle":"2023-08-06T15:24:28.720137Z","shell.execute_reply.started":"2023-08-06T15:24:28.710904Z","shell.execute_reply":"2023-08-06T15:24:28.719111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## SEED","metadata":{}},{"cell_type":"code","source":"def set_seed(seed=2022):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n\nset_seed(config['seed'])","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:28.721626Z","iopub.execute_input":"2023-08-06T15:24:28.722146Z","iopub.status.idle":"2023-08-06T15:24:28.732379Z","shell.execute_reply.started":"2023-08-06T15:24:28.722109Z","shell.execute_reply":"2023-08-06T15:24:28.731251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load Data","metadata":{}},{"cell_type":"code","source":"train_data_path = '/kaggle/input/bengali-eda-part2/train_normalized.csv'\n\ntrain = pd.read_csv(train_data_path)\ntrain.drop(['sentence'], axis = 1, inplace = True)\ntrain.rename(columns = {'normal_sentence' : 'sentence'}, inplace = True)\nprint(train.shape)\ndisplay(train.head())","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:28.733881Z","iopub.execute_input":"2023-08-06T15:24:28.734210Z","iopub.status.idle":"2023-08-06T15:24:38.071335Z","shell.execute_reply.started":"2023-08-06T15:24:28.734178Z","shell.execute_reply":"2023-08-06T15:24:38.070389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# No Progress BAR ?","metadata":{}},{"cell_type":"code","source":"from huggingface_hub.utils import are_progress_bars_disabled, disable_progress_bars, enable_progress_bars\n\n# Disable progress bars globally\ndisable_progress_bars()","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:38.072689Z","iopub.execute_input":"2023-08-06T15:24:38.073176Z","iopub.status.idle":"2023-08-06T15:24:38.080841Z","shell.execute_reply.started":"2023-08-06T15:24:38.073138Z","shell.execute_reply":"2023-08-06T15:24:38.079725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# processor","metadata":{}},{"cell_type":"code","source":"processor = WhisperProcessor.from_pretrained(\n     config['model_name'], \n     language=config['language'],\n     task=config['task'],\n     model_max_length= config['max_length']\n)\n## WhisperfeatureExtractor, WhisperTokenizer\nprint(\"Processor Loaded\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:38.084165Z","iopub.execute_input":"2023-08-06T15:24:38.084469Z","iopub.status.idle":"2023-08-06T15:24:40.066687Z","shell.execute_reply.started":"2023-08-06T15:24:38.084444Z","shell.execute_reply":"2023-08-06T15:24:40.065651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Processor SAVE\nprocessor.save_pretrained('./processor/')\nprint(\"Processor SAVED\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:40.068269Z","iopub.execute_input":"2023-08-06T15:24:40.068668Z","iopub.status.idle":"2023-08-06T15:24:40.203804Z","shell.execute_reply.started":"2023-08-06T15:24:40.068632Z","shell.execute_reply":"2023-08-06T15:24:40.202687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# processor","metadata":{"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MyDataset","metadata":{}},{"cell_type":"code","source":"# from torch.utils.data import Dataset, DataLoader\n\nclass MyDataset(Dataset):\n    def __init__(self, df):\n        self.df = df\n\n    def __len__(self):\n        return self.df.shape[0]\n\n    def __getitem__(self, idx):\n        return {'path': self.df.path.values[idx], 'labels': self.df.sentence.values[idx]}","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:40.208371Z","iopub.execute_input":"2023-08-06T15:24:40.208695Z","iopub.status.idle":"2023-08-06T15:24:40.214872Z","shell.execute_reply.started":"2023-08-06T15:24:40.208669Z","shell.execute_reply":"2023-08-06T15:24:40.213468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# idx = 11\n\n# # load orignial audio\n# audio, sr = librosa.load(train.path.values[idx], sr = None) \n# print(\"AUDIO SHAPE: \", audio.shape, sr)\n\n# # resample\n# audio = librosa.resample(audio, orig_sr=sr, target_sr = config['sample_rate']) \n# print(\"RESAMPLED AUDIO SHAPE: \", audio.shape)\n\n# # normalize\n# audio = librosa.util.normalize(audio) \n# print(\"NORMALIZED AUDIO SHAPE: \", audio.shape)\n\n# # processor\n# inputs = processor(audio = audio, sampling_rate = config['sample_rate'], return_tensors = 'pt')\n# print(\"INPUTS' input_features SHAPE: \", inputs.input_features.shape)\n\n#### AUDIO SHAPE:  (126720,) 32000\n#### RESAMPLED AUDIO SHAPE:  (63360,)\n#### NORMALIZED AUDIO SHAPE:  (63360,)\n#### INPUTS' input_features SHAPE:  torch.Size([1, 80, 3000])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Collate_fn","metadata":{}},{"cell_type":"code","source":"class MyDataCollatorCTCWithPadding:\n    def __init__(self, padding, processor):\n        self.processor = processor\n        self.padding = padding\n        \n    def __call__(self, features):\n        # split inputs and labels since they have to be of different lengths and need\n        # different padding methods\n        input_features = []\n        for feature in features:\n            # load orignial audio\n            audio, sr = librosa.load(feature['path'], sr = None) \n            # resample\n            audio = librosa.resample(audio, orig_sr=sr, target_sr = config['sample_rate']) \n            # normalize\n            audio = librosa.util.normalize(audio) \n            # processor\n            inputs = self.processor(audio = audio, \n                                    sampling_rate = config['sample_rate'], \n                                    return_tensors = 'pt')\n            input_features.append({'input_features': inputs.input_features.squeeze(0)})\n\n        batch = self.processor.feature_extractor.pad(input_features, padding=self.padding, return_tensors=\"pt\")\n        label_features = [{\"input_ids\": self.processor(text = feature[\"labels\"], return_tensors ='pt').input_ids.squeeze(0)} for feature in features]\n        labels_batch = self.processor.tokenizer.pad(label_features, padding=self.padding, return_tensors=\"pt\")\n\n        # replace padding with -100 to ignore loss correctly\n        labels = labels_batch[\"input_ids\"].masked_fill(labels_batch.attention_mask.ne(1), -100)\n\n        batch[\"labels\"] = labels\n\n        return batch\n\n\ncollate_fn = MyDataCollatorCTCWithPadding(padding = True, processor=processor)","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:40.216421Z","iopub.execute_input":"2023-08-06T15:24:40.217069Z","iopub.status.idle":"2023-08-06T15:24:40.228152Z","shell.execute_reply.started":"2023-08-06T15:24:40.217037Z","shell.execute_reply":"2023-08-06T15:24:40.227054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# prepare_loaders","metadata":{}},{"cell_type":"code","source":"# Split\ntrain_df = train.loc[train.split == 'train'].reset_index(drop = True)\nvalid_df = train.loc[train.split != 'train'].reset_index(drop = True)\n\nprint(train_df.shape, valid_df.shape)","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:40.229527Z","iopub.execute_input":"2023-08-06T15:24:40.229928Z","iopub.status.idle":"2023-08-06T15:24:40.696410Z","shell.execute_reply.started":"2023-08-06T15:24:40.229898Z","shell.execute_reply":"2023-08-06T15:24:40.694550Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Sample Test\ntrain_df = train_df[:10000].reset_index(drop = True)\nvalid_df = valid_df[:2000].reset_index(drop = True)\n\nprint(train_df.shape, valid_df.shape)","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:40.697705Z","iopub.execute_input":"2023-08-06T15:24:40.698071Z","iopub.status.idle":"2023-08-06T15:24:40.731963Z","shell.execute_reply.started":"2023-08-06T15:24:40.698037Z","shell.execute_reply":"2023-08-06T15:24:40.727184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### import os\nnums = os.cpu_count()\nif nums > 2:\n    nums -=1\nprint(nums)\n\n### MyDataset\ntrain_ds = MyDataset(df = train_df,)\nvalid_ds = MyDataset(df = valid_df,)\n\nprint(\"Dataset Completed\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:40.733605Z","iopub.execute_input":"2023-08-06T15:24:40.736805Z","iopub.status.idle":"2023-08-06T15:24:40.744111Z","shell.execute_reply.started":"2023-08-06T15:24:40.736771Z","shell.execute_reply":"2023-08-06T15:24:40.743201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### DataLoader\ntrain_loader = DataLoader(train_ds, \n                          collate_fn = collate_fn, \n                          batch_size = config['train_batch_size'], \n                          shuffle = True, \n                          pin_memory = True, \n                          num_workers = 1 ## In this notebook, num_workers = 1 or 2,\n                         )\nvalid_loader = DataLoader(valid_ds, \n                          collate_fn = collate_fn, \n                          batch_size = config['valid_batch_size'], \n                          shuffle = False, \n                          pin_memory = True, \n                          num_workers = 1 ## In this notebook, num_workers = 1 or 2,\n                         )\n\nprint(\"DataLoader Completed\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:40.745760Z","iopub.execute_input":"2023-08-06T15:24:40.746171Z","iopub.status.idle":"2023-08-06T15:24:40.760104Z","shell.execute_reply.started":"2023-08-06T15:24:40.746134Z","shell.execute_reply":"2023-08-06T15:24:40.758960Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# sample = next(iter(train_loader)) ### Batch Size: 4\n# sample.keys()\n#### dict_keys(['input_features', 'labels'])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# sample['input_features'].shape, sample[\"labels\"].shape, # sample[\"input_ids\"].shape, \n#### (torch.Size([4, 80, 3000]), torch.Size([4, 104]))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# WER","metadata":{}},{"cell_type":"code","source":"## WER: Evaluation Metric \nwer = evaluate.load(\"wer\")\nprint(\"WER Loaded\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:40.761332Z","iopub.execute_input":"2023-08-06T15:24:40.761847Z","iopub.status.idle":"2023-08-06T15:24:41.300878Z","shell.execute_reply.started":"2023-08-06T15:24:40.761815Z","shell.execute_reply":"2023-08-06T15:24:41.299863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"## Download Base Model \nprint(f\"BASE MODEL: {config['model_name']}\")\nbase_model = WhisperForConditionalGeneration.from_pretrained(config['model_name'], \n                #                                              load_in_8bit=True, ### Error\n                #                                              device_map=\"auto\"  ### Error\n                                                            )\n                                            \nprint(\"BASE MODEL COMPLETED\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:41.302326Z","iopub.execute_input":"2023-08-06T15:24:41.303015Z","iopub.status.idle":"2023-08-06T15:24:45.059680Z","shell.execute_reply.started":"2023-08-06T15:24:41.302981Z","shell.execute_reply":"2023-08-06T15:24:45.058727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### Override generation arguments - no tokens are forced as decoder outputs (see forced_decoder_ids), \n###  no tokens are suppressed during generation (see suppress_tokens):\nbase_model.config.forced_decoder_ids = None\nbase_model.config.suppress_tokens = []\nbase_model.config.use_cache = False\nprint(\"BASEMODEL Override generation arguments ... Done\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:45.061372Z","iopub.execute_input":"2023-08-06T15:24:45.062388Z","iopub.status.idle":"2023-08-06T15:24:45.068892Z","shell.execute_reply.started":"2023-08-06T15:24:45.062348Z","shell.execute_reply":"2023-08-06T15:24:45.067783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## BASE MODEL SAVE\nbase_model.save_pretrained('./basemodel/')\nprint(\"BASE MODEL SAVED\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:45.073115Z","iopub.execute_input":"2023-08-06T15:24:45.073445Z","iopub.status.idle":"2023-08-06T15:24:46.584954Z","shell.execute_reply.started":"2023-08-06T15:24:45.073416Z","shell.execute_reply":"2023-08-06T15:24:46.583949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## prepare_model for int8 training\n# from peft import prepare_model_for_int8_training\n\nmodel = prepare_model_for_int8_training(base_model)\nprint(\"prepare_model_for_int8_training_Completed\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:46.586395Z","iopub.execute_input":"2023-08-06T15:24:46.586764Z","iopub.status.idle":"2023-08-06T15:24:46.596061Z","shell.execute_reply.started":"2023-08-06T15:24:46.586731Z","shell.execute_reply":"2023-08-06T15:24:46.594988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Lora Config\n# from peft import LoraConfig, PeftModel, LoraModel, LoraConfig, get_peft_model\n\nlora_config = LoraConfig(r= config['lora_r'], \n                         lora_alpha= config['lora_alpha'], \n                         target_modules=config['lora_target_modules'], # [\"q_proj\", \"v_proj\"],\n                         lora_dropout= config['lora_dropout_p'], # 0.1\n                         bias=\"none\")\nprint(\"LoRA Config defined\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:46.597667Z","iopub.execute_input":"2023-08-06T15:24:46.598784Z","iopub.status.idle":"2023-08-06T15:24:46.609453Z","shell.execute_reply.started":"2023-08-06T15:24:46.598750Z","shell.execute_reply":"2023-08-06T15:24:46.608013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Get PeftModel\nmodel = get_peft_model(base_model, lora_config)\n# model = model.to(device) # CUDA --> Accelerate will do this soon\nmodel.print_trainable_parameters()","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:46.610716Z","iopub.execute_input":"2023-08-06T15:24:46.611164Z","iopub.status.idle":"2023-08-06T15:24:46.850781Z","shell.execute_reply.started":"2023-08-06T15:24:46.611131Z","shell.execute_reply":"2023-08-06T15:24:46.849786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## model cuda? Check!\n# next(model.parameters()).is_cuda","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# sam_out = model(**sample)\n# sam_out.keys()\n\n#### odict_keys(['loss', 'logits', 'encoder_last_hidden_state'])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# sam_out.loss\n#### tensor(0.4957, grad_fn=<NllLossBackward0>)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# sam_out.logits.shape, sam_out.logits[0, 0]\n##### (torch.Size([4, 104, 51865]),\n##### tensor([-6.3497, -4.7272, -2.5783,  ..., -2.7204, -3.8126, -2.3563], grad_fn=<SelectBackward0>))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# #### Whisper in Transformers.ipynb \n# ## Transcribe Sample from: https://colab.research.google.com/drive/16HO7if9iwfpSJzhqlaNOu6iiMhUBLMKE?usp=sharing\n# ## input_features = processor.feature_extractor(arrays, padding=\"max_length\", max_length=480_000, return_tensors=\"pt\").input_features.to(device)\n# ## sequences = model.generate(input_features, max_length=480_000, use_cache=True)\n# ## results = processor.tokenizer.batch_decode(sequences, skip_special_tokens=True, normalize=True)\n\n# forced_decoder_ids = processor.get_decoder_prompt_ids(language=\"bengali\", task=\"transcribe\")\n# sequences = model.generate(**sample, use_cache=True, max_length= 256, num_beams = 3, forced_decoder_ids=forced_decoder_ids)\n# results = processor.tokenizer.batch_decode(sequences, skip_special_tokens=True, normalize = True)\n# print(results)\n\n\n###### RESULTS decoded\n###### ['ল ম তম ক সব দ খয দ ব', 'পর কত র র স প ল সট জ র উপর ব শ য ছ ল ন',\n######  'তর প র সপট ন সপট র জ ব তর র সপ ল র সম লয প সম প সম প সম প সম প সম প সম প সম প সম প র সম প', \n######  'ত ন অসটর ল য র পদ হ য প নর য ট সট অ শ ন ন']\n\n##### LABELS decoded\n##### ['ল ম ত ম ক সব দ খ য দ ব', 'বকত র সকল ই ষট জ র উপর বস য ছ ল ন', \n#####  'ওক এর মলয দ ত হব', 'ত ন অসটর ল য র পকষ পনর য ট সট অ শ ন ন']","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### SAMPLES - LABELS\n## Labels -> Sentences\n# eval_labels = []\n# labels = sample['labels'].detach().cpu().numpy()\n# labels[labels == -100] = processor.tokenizer.pad_token_id\n# labels = processor.tokenizer.batch_decode(labels, skip_special_tokens=True, normalize = True)\n# print(labels)\n\n##### LABELS\n#### ['ল ম ত ম ক সব দ খ য দ ব', 'বকত র সকল ই ষট জ র উপর বস য ছ ল ন', \n#### 'ওক এর মলয দ ত হব', 'ত ন অসটর ল য র পকষ পনর য ট সট অ শ ন ন']","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# #### CASE 1) SAMPLES->GENERATE-> WER\n\n# wer_score = wer.compute(predictions=results, references=labels)\n# print(wer_score)\n#### 1.2","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"##### CASE 2) SAMPLE -> FORWARD -> WER\n# preds = processor.tokenizer.batch_decode(sam_out.logits.argmax(-1), skip_special_tokens=True, normalize = True)\n# print(preds) \n\n#### FORWARD decoded\n###['লধ ম তধ ম ক সব দ খ য দ ব', 'দ ত র স সটজ র উপর দ য ল ন', \n### 'এ ধ সদ সল গ এ সধ এধধ', 'ত ন অসটর ল য র পদধ পধনর য ট সট অ শ ন ন']\n\n##### LABELS decoded\n##### ['ল ম ত ম ক সব দ খ য দ ব', 'বকত র সকল ই ষট জ র উপর বস য ছ ল ন', \n#####  'ওক এর মলয দ ত হব', 'ত ন অসটর ল য র পকষ পনর য ট সট অ শ ন ন']\n\n# wer_score_f = wer.compute(predictions=preds, references=labels)\n# print(wer_score_f)\n##### 0.4444444444444444","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optimizer, Scheduler\n : moved to the run_train part","metadata":{}},{"cell_type":"code","source":"# optimizer \n# from transformers import AdamW\n\n### optimizer 01\n# optimizer = torch.optim.Adam(model.parameters(), lr = config['lr'], weight_decay = config['weight_decay'])\n\n### optimizer 02\noptimizer = AdamW(model.parameters(), \n                  lr = config['learning_rate'], \n                  weight_decay = config['weight_decay'])\nprint(\"optimizer defined\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:46.852536Z","iopub.execute_input":"2023-08-06T15:24:46.853093Z","iopub.status.idle":"2023-08-06T15:24:46.865434Z","shell.execute_reply.started":"2023-08-06T15:24:46.853058Z","shell.execute_reply":"2023-08-06T15:24:46.864569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def fetch_scheduler(optimizer):\n    \n#     if config['scheduler'] == 'CosineAnnealingLR':\n#         scheduler = lr_scheduler.CosineAnnealingLR(optimizer, \n#                                                    T_max = config['T_max'], \n#                                                    eta_min= config['min_lr'])\n        \n#     elif config['scheduler'] == 'CosineAnnealingWarmRestarts':\n#         scheduler = lr_scheduler.CosineAnnealingWarmRestarts(optimizer, \n#                                                    T_0 = config['T_0'], \n#                                                    eta_min= config['min_lr'])\n#     elif config['scheduler'] == None:\n#         return None\n    \n#     return scheduler\n\n### scheduler 01\n# scheduler = fetch_scheduler(optimizer)\n# print(\"Scheduler Loaded\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch.optim as optim\n\nclass CosineWarmupScheduler(optim.lr_scheduler._LRScheduler):\n    def __init__(self, optimizer, warmup, max_iters):\n        self.warmup = warmup\n        self.max_num_iters = max_iters\n        super().__init__(optimizer)\n\n    def get_lr(self):\n        lr_factor = self.get_lr_factor(epoch=self.last_epoch)\n        return [base_lr * lr_factor for base_lr in self.base_lrs]\n\n    def get_lr_factor(self, epoch):\n        lr_factor = 0.5 * (1 + np.cos(np.pi * epoch / self.max_num_iters))\n        if epoch <= self.warmup:\n            lr_factor *= epoch * 1.0 / self.warmup\n        return lr_factor\n    \n    \n### scheduler 02\nscheduler = CosineWarmupScheduler(optimizer=optimizer, warmup=100, max_iters=2000)\nprint(\"Scheduler Loaded\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:46.866978Z","iopub.execute_input":"2023-08-06T15:24:46.867541Z","iopub.status.idle":"2023-08-06T15:24:46.884421Z","shell.execute_reply.started":"2023-08-06T15:24:46.867508Z","shell.execute_reply":"2023-08-06T15:24:46.883451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Accelerator","metadata":{}},{"cell_type":"code","source":"# from accelerate import Accelerator\n\n# Needs\naccelerator = Accelerator()\n\n# prepare\nmodel, train_loader, valid_loader, optimizer, scheduler = accelerator.prepare(\n    model, train_loader, valid_loader, optimizer, scheduler)","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:46.885903Z","iopub.execute_input":"2023-08-06T15:24:46.886590Z","iopub.status.idle":"2023-08-06T15:24:52.644669Z","shell.execute_reply.started":"2023-08-06T15:24:46.886554Z","shell.execute_reply":"2023-08-06T15:24:52.643683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## model cuda? Check!\nnext(model.parameters()).is_cuda","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:52.646197Z","iopub.execute_input":"2023-08-06T15:24:52.646585Z","iopub.status.idle":"2023-08-06T15:24:52.653136Z","shell.execute_reply.started":"2023-08-06T15:24:52.646550Z","shell.execute_reply":"2023-08-06T15:24:52.651923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TorchTracemalloc: I don't know what this is\n- [source](TorchTracemalloc: I don't know what this is)","metadata":{}},{"cell_type":"code","source":"import sys\nimport threading\nimport psutil\n\n# Converting Bytes to Megabytes\ndef b2mb(x):\n    return int(x / 2**20)\n\n# This context manager is used to track the peak memory usage of the process\nclass TorchTracemalloc:\n    def __enter__(self):\n        gc.collect()\n        torch.cuda.empty_cache()\n        torch.cuda.reset_max_memory_allocated()  # reset the peak gauge to zero\n        self.begin = torch.cuda.memory_allocated()\n        self.process = psutil.Process()\n\n        self.cpu_begin = self.cpu_mem_used()\n        self.peak_monitoring = True\n        peak_monitor_thread = threading.Thread(target=self.peak_monitor_func)\n        peak_monitor_thread.daemon = True\n        peak_monitor_thread.start()\n        return self\n\n    def cpu_mem_used(self):\n        \"\"\"get resident set size memory for the current process\"\"\"\n        return self.process.memory_info().rss\n\n    def peak_monitor_func(self):\n        self.cpu_peak = -1\n\n        while True:\n            self.cpu_peak = max(self.cpu_mem_used(), self.cpu_peak)\n\n            # can't sleep or will not catch the peak right (this comment is here on purpose)\n            # time.sleep(0.001) # 1msec\n\n            if not self.peak_monitoring:\n                break\n\n    def __exit__(self, *exc):\n        self.peak_monitoring = False\n\n        gc.collect()\n        torch.cuda.empty_cache()\n        self.end = torch.cuda.memory_allocated()\n        self.peak = torch.cuda.max_memory_allocated()\n        self.used = b2mb(self.end - self.begin)\n        self.peaked = b2mb(self.peak - self.begin)\n\n        self.cpu_end = self.cpu_mem_used()\n        self.cpu_used = b2mb(self.cpu_end - self.cpu_begin)\n        self.cpu_peaked = b2mb(self.cpu_peak - self.cpu_begin)\n        # print(f\"delta used/peak {self.used:4d}/{self.peaked:4d}\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:52.654750Z","iopub.execute_input":"2023-08-06T15:24:52.655382Z","iopub.status.idle":"2023-08-06T15:24:52.668228Z","shell.execute_reply.started":"2023-08-06T15:24:52.655350Z","shell.execute_reply":"2023-08-06T15:24:52.667262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# train_one_epoch","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model,\n                    epoch,\n                    accelerator, \n                    train_dataloader,\n                    optimizer,\n                    lr_scheduler,\n                    grad_clipping,\n                    wer):\n\n    with TorchTracemalloc() as tracemalloc:\n        model.train()\n        total_loss = 0\n#         eval_preds, eval_labels = [], []\n\n        for step, batch in enumerate(tqdm(train_dataloader)):\n            # outputs = model(**batch.to(device))\n            outputs = model(**batch)\n            loss = outputs.loss\n            total_loss += loss.detach().float()\n            accelerator.backward(loss)\n\n            # Gradient-Clipping | source: https://velog.io/@seven7724/Transformer-계열의-훈련-Tricks\n            max_norm = 5\n            if grad_clipping:\n                #print(\"Gradient Clipping Turned On | max_norm: \", max_norm)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)\n\n            optimizer.step()\n            lr_scheduler.step()\n            optimizer.zero_grad()\n\n            # Generate(=Transcribe) Process to the eval_preds\n#             outputs = model.generate(**batch, max_length=225, use_cache=True)\n#             outputs = accelerator.pad_across_processes(outputs, dim=1, pad_index=processor.tokenizer.pad_token_id)\n#             outputs = accelerator.gather_for_metrics(outputs)\n#             eval_preds.extend(processor.tokenizer.batch_decode(outputs, skip_special_tokens=True, normalize=True))\n\n#             labels = batch['labels'].detach().cpu().numpy()\n#             labels[labels == -100] = processor.tokenizer.pad_token_id\n#             labels = processor.tokenizer.batch_decode(labels)\n\n#             eval_preds += preds\n#             eval_labels += labels\n\n    # Printing the GPU memory usage details such as allocated memory, peak memory, and total memory usage\n    accelerator.print(\"GPU Memory before entering the train : {}\".format(b2mb(tracemalloc.begin)))\n    accelerator.print(\"GPU Memory consumed at the end of the train (end-begin): {}\".format(tracemalloc.used))\n    accelerator.print(\"GPU Peak Memory consumed during the train (max-begin): {}\".format(tracemalloc.peaked))\n    accelerator.print(\n        \"GPU Total Peak Memory consumed during the train (max): {}\".format(\n            tracemalloc.peaked + b2mb(tracemalloc.begin)\n        )\n    )\n\n    accelerator.print(\"CPU Memory before entering the train : {}\".format(b2mb(tracemalloc.cpu_begin)))\n    accelerator.print(\"CPU Memory consumed at the end of the train (end-begin): {}\".format(tracemalloc.cpu_used))\n    accelerator.print(\"CPU Peak Memory consumed during the train (max-begin): {}\".format(tracemalloc.cpu_peaked))\n    accelerator.print(\n        \"CPU Total Peak Memory consumed during the train (max): {}\".format(\n            tracemalloc.cpu_peaked + b2mb(tracemalloc.cpu_begin)\n            )\n        )\n\n    ## Train Epoch Loss\n    train_epoch_loss = total_loss / len(train_dataloader)\n    accelerator.print(f\"{epoch=}: {train_epoch_loss=}\")\n\n#     eval_preds = [pred.strip().replace(\".\", \"\") for pred in eval_preds]\n#     eval_labels = [label.strip().replace(\".\", \"\") for label in eval_labels]\n    \n    ## WER SCORE\n#     wer_score = wer.compute(predictions=eval_preds, references=eval_labels)\n#     accelerator.print(f\"{wer_score=}\")\n    \n#     del eval_preds, eval_labels\n    \n    torch.cuda.empty_cache()\n    _ = gc.collect()\n\n    return train_epoch_loss#, wer_score","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:52.673965Z","iopub.execute_input":"2023-08-06T15:24:52.674227Z","iopub.status.idle":"2023-08-06T15:24:52.687586Z","shell.execute_reply.started":"2023-08-06T15:24:52.674204Z","shell.execute_reply":"2023-08-06T15:24:52.686318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# loss = train_one_epoch(model = model,\n#                        epoch = 1,\n#                        accelerator = accelerator,\n#                        train_dataloader = train_loader,\n#                        optimizer = optimizer,\n#                        lr_scheduler = scheduler,\n#                        grad_clipping = False,\n#                        wer= wer)\n\n###### train_one_epoch() #######\n## GPU[T4 x2] BS:8 | EPOCHS: 1 | train = 10000\n##  >> VRAM: 10.4GB | 3MB\n##  >> train: 30 m","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# valid_one_epoch","metadata":{}},{"cell_type":"code","source":"def valid_one_epoch(model,\n                    epoch,\n                    accelerator,\n                    eval_dataloader,\n                    wer\n                    ):\n    \n    ### This is for model.generate()\n    forced_decoder_ids = processor.get_decoder_prompt_ids(language=\"bengali\", task=\"transcribe\")\n    \n    with TorchTracemalloc() as tracemalloc:\n        model.eval()\n        total_loss = 0\n        eval_preds = []\n        eval_labels = []\n        \n        \n        for step, batch in enumerate(tqdm(eval_dataloader)):\n            with torch.no_grad():\n                # outputs = model(**batch.to(device))\n                outputs = model(**batch)\n                loss = outputs.loss\n                total_loss += loss.detach().float()\n\n                # Generate(=Transcribe) Process to the eval_preds\n                outputs = accelerator.unwrap_model(model).generate(**batch, \n                                                                   forced_decoder_ids=forced_decoder_ids,\n                                                                   # max_length=512, \n                                                                   # num_beams = 3, \n                                                                   # early_stopping=True,\n                                                                   use_cache=True)\n                outputs = accelerator.pad_across_processes(outputs, dim=1, pad_index=processor.tokenizer.pad_token_id)\n                outputs = accelerator.gather_for_metrics(outputs)\n                eval_preds.extend(processor.tokenizer.batch_decode(outputs, skip_special_tokens=True, normalize=True))\n\n                # WER, ACC\n                labels = batch['labels'].detach().cpu().numpy()\n                labels[labels == -100] = processor.tokenizer.pad_token_id\n                labels = processor.tokenizer.batch_decode(labels, skip_special_tokens=True, normalize=True)\n\n                eval_labels += labels\n\n\n    # Printing the GPU memory usage details such as allocated memory, peak memory, and total memory usage\n    accelerator.print(\"GPU Memory before entering the eval : {}\".format(b2mb(tracemalloc.begin)))\n    accelerator.print(\"GPU Memory consumed at the end of the eval (end-begin): {}\".format(tracemalloc.used))\n    accelerator.print(\"GPU Peak Memory consumed during the eval (max-begin): {}\".format(tracemalloc.peaked))\n    accelerator.print(\n        \"GPU Total Peak Memory consumed during the eval (max): {}\".format(\n            tracemalloc.peaked + b2mb(tracemalloc.begin)\n        )\n    )\n\n    accelerator.print(\"CPU Memory before entering the eval : {}\".format(b2mb(tracemalloc.cpu_begin)))\n    accelerator.print(\"CPU Memory consumed at the end of the eval (end-begin): {}\".format(tracemalloc.cpu_used))\n    accelerator.print(\"CPU Peak Memory consumed during the eval (max-begin): {}\".format(tracemalloc.cpu_peaked))\n    accelerator.print(\n        \"CPU Total Peak Memory consumed during the eval (max): {}\".format(\n            tracemalloc.cpu_peaked + b2mb(tracemalloc.cpu_begin)\n        )\n    )\n\n    # Epoch Loss\n    eval_epoch_loss = total_loss / len(eval_dataloader)\n    accelerator.print(f\"{epoch=}: {eval_epoch_loss=}\")\n\n    eval_preds = [pred.strip().replace(\".\", \"\") for pred in eval_preds]\n    eval_labels = [label.strip().replace(\".\", \"\") for label in eval_labels]\n\n    # Print Samples\n    accelerator.print(f\"{eval_preds[:2]=}\")\n    accelerator.print(f\"{eval_labels[:2]=}\")\n    \n    # WER SCORE\n    wer_score = wer.compute(predictions=eval_preds, references=eval_labels)\n    accelerator.print(f\"{wer_score=}\")\n    \n    del eval_preds, eval_labels#, accuracy\n\n    torch.cuda.empty_cache()\n    _ = gc.collect()\n\n    return eval_epoch_loss, wer_score","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:52.689966Z","iopub.execute_input":"2023-08-06T15:24:52.690685Z","iopub.status.idle":"2023-08-06T15:24:52.707319Z","shell.execute_reply.started":"2023-08-06T15:24:52.690651Z","shell.execute_reply":"2023-08-06T15:24:52.706221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# loss, wer = valid_one_epoch(model = model,\n#                             epoch = 1,\n#                             accelerator = accelerator,\n#                             eval_dataloader = valid_loader,\n#                             wer=wer)\n\n###### valid one epoch() #######\n## GPU[T4 x2] BS:32 | generate: max_length: 225) | EPOCHS: 1 | valid = 2000\n##  >> VRAM: 8.8GB | 3MB\n##  >> Valid: 15 m","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# run_train","metadata":{}},{"cell_type":"code","source":"def run_train(model = model,\n              logging = True,\n              accelerator = accelerator,\n              train_loader= train_loader,\n              valid_loader = valid_loader,\n              optimizer = optimizer,\n              scheduler = scheduler,\n              wer= wer,\n              n_epochs = config['n_epochs']):\n\n    if logging:\n        # To automatically log gradients\n        wandb.watch(model, log_freq=100)\n\n    if torch.cuda.is_available():\n        print(\"INFO: GPU - {}\\n\".format(torch.cuda.get_device_name()))\n\n    start = time.time()\n    train_hs, valid_hs, eval_accs, train_accs, eval_wers, train_wers = [], [], [], [], [], []\n\n    lowest_epoch, lowest_loss, best_epoch, best_wer = np.inf, np.inf, np.inf, np.inf\n\n    for epoch in range(n_epochs):\n        train_epoch_loss = train_one_epoch(model = model,\n                                           epoch = epoch,\n                                           accelerator = accelerator, \n                                           train_dataloader= train_loader,\n                                           optimizer = optimizer, \n                                           lr_scheduler = scheduler,\n                                           grad_clipping = False,\n                                           wer=wer)\n\n        eval_epoch_loss, eval_wer = valid_one_epoch(model=model,\n                                                    epoch = epoch,\n                                                    accelerator = accelerator,\n                                                    eval_dataloader = valid_loader,\n                                                    wer=wer)\n\n        if logging:\n            # Log the metrics\n            wandb.log({\"train/loss\": train_epoch_loss})\n            wandb.log({\"eval/loss\": eval_epoch_loss})\n\n            # Log the metrics\n            # wandb.log({\"train/WER\": train_wer})\n            wandb.log({\"eval/WER\": eval_wer})\n\n        train_hs.append(train_epoch_loss)\n        valid_hs.append(eval_epoch_loss)\n        \n        # train_wers.append(train_wer)\n        eval_wers.append(eval_wer)\n\n        print()\n        print(f\"Epoch:{epoch} | TL:{train_epoch_loss:.3e} | VL:{eval_epoch_loss:.3e} | LL:{lowest_loss: .3e}|\")\n#         print(f\"Train WER: {train_wer:.3f} | Valid WER: {eval_wer:.2f} |\")\n        print(f\"Valid WER: {eval_wer:.2f} |\")\n        print()\n\n        if eval_wer <= best_wer:\n            print(f\"{b_}Eval WER Improved({best_wer:.2f}) --> ({eval_wer:.2f})\")\n            best_wer = eval_wer\n            best_epoch = epoch\n\n            # peft_model save!\n            PEFT_MODEL_PATH = './peft'\n            model.save_pretrained(PEFT_MODEL_PATH)\n            #\n            UNWRAPPED_PEFT_MODEL_PATH = './unwrapped_peft'\n            accelerator.wait_for_everyone()\n            unwrapped_model = accelerator.unwrap_model(model)\n            unwrapped_model.save_pretrained(UNWRAPPED_PEFT_MODEL_PATH, \n#                                             save_function=accelerator.save, \n#                                             state_dict=accelerator.get_state_dict(model)\n                                           )\n            ## without this, model could be saved :) \n            print(f\"{y_}PEFT LORA Model (whose perfomance is better) Saved at {epoch} EPOCH\")\n\n        if eval_epoch_loss < lowest_loss:\n            print(f\"{b_}Eval Loss Improved({lowest_loss:.3e}) --> ({eval_epoch_loss:.3e})\")\n            lowest_loss = eval_epoch_loss\n            lowest_epoch = epoch\n\n    end = time.time()\n    time_elapsed = end - start\n    print('Training complete in {:.0f}h {:.0f}m {:.0f}s'.format(time_elapsed // 3600, (time_elapsed % 3600) // 60, (time_elapsed % 3600) % 60))\n    print(\"Best Loss : %.4e at %d th Epoch\" % (lowest_loss, lowest_epoch))\n    print(\"Best WER : %.2f at %d th Epoch\" % (best_wer, best_epoch))\n\n    result = dict()\n    # Loss\n    result[\"train/loss\"] = train_hs\n    result[\"eval/loss\"] = valid_hs\n\n    # WER\n    # result[\"train/WER\"] = train_wers\n    result[\"eval/WER\"] = eval_wers\n    \n    \n    return result","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:52.708846Z","iopub.execute_input":"2023-08-06T15:24:52.709173Z","iopub.status.idle":"2023-08-06T15:24:52.726869Z","shell.execute_reply.started":"2023-08-06T15:24:52.709143Z","shell.execute_reply":"2023-08-06T15:24:52.725383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize","metadata":{}},{"cell_type":"code","source":"################### Visualize #########################\ndef make_plot(result, stage = \"Loss\"):\n\n    plot_from = 0\n\n    if stage == \"Loss\":\n        trains = \"train/loss\"\n        valids = \"eval/loss\"\n\n    elif stage == \"WER\":\n        #trains = \"train/WER\" \n        trains = \"eval/WER\"\n        valids = \"eval/WER\"\n\n\n    plt.figure(figsize=(10, 6))\n\n    plt.title(f\"Train/Valid {stage} History\", fontsize = 20)\n\n    ## Modified for converting Type\n    if type(result[trains][0]) == torch.Tensor:\n        result[trains] = [num.detach().cpu().item() for num in result[trains]]\n        result[valids] = [num.detach().cpu().item() for num in result[valids]]\n\n    plt.plot(\n        range(0, len(result[trains][plot_from:])),\n        result[trains][plot_from:],\n        label = trains\n        )\n\n    plt.plot(\n        range(0, len(result[valids][plot_from:])),\n        result[valids][plot_from:],\n        label = valids\n        )\n\n    plt.legend()\n    if stage == \"Loss\":\n        plt.yscale('log')\n    plt.grid(True)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:52.728044Z","iopub.execute_input":"2023-08-06T15:24:52.728568Z","iopub.status.idle":"2023-08-06T15:24:52.741205Z","shell.execute_reply.started":"2023-08-06T15:24:52.728533Z","shell.execute_reply":"2023-08-06T15:24:52.740132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Wandb init()","metadata":{}},{"cell_type":"code","source":"name = f\"[GEN][{config['model_name'].split('/')[1]}] Peft/LoRA_EP:{config['n_epochs']}_BS:({config['train_batch_size']} & {config['valid_batch_size']})_LR:{config['learning_rate']}\"\nname ","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:24:52.742639Z","iopub.execute_input":"2023-08-06T15:24:52.743166Z","iopub.status.idle":"2023-08-06T15:24:52.760424Z","shell.execute_reply.started":"2023-08-06T15:24:52.743134Z","shell.execute_reply":"2023-08-06T15:24:52.759210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run = wandb.init(project=config['competition'], \n                 config= config,\n                 job_type='Train',\n                 tags=[config['model_name']],\n                 name=name,\n                 anonymous='must')","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:25:19.902579Z","iopub.execute_input":"2023-08-06T15:25:19.903164Z","iopub.status.idle":"2023-08-06T15:25:53.345166Z","shell.execute_reply.started":"2023-08-06T15:25:19.903128Z","shell.execute_reply":"2023-08-06T15:25:53.344206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Let's Train: notebook_launcher","metadata":{}},{"cell_type":"code","source":"result = run_train( model = model,\n                    logging = True,\n                    accelerator = accelerator,\n                    train_loader= train_loader,\n                    valid_loader = valid_loader,\n                    optimizer = optimizer,\n                    scheduler = scheduler,\n                    wer= wer,\n                    n_epochs = config['n_epochs'])\n\n## GPU[T4 x2] BS:8, 32 | EPOCHS: 3 | train, valid = 10000, 2000\n##  >> VRAM: (train) 13.1GB | (valid) 8.9GB\n##  >> Time: (train)  30m   | (valid) 19m ","metadata":{"execution":{"iopub.status.busy":"2023-08-06T15:25:53.349772Z","iopub.execute_input":"2023-08-06T15:25:53.352817Z","iopub.status.idle":"2023-08-06T17:44:18.578729Z","shell.execute_reply.started":"2023-08-06T15:25:53.352779Z","shell.execute_reply":"2023-08-06T17:44:18.577225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train Completed","metadata":{}},{"cell_type":"code","source":"torch.cuda.empty_cache()\n_ = gc.collect()\n\nprint(\"Train Completed\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:44:18.580222Z","iopub.execute_input":"2023-08-06T17:44:18.585072Z","iopub.status.idle":"2023-08-06T17:44:18.938779Z","shell.execute_reply.started":"2023-08-06T17:44:18.585035Z","shell.execute_reply":"2023-08-06T17:44:18.937515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visulization of Loss, WER","metadata":{}},{"cell_type":"code","source":"make_plot(result, \"Loss\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:44:49.502968Z","iopub.execute_input":"2023-08-06T17:44:49.503356Z","iopub.status.idle":"2023-08-06T17:44:50.107162Z","shell.execute_reply.started":"2023-08-06T17:44:49.503325Z","shell.execute_reply":"2023-08-06T17:44:50.106158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"make_plot(result, \"WER\")","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:44:50.109216Z","iopub.execute_input":"2023-08-06T17:44:50.109889Z","iopub.status.idle":"2023-08-06T17:44:50.461186Z","shell.execute_reply.started":"2023-08-06T17:44:50.109854Z","shell.execute_reply":"2023-08-06T17:44:50.455913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Run Finished","metadata":{}},{"cell_type":"code","source":"run.finish()","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:44:57.207320Z","iopub.execute_input":"2023-08-06T17:44:57.208033Z","iopub.status.idle":"2023-08-06T17:45:01.900522Z","shell.execute_reply.started":"2023-08-06T17:44:57.207997Z","shell.execute_reply":"2023-08-06T17:45:01.899683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Display","metadata":{}},{"cell_type":"code","source":"## This is just to display the W&B run page in this interactive session\nfrom IPython import display\n\n# we create an IFrame and set the width and height\niF = display.IFrame(run.url, width = 1080, height = 720)\niF","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:45:26.661809Z","iopub.execute_input":"2023-08-06T17:45:26.662191Z","iopub.status.idle":"2023-08-06T17:45:26.673133Z","shell.execute_reply.started":"2023-08-06T17:45:26.662158Z","shell.execute_reply":"2023-08-06T17:45:26.671964Z"},"trusted":true},"execution_count":null,"outputs":[]}]}