{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":16880,"databundleVersionId":858837,"sourceType":"competition"},{"sourceId":924245,"sourceType":"datasetVersion","datasetId":464091}],"dockerImageVersionId":30627,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-18T12:55:04.054373Z","iopub.execute_input":"2023-12-18T12:55:04.054628Z","iopub.status.idle":"2023-12-18T12:56:51.010109Z","shell.execute_reply.started":"2023-12-18T12:55:04.054603Z","shell.execute_reply":"2023-12-18T12:56:51.009020Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_path = '/kaggle/input/deepfake-faces/metadata.csv'","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.012149Z","iopub.execute_input":"2023-12-18T12:56:51.012696Z","iopub.status.idle":"2023-12-18T12:56:51.017287Z","shell.execute_reply.started":"2023-12-18T12:56:51.012658Z","shell.execute_reply":"2023-12-18T12:56:51.016378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(dataset_path)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.018676Z","iopub.execute_input":"2023-12-18T12:56:51.019217Z","iopub.status.idle":"2023-12-18T12:56:51.201666Z","shell.execute_reply.started":"2023-12-18T12:56:51.019181Z","shell.execute_reply":"2023-12-18T12:56:51.200658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.203822Z","iopub.execute_input":"2023-12-18T12:56:51.204134Z","iopub.status.idle":"2023-12-18T12:56:51.225765Z","shell.execute_reply.started":"2023-12-18T12:56:51.204108Z","shell.execute_reply":"2023-12-18T12:56:51.224660Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.tail()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.226990Z","iopub.execute_input":"2023-12-18T12:56:51.227310Z","iopub.status.idle":"2023-12-18T12:56:51.237411Z","shell.execute_reply.started":"2023-12-18T12:56:51.227274Z","shell.execute_reply":"2023-12-18T12:56:51.236410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.shape","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.238727Z","iopub.execute_input":"2023-12-18T12:56:51.239397Z","iopub.status.idle":"2023-12-18T12:56:51.250826Z","shell.execute_reply.started":"2023-12-18T12:56:51.239360Z","shell.execute_reply":"2023-12-18T12:56:51.249947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.columns","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.251931Z","iopub.execute_input":"2023-12-18T12:56:51.252278Z","iopub.status.idle":"2023-12-18T12:56:51.267156Z","shell.execute_reply.started":"2023-12-18T12:56:51.252239Z","shell.execute_reply":"2023-12-18T12:56:51.266278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.duplicated().sum()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.268284Z","iopub.execute_input":"2023-12-18T12:56:51.268589Z","iopub.status.idle":"2023-12-18T12:56:51.334593Z","shell.execute_reply.started":"2023-12-18T12:56:51.268566Z","shell.execute_reply":"2023-12-18T12:56:51.333606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.isnull().sum()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.335839Z","iopub.execute_input":"2023-12-18T12:56:51.336213Z","iopub.status.idle":"2023-12-18T12:56:51.371259Z","shell.execute_reply.started":"2023-12-18T12:56:51.336176Z","shell.execute_reply":"2023-12-18T12:56:51.370351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.info()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.374855Z","iopub.execute_input":"2023-12-18T12:56:51.375767Z","iopub.status.idle":"2023-12-18T12:56:51.422285Z","shell.execute_reply.started":"2023-12-18T12:56:51.375729Z","shell.execute_reply":"2023-12-18T12:56:51.421360Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.nunique()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.423392Z","iopub.execute_input":"2023-12-18T12:56:51.423660Z","iopub.status.idle":"2023-12-18T12:56:51.471265Z","shell.execute_reply.started":"2023-12-18T12:56:51.423636Z","shell.execute_reply":"2023-12-18T12:56:51.470381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"object_columns = df.select_dtypes(include=['object']).columns\nprint(\"Object type columns:\")\nprint(object_columns)\n\nnumerical_columns = df.select_dtypes(include=['int64', 'float64']).columns\nprint(\"\\nNumerical type columns:\")\nprint(numerical_columns)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.472480Z","iopub.execute_input":"2023-12-18T12:56:51.472826Z","iopub.status.idle":"2023-12-18T12:56:51.483191Z","shell.execute_reply.started":"2023-12-18T12:56:51.472790Z","shell.execute_reply":"2023-12-18T12:56:51.482325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def classify_features(df):\n    categorical_features = []\n    non_categorical_features = []\n    discrete_features = []\n    continuous_features = []\n\n    for column in df.columns:\n        if df[column].dtype == 'object':\n            if df[column].nunique() < 10:\n                categorical_features.append(column)\n            else:\n                non_categorical_features.append(column)\n        elif df[column].dtype in ['int64', 'float64']:\n            if df[column].nunique() < 10:\n                discrete_features.append(column)\n            else:\n                continuous_features.append(column)\n\n    return categorical_features, non_categorical_features, discrete_features, continuous_features","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.484322Z","iopub.execute_input":"2023-12-18T12:56:51.484591Z","iopub.status.idle":"2023-12-18T12:56:51.496213Z","shell.execute_reply.started":"2023-12-18T12:56:51.484568Z","shell.execute_reply":"2023-12-18T12:56:51.495358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"categorical, non_categorical, discrete, continuous = classify_features(df)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.497616Z","iopub.execute_input":"2023-12-18T12:56:51.497963Z","iopub.status.idle":"2023-12-18T12:56:51.558678Z","shell.execute_reply.started":"2023-12-18T12:56:51.497931Z","shell.execute_reply":"2023-12-18T12:56:51.557688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Categorical Features:\", categorical)\nprint(\"Non-Categorical Features:\", non_categorical)\nprint(\"Discrete Features:\", discrete)\nprint(\"Continuous Features:\", continuous)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.559955Z","iopub.execute_input":"2023-12-18T12:56:51.560288Z","iopub.status.idle":"2023-12-18T12:56:51.565780Z","shell.execute_reply.started":"2023-12-18T12:56:51.560252Z","shell.execute_reply":"2023-12-18T12:56:51.564648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = df.fillna(\"Not Available\")","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.567063Z","iopub.execute_input":"2023-12-18T12:56:51.567422Z","iopub.status.idle":"2023-12-18T12:56:51.612407Z","shell.execute_reply.started":"2023-12-18T12:56:51.567386Z","shell.execute_reply":"2023-12-18T12:56:51.611488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in categorical:\n    print(i,':', df[i].unique())\n    print()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.614065Z","iopub.execute_input":"2023-12-18T12:56:51.614424Z","iopub.status.idle":"2023-12-18T12:56:51.626047Z","shell.execute_reply.started":"2023-12-18T12:56:51.614391Z","shell.execute_reply":"2023-12-18T12:56:51.625161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in categorical:\n    print(df[i].value_counts())\n    print()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.627511Z","iopub.execute_input":"2023-12-18T12:56:51.627834Z","iopub.status.idle":"2023-12-18T12:56:51.648060Z","shell.execute_reply.started":"2023-12-18T12:56:51.627803Z","shell.execute_reply":"2023-12-18T12:56:51.647161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:51.649167Z","iopub.execute_input":"2023-12-18T12:56:51.649435Z","iopub.status.idle":"2023-12-18T12:56:52.224563Z","shell.execute_reply.started":"2023-12-18T12:56:51.649412Z","shell.execute_reply":"2023-12-18T12:56:52.223717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:52.225661Z","iopub.execute_input":"2023-12-18T12:56:52.225987Z","iopub.status.idle":"2023-12-18T12:56:52.230948Z","shell.execute_reply.started":"2023-12-18T12:56:52.225956Z","shell.execute_reply":"2023-12-18T12:56:52.229972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"real_df = df[df[\"label\"] == \"REAL\"]\nfake_df = df[df[\"label\"] == \"FAKE\"]\nsample_size = 10000\n\nreal_df = real_df.sample(sample_size, random_state=42)\nfake_df = fake_df.sample(sample_size, random_state=42)\n\nsample_meta = pd.concat([real_df, fake_df])","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:52.232080Z","iopub.execute_input":"2023-12-18T12:56:52.232435Z","iopub.status.idle":"2023-12-18T12:56:52.298444Z","shell.execute_reply.started":"2023-12-18T12:56:52.232396Z","shell.execute_reply":"2023-12-18T12:56:52.297491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\nTrain_set, Test_set = train_test_split(sample_meta,test_size=0.2,random_state=42,stratify=sample_meta['label'])\nTrain_set, Val_set  = train_test_split(Train_set,test_size=0.3,random_state=42,stratify=Train_set['label'])","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:52.299672Z","iopub.execute_input":"2023-12-18T12:56:52.299967Z","iopub.status.idle":"2023-12-18T12:56:52.709245Z","shell.execute_reply.started":"2023-12-18T12:56:52.299944Z","shell.execute_reply":"2023-12-18T12:56:52.708330Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Train_set.shape,Val_set.shape,Test_set.shape","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:52.710515Z","iopub.execute_input":"2023-12-18T12:56:52.710855Z","iopub.status.idle":"2023-12-18T12:56:52.718394Z","shell.execute_reply.started":"2023-12-18T12:56:52.710826Z","shell.execute_reply":"2023-12-18T12:56:52.717480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:52.719608Z","iopub.execute_input":"2023-12-18T12:56:52.719888Z","iopub.status.idle":"2023-12-18T12:56:53.116382Z","shell.execute_reply.started":"2023-12-18T12:56:52.719863Z","shell.execute_reply":"2023-12-18T12:56:53.115450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\n\ndef retrieve_dataset(set_name):\n    dataset = []\n    for (img, imclass) in zip(set_name['videoname'], set_name['label']):\n        # Remove file extension from 'img'\n        img = img.split('.')[0]\n\n        # Assuming 'img' contains the filename without extension\n        image_path = f\"/kaggle/input/deepfake-faces/faces_224/{img}.jpg\"\n        try:\n            image = Image.open(image_path)\n        except FileNotFoundError:\n            print(f\"Image not found: {image_path}\")\n            continue\n\n        # Assuming binary classification (FAKE: 1, REAL: 0)\n        label = 1 if imclass == 'FAKE' else 0\n        \n        # Create the dictionary for each sample\n        sample = {\n            'image': image,\n            'image_file_path': image_path,\n            'labels': label\n        }\n\n        # Append the sample to the dataset\n        dataset.append(sample)\n\n    return dataset\n\nds = retrieve_dataset(Train_set)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:56:53.117561Z","iopub.execute_input":"2023-12-18T12:56:53.117870Z","iopub.status.idle":"2023-12-18T12:57:48.782903Z","shell.execute_reply.started":"2023-12-18T12:56:53.117844Z","shell.execute_reply":"2023-12-18T12:57:48.782083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import ViTImageProcessor, ViTForImageClassification\nimport requests","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:57:48.783944Z","iopub.execute_input":"2023-12-18T12:57:48.784242Z","iopub.status.idle":"2023-12-18T12:58:12.220060Z","shell.execute_reply.started":"2023-12-18T12:57:48.784217Z","shell.execute_reply":"2023-12-18T12:58:12.219074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"def collate_fn(batch):\n    images = [item['pixel_values'] for item in batch]\n    labels = [item['labels'] for item in batch]\n\n    inputs = {\n        'pixel_values': torch.stack(images),\n        'labels': torch.tensor(labels)\n    }\n\n    return inputs","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:58:12.221148Z","iopub.execute_input":"2023-12-18T12:58:12.221697Z","iopub.status.idle":"2023-12-18T12:58:12.226944Z","shell.execute_reply.started":"2023-12-18T12:58:12.221670Z","shell.execute_reply":"2023-12-18T12:58:12.225902Z"}}},{"cell_type":"code","source":"def collate_fn(batch):\n    # Convert the list of images and labels to NumPy arrays\n    images = [item['image'] for item in batch]\n    labels = [item['labels'] for item in batch]\n\n    # Print the shape of each image\n    for image in images:\n        print(image.shape)\n\n    # Convert the list of images and labels to NumPy arrays\n    pixel_values = torch.stack([item['pixel_values'] for item in batch])\n    return {\n        'pixel_values': pixel_values,\n        'labels': torch.tensor(labels)\n    }","metadata":{"execution":{"iopub.status.busy":"2023-12-18T13:06:46.338917Z","iopub.execute_input":"2023-12-18T13:06:46.339559Z","iopub.status.idle":"2023-12-18T13:06:46.345534Z","shell.execute_reply.started":"2023-12-18T13:06:46.339525Z","shell.execute_reply":"2023-12-18T13:06:46.344539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_name_or_path = 'google/vit-base-patch16-224-in21k'\nprocessor = ViTImageProcessor.from_pretrained(model_name_or_path)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:58:12.233973Z","iopub.execute_input":"2023-12-18T12:58:12.234899Z","iopub.status.idle":"2023-12-18T12:58:12.450350Z","shell.execute_reply.started":"2023-12-18T12:58:12.234845Z","shell.execute_reply":"2023-12-18T12:58:12.449338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process_example(example):\n    inputs = processor(example['image'], return_tensors='pt')\n    inputs['labels'] = example['labels']\n    return inputs\n\n# processed_dataset = [process_example(example) for example in ds]","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:58:12.451978Z","iopub.execute_input":"2023-12-18T12:58:12.452381Z","iopub.status.idle":"2023-12-18T12:58:12.458291Z","shell.execute_reply.started":"2023-12-18T12:58:12.452345Z","shell.execute_reply":"2023-12-18T12:58:12.456711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 32  \nnum_batches = len(ds) // batch_size\n\nprocessed_dataset = []\n\nfor i in range(num_batches):\n    batch_examples = ds[i * batch_size: (i + 1) * batch_size]\n    batch_processed = [process_example(example) for example in batch_examples]\n    processed_dataset.extend(batch_processed)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T12:58:12.459787Z","iopub.execute_input":"2023-12-18T12:58:12.460121Z","iopub.status.idle":"2023-12-18T12:59:53.824006Z","shell.execute_reply.started":"2023-12-18T12:58:12.460095Z","shell.execute_reply":"2023-12-18T12:59:53.822943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch","metadata":{"execution":{"iopub.status.busy":"2023-12-18T13:01:53.254917Z","iopub.execute_input":"2023-12-18T13:01:53.255869Z","iopub.status.idle":"2023-12-18T13:01:53.260187Z","shell.execute_reply.started":"2023-12-18T13:01:53.255833Z","shell.execute_reply":"2023-12-18T13:01:53.259099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader = torch.utils.data.DataLoader(processed_dataset, batch_size=8, shuffle=True, collate_fn=collate_fn)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T13:01:56.847735Z","iopub.execute_input":"2023-12-18T13:01:56.848513Z","iopub.status.idle":"2023-12-18T13:01:56.853383Z","shell.execute_reply.started":"2023-12-18T13:01:56.848480Z","shell.execute_reply":"2023-12-18T13:01:56.852459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from datasets import load_metric","metadata":{"execution":{"iopub.status.busy":"2023-12-18T13:02:08.378807Z","iopub.execute_input":"2023-12-18T13:02:08.379203Z","iopub.status.idle":"2023-12-18T13:02:09.228108Z","shell.execute_reply.started":"2023-12-18T13:02:08.379171Z","shell.execute_reply":"2023-12-18T13:02:09.227092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metric = load_metric(\"accuracy\")\ndef compute_metrics(p):\n    return metric.compute(predictions=np.argmax(p.predictions, axis=1), references=p.label_ids)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T13:02:11.140675Z","iopub.execute_input":"2023-12-18T13:02:11.141944Z","iopub.status.idle":"2023-12-18T13:02:11.549882Z","shell.execute_reply.started":"2023-12-18T13:02:11.141907Z","shell.execute_reply":"2023-12-18T13:02:11.548819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = [0, 1]  \nmodel = ViTForImageClassification.from_pretrained(\n    model_name_or_path,\n    num_labels=len(labels),\n    id2label={str(i): str(i) for i in labels},\n    label2id={str(i): i for i in labels}\n)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T13:02:14.370015Z","iopub.execute_input":"2023-12-18T13:02:14.370938Z","iopub.status.idle":"2023-12-18T13:02:17.294332Z","shell.execute_reply.started":"2023-12-18T13:02:14.370885Z","shell.execute_reply":"2023-12-18T13:02:17.293442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import TrainingArguments","metadata":{"execution":{"iopub.status.busy":"2023-12-18T13:03:35.241206Z","iopub.execute_input":"2023-12-18T13:03:35.241902Z","iopub.status.idle":"2023-12-18T13:03:35.269378Z","shell.execute_reply.started":"2023-12-18T13:03:35.241869Z","shell.execute_reply":"2023-12-18T13:03:35.268354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_args = TrainingArguments(\n    output_dir=\"/kaggle/working/\",\n    per_device_train_batch_size=8,\n    evaluation_strategy=\"steps\",\n    num_train_epochs=1,\n    fp16=True,\n    save_steps=100,\n    eval_steps=100,\n    logging_steps=10,\n    learning_rate=2e-4,\n    save_total_limit=2,\n    remove_unused_columns=False,\n    push_to_hub=False,\n    report_to='tensorboard',\n    load_best_model_at_end=True,\n)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T13:03:37.399804Z","iopub.execute_input":"2023-12-18T13:03:37.400554Z","iopub.status.idle":"2023-12-18T13:03:37.488833Z","shell.execute_reply.started":"2023-12-18T13:03:37.400520Z","shell.execute_reply":"2023-12-18T13:03:37.487696Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader = torch.utils.data.DataLoader(processed_dataset, batch_size=16, shuffle=True, collate_fn=collate_fn)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T13:03:42.376740Z","iopub.execute_input":"2023-12-18T13:03:42.377489Z","iopub.status.idle":"2023-12-18T13:03:42.382013Z","shell.execute_reply.started":"2023-12-18T13:03:42.377455Z","shell.execute_reply":"2023-12-18T13:03:42.381188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import Trainer","metadata":{"execution":{"iopub.status.busy":"2023-12-18T13:05:19.538478Z","iopub.execute_input":"2023-12-18T13:05:19.538859Z","iopub.status.idle":"2023-12-18T13:05:19.615565Z","shell.execute_reply.started":"2023-12-18T13:05:19.538833Z","shell.execute_reply":"2023-12-18T13:05:19.614648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = Trainer(\n    model=model,\n    args=training_args,\n    data_collator=collate_fn,\n    compute_metrics=compute_metrics,\n    train_dataset=train_dataloader.dataset,)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T13:05:22.301057Z","iopub.execute_input":"2023-12-18T13:05:22.301433Z","iopub.status.idle":"2023-12-18T13:05:29.754906Z","shell.execute_reply.started":"2023-12-18T13:05:22.301402Z","shell.execute_reply":"2023-12-18T13:05:29.754033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"train_results = trainer.train()\ntrainer.save_model()\ntrainer.log_metrics(\"train\", train_results.metrics)\ntrainer.save_metrics(\"train\", train_results.metrics)\ntrainer.save_state()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T13:07:07.466454Z","iopub.execute_input":"2023-12-18T13:07:07.467529Z","iopub.status.idle":"2023-12-18T13:07:08.127595Z","shell.execute_reply.started":"2023-12-18T13:07:07.467483Z","shell.execute_reply":"2023-12-18T13:07:08.126314Z"}}}]}