{"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":6707460,"sourceType":"datasetVersion","datasetId":3865741}],"dockerImageVersionId":30674,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<h1>This notebook is purely inspired from <a href=\"https://www.kaggle.com/code/smjishanulislam/quickstart-with-whisper-small/notebook\"> here</a><h1>\n   <h1>This notebook is created for utilizing GPU time. Because of limited GPU resources in Kaggle, training and inference in the same notebook is not a good idea. Use this notebook only for training. After training, import the model as a dataset. For inference, use the code in the <a href=\"https://www.kaggle.com/code/samratabduljalil/bengali-speech-recognition-inference/\"> link</a>  <h1> ","metadata":{}},{"cell_type":"code","source":"import os\n\nimport pandas as pd\n\nimport librosa\nimport librosa.display\n\nimport numpy as np\n\nimport IPython.display as ipd\n\nimport matplotlib.pyplot as plt\n\nimport random\n\nfrom collections import Counter\n\nfrom sklearn.model_selection import train_test_split\n\nimport torch\nimport torchaudio\n\nfrom dataclasses import dataclass\nfrom typing import Any, Dict, List, Union\nfrom datasets import DatasetDict\nfrom datasets import Dataset as DS\n\nfrom transformers import (\n    WhisperFeatureExtractor,\n    WhisperTokenizer,\n    WhisperProcessor,\n    WhisperForConditionalGeneration,\n    Seq2SeqTrainingArguments,\n    Seq2SeqTrainer,\n    TrainerCallback,\n    TrainingArguments,\n    TrainerState,\n    TrainerControl,\n    EarlyStoppingCallback,\n    pipeline\n)\n\nfrom torchmetrics.text import WordErrorRate, CharErrorRate","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:47:43.677506Z","iopub.execute_input":"2024-04-06T01:47:43.677846Z","iopub.status.idle":"2024-04-06T01:48:07.955949Z","shell.execute_reply.started":"2024-04-06T01:47:43.677820Z","shell.execute_reply":"2024-04-06T01:48:07.955098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_DIR = '/kaggle/input/ben10/ben10'\ntrain_data_dir = f\"{BASE_DIR}/16_kHz_train_audio/\"\ntest_data_dir = f\"{BASE_DIR}/16_kHz_valid_audio/\"\ndata_path = f\"{BASE_DIR}/train.csv\"","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:07.957505Z","iopub.execute_input":"2024-04-06T01:48:07.958103Z","iopub.status.idle":"2024-04-06T01:48:07.962398Z","shell.execute_reply.started":"2024-04-06T01:48:07.958076Z","shell.execute_reply":"2024-04-06T01:48:07.961535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"split2path = {\n    \"train\": train_data_dir,\n    \"test\": test_data_dir,\n}","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:07.963621Z","iopub.execute_input":"2024-04-06T01:48:07.964095Z","iopub.status.idle":"2024-04-06T01:48:07.973154Z","shell.execute_reply.started":"2024-04-06T01:48:07.964068Z","shell.execute_reply":"2024-04-06T01:48:07.972244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = pd.read_csv(data_path)\ndata.sample(10)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:07.975259Z","iopub.execute_input":"2024-04-06T01:48:07.975538Z","iopub.status.idle":"2024-04-06T01:48:08.210253Z","shell.execute_reply.started":"2024-04-06T01:48:07.975514Z","shell.execute_reply":"2024-04-06T01:48:08.209149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def extract_split(filename):\n    filename_ = filename.split(\"_\")\n    split = filename_[0]\n    return split\n\ndef extract_district(filename):\n    filename_ = filename.split(\" \")[0]\n    district = filename_.split(\"_\")[1]\n    return district\n\ndef beautify_dataset(data):\n    splits = []\n    districts = []\n    newpaths = []\n    transcripts = []\n    \n    for i in range(len(data)):\n        filename, transcript = data.iloc[i]\n        split = extract_split(filename)\n        district = extract_district(filename)\n        dir_path = split2path[split]\n        composed_path = f\"{dir_path}{filename}\"\n        \n        if os.path.exists(composed_path) == False:\n            print(f\"{composed_path} does not exist.\")\n            continue\n        \n        # replace any newline characters\n        transcript = transcript.replace(\"\\n\", \" \")\n        transcript = \" \".join(transcript.split())\n        \n        splits.append(split)\n        districts.append(district)\n        newpaths.append(composed_path)\n        transcripts.append(transcript)\n    \n    data['file_path'] = newpaths\n    data['district'] = districts\n    data['split'] = splits\n    data['transcripts'] = transcripts\n    \n#     data.drop(columns=['file_name'], inplace=True)\n    \n    return data","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:08.211369Z","iopub.execute_input":"2024-04-06T01:48:08.211659Z","iopub.status.idle":"2024-04-06T01:48:08.220996Z","shell.execute_reply.started":"2024-04-06T01:48:08.211632Z","shell.execute_reply":"2024-04-06T01:48:08.220154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = beautify_dataset(data)\ndata.sample(20)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:08.223607Z","iopub.execute_input":"2024-04-06T01:48:08.224069Z","iopub.status.idle":"2024-04-06T01:48:37.674049Z","shell.execute_reply.started":"2024-04-06T01:48:08.224023Z","shell.execute_reply":"2024-04-06T01:48:37.673040Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data[data[\"transcripts\"] == \"<>\"]\ndata[data[\"transcripts\"] == \"\"]\ndata[data[\"transcripts\"] == \"..\"]","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:37.675220Z","iopub.execute_input":"2024-04-06T01:48:37.675521Z","iopub.status.idle":"2024-04-06T01:48:37.694960Z","shell.execute_reply.started":"2024-04-06T01:48:37.675494Z","shell.execute_reply":"2024-04-06T01:48:37.694035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(list(data[data['transcripts'] == ''].index))\ndata.drop(data[data['transcripts'] == ''].index, inplace=True)\n      \n# print(list(data[data['transcripts'] == '<>'].index))\ndata.drop(data[data['transcripts'] == \"<>\"].index, inplace=True)\n      \n# print(list(data[data['transcripts'] == '..'].index))\ndata.drop(data[data['transcripts'] == \"..\"].index, inplace=True)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:37.696982Z","iopub.execute_input":"2024-04-06T01:48:37.697363Z","iopub.status.idle":"2024-04-06T01:48:37.760577Z","shell.execute_reply.started":"2024-04-06T01:48:37.697333Z","shell.execute_reply":"2024-04-06T01:48:37.759774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data[\"transcripts\"] = data[\"transcripts\"].str.strip()","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:37.761771Z","iopub.execute_input":"2024-04-06T01:48:37.762564Z","iopub.status.idle":"2024-04-06T01:48:37.772205Z","shell.execute_reply.started":"2024-04-06T01:48:37.762525Z","shell.execute_reply":"2024-04-06T01:48:37.771271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TASK = \"transcribe\"\nMODEL_NAME = \"/kaggle/input/bengali-ai-asr-submission/bengali-whisper-medium\"","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:37.776378Z","iopub.execute_input":"2024-04-06T01:48:37.777108Z","iopub.status.idle":"2024-04-06T01:48:37.781370Z","shell.execute_reply.started":"2024-04-06T01:48:37.777080Z","shell.execute_reply":"2024-04-06T01:48:37.780465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"feature_extractor = WhisperFeatureExtractor.from_pretrained(MODEL_NAME)\ntokenizer = WhisperTokenizer.from_pretrained(MODEL_NAME, language='bn', task=TASK)\nprocessor = WhisperProcessor.from_pretrained(MODEL_NAME, language='bn', task=TASK)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:37.782410Z","iopub.execute_input":"2024-04-06T01:48:37.782677Z","iopub.status.idle":"2024-04-06T01:48:38.725553Z","shell.execute_reply.started":"2024-04-06T01:48:37.782654Z","shell.execute_reply":"2024-04-06T01:48:38.724774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ids = tokenizer.encode(\"\")\nids","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:38.726636Z","iopub.execute_input":"2024-04-06T01:48:38.726943Z","iopub.status.idle":"2024-04-06T01:48:38.733377Z","shell.execute_reply.started":"2024-04-06T01:48:38.726918Z","shell.execute_reply":"2024-04-06T01:48:38.732521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tokenizer.decode(ids)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:38.734412Z","iopub.execute_input":"2024-04-06T01:48:38.734707Z","iopub.status.idle":"2024-04-06T01:48:38.767744Z","shell.execute_reply.started":"2024-04-06T01:48:38.734684Z","shell.execute_reply":"2024-04-06T01:48:38.766797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@dataclass\nclass DataCollatorSpeechSeq2SeqWithPadding:\n    processor: Any\n\n    def __call__(self, features: List[Dict[str, Union[List[int], torch.Tensor]]]) -> Dict[str, torch.Tensor]:\n        # split inputs and labels since they have to be of different lengths and need different padding methods\n        # first treat the audio inputs by simply returning torch tensors\n        input_features = [{\"input_features\": feature[\"input_features\"]} for feature in features]\n        batch = self.processor.feature_extractor.pad(input_features, return_tensors=\"pt\")\n\n        # get the tokenized label sequences\n        label_features = [{\"input_ids\": feature[\"labels\"]} for feature in features]\n        # pad the labels to max length\n        labels_batch = self.processor.tokenizer.pad(label_features, 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        # if bos token is appended in previous tokenization step,\n        # cut bos token here as it's append later anyways\n        if (labels[:, 0] == self.processor.tokenizer.bos_token_id).all().cpu().item():\n            labels = labels[:, 1:]\n\n        batch[\"labels\"] = labels\n        \n        torch.cuda.empty_cache()\n\n        return batch","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:38.769091Z","iopub.execute_input":"2024-04-06T01:48:38.769922Z","iopub.status.idle":"2024-04-06T01:48:38.778644Z","shell.execute_reply.started":"2024-04-06T01:48:38.769895Z","shell.execute_reply":"2024-04-06T01:48:38.777773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_collator = DataCollatorSpeechSeq2SeqWithPadding(processor=processor)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:38.779774Z","iopub.execute_input":"2024-04-06T01:48:38.780106Z","iopub.status.idle":"2024-04-06T01:48:38.787863Z","shell.execute_reply.started":"2024-04-06T01:48:38.780063Z","shell.execute_reply":"2024-04-06T01:48:38.786975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_dataset(example):\n    audio_path = example[\"file_path\"]\n    \n    # load the audio using librosa or torch audio (as you wish)\n    audio, sr = librosa.load(audio_path, sr=16_000)\n    \n    example[\"input_features\"] = feature_extractor(audio, sampling_rate=sr).input_features[0]\n    \n    example[\"labels\"] = tokenizer(f\"{example['transcripts']}\", max_length=448, padding=True, truncation=True).input_ids\n    \n    return example\n\n\ndef filter_inputs(input_audio):\n    \"\"\"filter inputs with zero input length\"\"\"\n    return 0 < len(input_audio)\n\n\ndef filter_labels(input_labels):\n    \"\"\"filter empty label sequences\"\"\"\n    return 0 < len(input_labels)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:38.788924Z","iopub.execute_input":"2024-04-06T01:48:38.789165Z","iopub.status.idle":"2024-04-06T01:48:38.798092Z","shell.execute_reply.started":"2024-04-06T01:48:38.789143Z","shell.execute_reply":"2024-04-06T01:48:38.797176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = data[data[\"split\"] == \"train\"]","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:38.799221Z","iopub.execute_input":"2024-04-06T01:48:38.799819Z","iopub.status.idle":"2024-04-06T01:48:38.820095Z","shell.execute_reply.started":"2024-04-06T01:48:38.799786Z","shell.execute_reply":"2024-04-06T01:48:38.819180Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\n    adjust test size accordingly.\n\"\"\"\ntrain_df, eval_df = train_test_split(train_df, test_size=0.01, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:38.821310Z","iopub.execute_input":"2024-04-06T01:48:38.821579Z","iopub.status.idle":"2024-04-06T01:48:38.829253Z","shell.execute_reply.started":"2024-04-06T01:48:38.821556Z","shell.execute_reply":"2024-04-06T01:48:38.828269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_df), len(eval_df)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:38.830409Z","iopub.execute_input":"2024-04-06T01:48:38.830775Z","iopub.status.idle":"2024-04-06T01:48:38.837387Z","shell.execute_reply.started":"2024-04-06T01:48:38.830719Z","shell.execute_reply":"2024-04-06T01:48:38.836326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ben_reg_voice_ds = DatasetDict()\n\ntrain_split = DS.from_pandas(train_df)\neval_split = DS.from_pandas(eval_df)\n\nds_splits = DatasetDict({\n    'train': train_split,\n    'eval': eval_split\n})","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:38.838533Z","iopub.execute_input":"2024-04-06T01:48:38.838953Z","iopub.status.idle":"2024-04-06T01:48:38.901355Z","shell.execute_reply.started":"2024-04-06T01:48:38.838922Z","shell.execute_reply":"2024-04-06T01:48:38.900555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_splits = ds_splits.remove_columns([\"split\"])\nprint(ds_splits)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:38.902490Z","iopub.execute_input":"2024-04-06T01:48:38.902877Z","iopub.status.idle":"2024-04-06T01:48:38.912625Z","shell.execute_reply.started":"2024-04-06T01:48:38.902839Z","shell.execute_reply":"2024-04-06T01:48:38.911741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.object = object","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:38.913847Z","iopub.execute_input":"2024-04-06T01:48:38.914455Z","iopub.status.idle":"2024-04-06T01:48:38.918776Z","shell.execute_reply.started":"2024-04-06T01:48:38.914420Z","shell.execute_reply":"2024-04-06T01:48:38.917775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_splits = ds_splits.map(prepare_dataset, remove_columns=ds_splits.column_names[\"train\"],\n                          num_proc=2 # open for multithreadding\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:48:38.921073Z","iopub.execute_input":"2024-04-06T01:48:38.921344Z","iopub.status.idle":"2024-04-06T01:54:40.868281Z","shell.execute_reply.started":"2024-04-06T01:48:38.921321Z","shell.execute_reply":"2024-04-06T01:54:40.867165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(ds_splits[\"train\"]), len(ds_splits[\"eval\"])","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:54:40.869972Z","iopub.execute_input":"2024-04-06T01:54:40.870288Z","iopub.status.idle":"2024-04-06T01:54:40.877498Z","shell.execute_reply.started":"2024-04-06T01:54:40.870257Z","shell.execute_reply":"2024-04-06T01:54:40.876472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cer = CharErrorRate()\nwer = WordErrorRate()","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:54:40.878553Z","iopub.execute_input":"2024-04-06T01:54:40.878929Z","iopub.status.idle":"2024-04-06T01:54:40.902906Z","shell.execute_reply.started":"2024-04-06T01:54:40.878895Z","shell.execute_reply":"2024-04-06T01:54:40.901989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def compute_metrics(pred):\n    pred_ids = pred.predictions\n    label_ids = pred.label_ids\n\n    label_ids[label_ids == -100] = tokenizer.pad_token_id\n\n    pred_str = tokenizer.batch_decode(pred_ids, skip_special_tokens=True)\n    label_str = tokenizer.batch_decode(label_ids, skip_special_tokens=True)\n\n    wer_res = wer(pred_str, label_str)\n    cer_res = cer(pred_str, label_str)\n    \n    \"\"\"\n        uncomment the next 3 lines if you want to see how the examples look like during eval \n    \"\"\"\n    print(\"WER:\",wer_res,\"| CER:\", cer_res) # to show up during running logs\n    print(\"Pred:\",pred_str[0])\n    print(\"Label:\",label_str[0])\n    \n    return {\"wer\": wer_res, \"cer\": cer_res}","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:54:40.904675Z","iopub.execute_input":"2024-04-06T01:54:40.905036Z","iopub.status.idle":"2024-04-06T01:54:40.913621Z","shell.execute_reply.started":"2024-04-06T01:54:40.904978Z","shell.execute_reply":"2024-04-06T01:54:40.912889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = WhisperForConditionalGeneration.from_pretrained(MODEL_NAME, device_map=\"auto\")","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:54:40.914861Z","iopub.execute_input":"2024-04-06T01:54:40.915459Z","iopub.status.idle":"2024-04-06T01:54:52.583918Z","shell.execute_reply.started":"2024-04-06T01:54:40.915426Z","shell.execute_reply":"2024-04-06T01:54:52.582958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_id = \"whisper-reg-ben\" #you can use different model from hugging face ","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:54:52.588419Z","iopub.execute_input":"2024-04-06T01:54:52.588759Z","iopub.status.idle":"2024-04-06T01:54:52.592814Z","shell.execute_reply.started":"2024-04-06T01:54:52.588704Z","shell.execute_reply":"2024-04-06T01:54:52.592016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#finetune this hyperparameter for getting the better result\n\ntraining_args = Seq2SeqTrainingArguments(\n    output_dir=model_id,\n    per_device_train_batch_size=2,\n    per_device_eval_batch_size=2,\n    gradient_accumulation_steps=1,\n    gradient_checkpointing=True,\n    fp16=True,\n    learning_rate=3e-4,\n    weight_decay=1e-2,\n    warmup_steps=1000,\n    num_train_epochs=1,\n    evaluation_strategy=\"steps\", # or \"epochs\"\n    predict_with_generate=True,\n#     generation_max_length=448,\n    save_steps=10000,\n    eval_steps=10000,\n    logging_steps=10000,\n    save_total_limit=1,\n    load_best_model_at_end=True,\n    metric_for_best_model=\"wer\",\n    greater_is_better=False,\n    push_to_hub=False,\n    report_to=\"none\",\n    remove_unused_columns=False,\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:54:52.593943Z","iopub.execute_input":"2024-04-06T01:54:52.594195Z","iopub.status.idle":"2024-04-06T01:54:52.607013Z","shell.execute_reply.started":"2024-04-06T01:54:52.594173Z","shell.execute_reply":"2024-04-06T01:54:52.606142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.generation_config.language = \"bn\"\nmodel.generation_config.task = \"transcribe\"\n\nmodel.generation_config.forced_decoder_ids = None\nmodel.config.suppress_tokens = [] # added later","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:54:52.608166Z","iopub.execute_input":"2024-04-06T01:54:52.608921Z","iopub.status.idle":"2024-04-06T01:54:52.618902Z","shell.execute_reply.started":"2024-04-06T01:54:52.608896Z","shell.execute_reply":"2024-04-06T01:54:52.618089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = Seq2SeqTrainer(\n    args=training_args,\n    model=model,\n    train_dataset=ds_splits[\"train\"],\n    eval_dataset=ds_splits[\"eval\"],\n    data_collator=data_collator,\n    tokenizer=processor.feature_extractor,\n    compute_metrics=compute_metrics,\n#     callbacks=[EarlyStoppingCallback(2, 1.0)]\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:54:52.619841Z","iopub.execute_input":"2024-04-06T01:54:52.620195Z","iopub.status.idle":"2024-04-06T01:54:52.647134Z","shell.execute_reply.started":"2024-04-06T01:54:52.620162Z","shell.execute_reply":"2024-04-06T01:54:52.646281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.train()\n\n# to use the high-level pipeline, ensure both the processor outputs and model outputs exist in the same dir\ntrainer.save_model(training_args.output_dir)\nprocessor.save_pretrained(training_args.output_dir)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T01:54:52.648124Z","iopub.execute_input":"2024-04-06T01:54:52.648444Z","iopub.status.idle":"2024-04-06T02:00:00.918393Z","shell.execute_reply.started":"2024-04-06T01:54:52.648413Z","shell.execute_reply":"2024-04-06T02:00:00.917006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1> go to output of the notebook you can find your best model .just create dataset just clicking on the three dot.then use that in inference notebook. inference notebook :<a href=\"\">Click here</a></h1>","metadata":{}}]}