{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":31254,"databundleVersionId":3103714,"sourceType":"competition"}],"dockerImageVersionId":30776,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install peft","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-10-02T00:35:36.894204Z","iopub.execute_input":"2024-10-02T00:35:36.894488Z","iopub.status.idle":"2024-10-02T00:35:50.441949Z","shell.execute_reply.started":"2024-10-02T00:35:36.89445Z","shell.execute_reply":"2024-10-02T00:35:50.440875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport random\nfrom collections import defaultdict\nfrom collections import Counter\nfrom matplotlib import pyplot as plt\nfrom tqdm.notebook import tqdm\nimport os\nfrom pathlib import Path\nfrom PIL import Image\nimport requests\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom datasets import load_dataset, Dataset\nimport itertools\nfrom sklearn.model_selection import StratifiedShuffleSplit,train_test_split\nfrom sklearn.metrics import confusion_matrix\nfrom peft import LoraConfig, get_peft_model\nfrom transformers import (\n    DataCollator,\n    CLIPProcessor, \n    CLIPModel, \n    TrainingArguments, \n    Trainer\n)","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:35:50.444305Z","iopub.execute_input":"2024-10-02T00:35:50.444631Z","iopub.status.idle":"2024-10-02T00:36:10.861721Z","shell.execute_reply.started":"2024-10-02T00:35:50.444591Z","shell.execute_reply":"2024-10-02T00:36:10.860796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def hf_clip_predict(model, processor, text_labels, images):\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    model.to(device)\n    text = [f\"A photo of a {label}\" for label in text_labels]\n    inputs = processor(text=text, images=images, return_tensors=\"pt\", padding=True).to(device)\n    \n    model.eval()\n    with torch.no_grad():\n        outputs = model(**inputs)\n    \n    logits_per_image = outputs.logits_per_image\n    probs = logits_per_image.softmax(dim=1)\n    return probs\n\ndef get_image_paths_and_labels_from_df(df, data_dir):\n    article_ids = df[\"article_id\"].values\n    image_paths = []\n    labels = []\n    \n    for article_id in article_ids:\n        image_path = f\"{data_dir}/images/0{str(article_id)[:2]}/0{article_id}.jpg\"\n        # Check if the image file exists\n        if os.path.exists(image_path):\n            image_paths.append(image_path)\n            # Add corresponding label only if the image exists\n            labels.append(df[df[\"article_id\"] == article_id])\n        else:\n            print(f\"Image not found for article_id: {article_id}\")\n    \n    return image_paths, labels\n\ndef get_image_paths_and_labels_ordered(df, data_dir):\n    article_ids = df[\"article_id\"].values\n    image_paths = []\n    labels = []\n    for article_id in article_ids:\n        image_path = f\"{data_dir}/images/0{str(article_id)[:2]}/0{article_id}.jpg\"\n        if os.path.exists(image_path):\n            image_paths.append(image_path)\n            labels.append(df[df[\"article_id\"] == article_id])\n    \n    return image_paths, labels\n\ndef get_image_paths_and_labels(df, data_dir):\n    image_paths = []\n    labels = []\n    for root, dirs, files in os.walk(data_dir):\n        for file in files:\n            if file.endswith(\".jpg\"):\n                image_path = os.path.join(root, file)\n                image_paths.append(image_path)\n                article_id = int(file.split(\".\")[0])\n                labels.append(df[df[\"article_id\"] == article_id])\n\n    return image_paths, labels\n\nclass ImageDataset(torch.utils.data.Dataset):\n    def __init__(self, image_paths, processor=None):\n        self.image_paths = image_paths\n        self.processor = processor\n        self.image_ids = []\n\n        for image_path in self.image_paths:\n            if not os.path.exists(image_path):\n                raise FileNotFoundError(f\"Image {image_path} not found.\")\n            else:\n                image_id = int(image_path.split(\"/\")[-1].split(\".\")[0])\n                self.image_ids.append(image_id)\n            \n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        image = Image.open(self.image_paths[idx])\n        if self.processor is not None:\n            inputs = self.processor(images=image, return_tensors=\"pt\", padding=True)\n            image = inputs[\"pixel_values\"][0]\n        return image, self.image_ids[idx]","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:36:10.863197Z","iopub.execute_input":"2024-10-02T00:36:10.86379Z","iopub.status.idle":"2024-10-02T00:36:10.882067Z","shell.execute_reply.started":"2024-10-02T00:36:10.863753Z","shell.execute_reply":"2024-10-02T00:36:10.880835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nmodel = CLIPModel.from_pretrained(\"openai/clip-vit-base-patch32\")\nprocessor = CLIPProcessor.from_pretrained(\"openai/clip-vit-base-patch32\")\nmodel = model.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:36:10.884122Z","iopub.execute_input":"2024-10-02T00:36:10.884449Z","iopub.status.idle":"2024-10-02T00:36:19.232722Z","shell.execute_reply.started":"2024-10-02T00:36:10.884396Z","shell.execute_reply":"2024-10-02T00:36:19.231869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"text_path = '/kaggle/input/h-and-m-personalized-fashion-recommendations/articles.csv'\narticles = pd.read_csv(text_path)\nprint(articles.shape) # 100k data points\narticles.head(1)","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:36:19.233987Z","iopub.execute_input":"2024-10-02T00:36:19.234324Z","iopub.status.idle":"2024-10-02T00:36:20.270599Z","shell.execute_reply.started":"2024-10-02T00:36:19.234289Z","shell.execute_reply":"2024-10-02T00:36:20.269646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# map from article_id to df index\narticle_id_to_idx = {article_id: idx for idx, article_id in enumerate(articles[\"article_id\"])}\n\n# get all classes of the dataframe\nclass_names = articles.columns.tolist()\nlabel_names = dict()\nlabel_names_to_idx = dict()\nfor class_name in class_names:\n    label_names[class_name] = articles[class_name].unique()\n    label_names_to_idx[class_name] = {label_name: idx for idx, label_name in enumerate(label_names[class_name])}\n\narticle_ids = label_names[\"article_id\"]\nselected_class_names = [\"product_group_name\", \"product_type_name\", \"graphical_appearance_name\", \"colour_group_name\", \"perceived_colour_value_name\", \"perceived_colour_master_name\", \"department_name\", \"index_name\", \"index_group_name\", \"section_name\", \"garment_group_name\"]","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:36:20.271883Z","iopub.execute_input":"2024-10-02T00:36:20.272267Z","iopub.status.idle":"2024-10-02T00:36:20.520638Z","shell.execute_reply.started":"2024-10-02T00:36:20.272227Z","shell.execute_reply":"2024-10-02T00:36:20.519519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# get label names in product group name with less than 10 samples\nproduct_group_name_cnts = articles[\"product_group_name\"].value_counts()\nremoved_label_names = product_group_name_cnts[product_group_name_cnts < 10]\n\n# remove data with the removed label name\nremoved_label_idxs = articles[articles[\"product_group_name\"].isin(removed_label_names.index)].index\narticles = articles.drop(removed_label_idxs)","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:36:20.521831Z","iopub.execute_input":"2024-10-02T00:36:20.522132Z","iopub.status.idle":"2024-10-02T00:36:20.573256Z","shell.execute_reply.started":"2024-10-02T00:36:20.522099Z","shell.execute_reply":"2024-10-02T00:36:20.5725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = \"/kaggle/input/h-and-m-personalized-fashion-recommendations\"\nimage_paths, labels = get_image_paths_and_labels_from_df(articles, data_dir)\nprint(f\"Number of images: {len(image_paths)}\")","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:36:20.574434Z","iopub.execute_input":"2024-10-02T00:36:20.574819Z","iopub.status.idle":"2024-10-02T00:43:11.873978Z","shell.execute_reply.started":"2024-10-02T00:36:20.574757Z","shell.execute_reply":"2024-10-02T00:43:11.873036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"现在的数据集是按照prod_name直接来随机分的，训练：验证 = 8:2， 测试集为200个。同一个prod_name应该只会在三个数据集中的一个，满足fredrik的要求（待验证？）","metadata":{}},{"cell_type":"code","source":"prod_name_to_images = defaultdict(list)\nfor idx, label_df in enumerate(labels):\n    prod_name = label_df['prod_name'].values[0]  # Extract scalar value\n    prod_name_to_images[prod_name].append(idx)  # Store the index of the image\n\nall_prod_names = list(prod_name_to_images.keys())\nrandom.seed(42)  # For reproducibility\nrandom.shuffle(all_prod_names)\n\n#Allocate 'prod_name's to the test set until we have at least 200 images\ntest_prod_names = []\ntest_image_indices = []\ntotal_test_images = 0\nfor prod_name in all_prod_names:\n    indices = prod_name_to_images[prod_name]\n    if total_test_images + len(indices) > 200:\n        # Adjust to get exactly 200 images\n        remaining_slots = 200 - total_test_images\n        test_image_indices.extend(indices[:remaining_slots])\n        total_test_images += remaining_slots\n        break\n    else:\n        test_prod_names.append(prod_name)\n        test_image_indices.extend(indices)\n        total_test_images += len(indices)\n    if total_test_images == 200:\n        break\n\n# Remove selected 'prod_name's from the list of all 'prod_name's\nremaining_prod_names = [pn for pn in all_prod_names if pn not in test_prod_names]\n\n# 8/2 split\nnum_train = int(0.8 * len(remaining_prod_names))\ntrain_prod_names = remaining_prod_names[:num_train]\nval_prod_names = remaining_prod_names[num_train:]\n\ntrain_image_indices = []\nfor prod_name in train_prod_names:\n    train_image_indices.extend(prod_name_to_images[prod_name])\n\nval_image_indices = []\nfor prod_name in val_prod_names:\n    val_image_indices.extend(prod_name_to_images[prod_name])\n\ntrain_image_paths = [image_paths[idx] for idx in train_image_indices]\ntrain_labels = [labels[idx] for idx in train_image_indices]\n\nval_image_paths = [image_paths[idx] for idx in val_image_indices]\nval_labels = [labels[idx] for idx in val_image_indices]\n\ntest_image_paths = [image_paths[idx] for idx in test_image_indices]\ntest_labels = [labels[idx] for idx in test_image_indices]\n\n# #256 for test\n# train_image_paths = train_image_paths[:256]\n# train_labels = train_labels[:256]\n\n# val_image_paths = val_image_paths[:256]\n# val_labels = val_labels[:256]\n\nprint(f\"Number of training images: {len(train_image_paths)}\")\nprint(f\"Number of validation images: {len(val_image_paths)}\")\nprint(f\"Number of test images: {len(test_image_paths)}\")","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:43:11.875102Z","iopub.execute_input":"2024-10-02T00:43:11.875418Z","iopub.status.idle":"2024-10-02T00:43:17.76594Z","shell.execute_reply.started":"2024-10-02T00:43:11.875383Z","shell.execute_reply":"2024-10-02T00:43:17.764863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = ImageDataset(train_image_paths, processor)\ntrain_dataloader = torch.utils.data.DataLoader(train_dataset, batch_size=256, shuffle=True)\n\nval_dataset = ImageDataset(val_image_paths, processor)\nval_dataloader = torch.utils.data.DataLoader(val_dataset, batch_size=256, shuffle=False)\n\ntest_dataset = ImageDataset(test_image_paths, processor)\ntest_dataloader = torch.utils.data.DataLoader(test_dataset, batch_size=256, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:43:17.770068Z","iopub.execute_input":"2024-10-02T00:43:17.770468Z","iopub.status.idle":"2024-10-02T00:44:04.033341Z","shell.execute_reply.started":"2024-10-02T00:43:17.770434Z","shell.execute_reply":"2024-10-02T00:44:04.03253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define LoRA configuration\nlora_config = LoraConfig(\n    r=8,                  # Low-rank dimension (adjustable)\n    lora_alpha=32,          # Scaling factor (adjustable)\n    target_modules=[\"q_proj\", \"v_proj\", \"k_proj\"],  # Specify which layers to apply LoRA to\n    lora_dropout=0.05,       # Dropout rate (optional)\n    bias=\"none\",            # Whether to include biases (\"none\", \"all\", \"lora_only\")\n    task_type=\"classification\"  # Task type (\"classification\" or \"regression\")\n)\n\n# Apply LoRA to the CLIP model\nmodel = get_peft_model(model, lora_config)","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:44:04.034552Z","iopub.execute_input":"2024-10-02T00:44:04.034879Z","iopub.status.idle":"2024-10-02T00:44:04.185767Z","shell.execute_reply.started":"2024-10-02T00:44:04.034845Z","shell.execute_reply":"2024-10-02T00:44:04.184836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"定义confusion matrix画几个类","metadata":{}},{"cell_type":"code","source":"N = 5  # Number of top classes to include in the confusion matrix","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:44:04.187158Z","iopub.execute_input":"2024-10-02T00:44:04.187457Z","iopub.status.idle":"2024-10-02T00:44:04.191316Z","shell.execute_reply.started":"2024-10-02T00:44:04.187424Z","shell.execute_reply":"2024-10-02T00:44:04.190395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_confusion_matrix(cm, class_labels, title='Confusion Matrix'):\n    plt.figure(figsize=(10, 8))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n                xticklabels=class_labels, yticklabels=class_labels)\n    plt.xlabel('Predicted Labels')\n    plt.ylabel('True Labels')\n    plt.title(title)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:44:04.192674Z","iopub.execute_input":"2024-10-02T00:44:04.193312Z","iopub.status.idle":"2024-10-02T00:44:04.205109Z","shell.execute_reply.started":"2024-10-02T00:44:04.193268Z","shell.execute_reply":"2024-10-02T00:44:04.204251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def validate(model, val_dataloader, criteria, device, tokenized_texts, selected_class_names, epoch):\n    model.eval()\n    total_loss = 0.0\n    total_correct = {class_name: 0 for class_name in selected_class_names}\n    total_samples = 0\n\n    # Initialize containers for confusion matrices\n    all_true_labels = {class_name: [] for class_name in selected_class_names}\n    all_preds = {class_name: [] for class_name in selected_class_names}\n\n    with torch.no_grad():\n        for images, image_ids in tqdm(val_dataloader):\n            images = images.to(device)\n            batch_size = images.size(0)\n\n            # Get true labels for all classes\n            true_labels = {}\n            for class_name in selected_class_names:\n                labels = [\n                    label_names_to_idx[class_name][\n                        articles.loc[article_id_to_idx[image_id.item()], class_name]\n                    ] for image_id in image_ids\n                ]\n                true_labels[class_name] = torch.tensor(labels).to(device)\n\n            # Compute image embeddings\n            image_features = model.get_image_features(images)\n            image_features = image_features / image_features.norm(dim=-1, keepdim=True)\n\n            loss = 0.0  # Reset loss for the batch\n\n            # Iterate over each class\n            for class_name in selected_class_names:\n                # Move tokenized text inputs to device\n                inputs = tokenized_texts[class_name]\n                inputs = {k: v.to(device) for k, v in inputs.items()}\n\n                # Compute text embeddings\n                text_features = model.get_text_features(**inputs)\n                text_features = text_features / text_features.norm(dim=-1, keepdim=True)\n\n                # Compute similarity logits\n                logits_per_image = image_features @ text_features.T  # Shape: [batch_size, num_labels]\n\n                # Compute loss for this class\n                class_loss = criteria(logits_per_image, true_labels[class_name])\n                loss += class_loss  # Sum losses from all classes\n\n                # Predictions and accuracy\n                _, preds = torch.max(logits_per_image, dim=1)\n                total_correct[class_name] += (preds == true_labels[class_name]).sum().item()\n\n                # Collect true labels and predictions for confusion matrix\n                all_true_labels[class_name].extend(true_labels[class_name].cpu().numpy())\n                all_preds[class_name].extend(preds.cpu().numpy())\n\n            total_loss += loss.item() * batch_size\n            total_samples += batch_size\n\n    avg_loss = total_loss / total_samples\n    accuracy = {class_name: total_correct[class_name] / total_samples for class_name in selected_class_names}\n\n    # Compute and display confusion matrices for validation\n    for class_name in selected_class_names:\n        true_labels_np = np.array(all_true_labels[class_name])\n        preds_np = np.array(all_preds[class_name])\n\n        # Compute the frequency of each class in true labels\n        label_counts = Counter(true_labels_np)\n        # Get the top N classes\n        top_N_classes = [label for label, _ in label_counts.most_common(N)]\n        top_N_set = set(top_N_classes)\n\n        # Create a mapping from old labels to new indices\n        label_to_new_index = {label: idx for idx, label in enumerate(top_N_classes)}\n        other_label_index = N  # Index for 'Other' category\n\n        # Remap labels\n        remapped_true_labels = np.array([\n            label_to_new_index.get(label, other_label_index) for label in true_labels_np\n        ])\n        remapped_preds = np.array([\n            label_to_new_index.get(label, other_label_index) for label in preds_np\n        ])\n\n        # Update class labels for the confusion matrix\n        class_labels_for_cm = [label_names[class_name][label] for label in top_N_classes] + ['Other']\n\n        # Compute confusion matrix\n        cm = confusion_matrix(\n            remapped_true_labels,\n            remapped_preds,\n            labels=list(range(N + 1))\n        )\n\n        # Plot confusion matrix\n        plot_confusion_matrix(\n            cm,\n            class_labels_for_cm,\n            title=f'Validation Confusion Matrix for {class_name} (Epoch {epoch+1})'\n        )\n\n    return avg_loss, accuracy","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:44:04.206212Z","iopub.execute_input":"2024-10-02T00:44:04.206518Z","iopub.status.idle":"2024-10-02T00:44:04.225474Z","shell.execute_reply.started":"2024-10-02T00:44:04.206486Z","shell.execute_reply":"2024-10-02T00:44:04.224525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs = 2  # Adjust as needed\ncriteria = torch.nn.CrossEntropyLoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-5)\n\nmodel.to(device)\n\n# Tokenize text inputs for all labels in all classes outside the batch loop\ntokenized_texts = {}\nfor class_name in selected_class_names:\n    labels = label_names[class_name]\n    texts = [f\"A photo of a {label}\" for label in labels]\n    tokenized_texts[class_name] = processor(\n        text=texts,\n        return_tensors=\"pt\",\n        padding=True\n    )\n\nfor epoch in range(num_epochs):\n    model.train()\n    total_loss = 0.0\n    total_correct = {class_name: 0 for class_name in selected_class_names}\n    total_samples = 0\n    \n    # Initialize containers for confusion matrices\n    all_true_labels = {class_name: [] for class_name in selected_class_names}\n    all_preds = {class_name: [] for class_name in selected_class_names}\n\n    for images, image_ids in tqdm(train_dataloader):\n        images = images.to(device)\n        batch_size = images.size(0)\n\n        # Get true labels for all classes\n        true_labels = {}\n        for class_name in selected_class_names:\n            labels = [\n                label_names_to_idx[class_name][\n                    articles.loc[article_id_to_idx[image_id.item()], class_name]\n                ] for image_id in image_ids\n            ]\n            true_labels[class_name] = torch.tensor(labels).to(device)\n\n        # Forward pass: compute image embeddings\n        image_features = model.get_image_features(images)\n        image_features = image_features / image_features.norm(dim=-1, keepdim=True)\n\n        loss = 0.0  # Reset loss for the batch\n\n        # Iterate over each class\n        for class_name in selected_class_names:\n            # Move tokenized text inputs to device\n            inputs = tokenized_texts[class_name]\n            inputs = {k: v.to(device) for k, v in inputs.items()}\n\n            # Compute text embeddings\n            text_features = model.get_text_features(**inputs)\n            text_features = text_features / text_features.norm(dim=-1, keepdim=True)\n\n            # Compute similarity logits\n            logits_per_image = image_features @ text_features.T  # Shape: [batch_size, num_labels]\n\n            # Compute loss for this class\n            class_loss = criteria(logits_per_image, true_labels[class_name])\n            loss += class_loss  # Sum losses from all classes\n\n            # Predictions and accuracy\n            _, preds = torch.max(logits_per_image, dim=1)\n            total_correct[class_name] += (preds == true_labels[class_name]).sum().item()\n            \n            # Collect true labels and predictions for confusion matrix\n            all_true_labels[class_name].extend(true_labels[class_name].cpu().numpy())\n            all_preds[class_name].extend(preds.cpu().numpy())\n\n        total_loss += loss.item() * batch_size\n        total_samples += batch_size\n\n        # Backward and optimize\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n    # Compute average loss and accuracy\n    avg_loss = total_loss / total_samples\n    accuracy = {class_name: total_correct[class_name] / total_samples for class_name in selected_class_names}\n    print(f\"Epoch [{epoch+1}/{num_epochs}], Loss: {avg_loss:.4f}\")\n    for class_name in selected_class_names:\n        print(f\"Accuracy for {class_name}: {accuracy[class_name]:.4f}\")\n    \n    # Compute and display confusion matrices for training\n    for class_name in selected_class_names:\n        true_labels_np = np.array(all_true_labels[class_name])\n        preds_np = np.array(all_preds[class_name])\n\n        # Compute the frequency of each class in true labels\n        label_counts = Counter(true_labels_np)\n        # Get the top N classes\n        top_N_classes = [label for label, _ in label_counts.most_common(N)]\n        top_N_set = set(top_N_classes)\n\n        # Create a mapping from old labels to new indices\n        label_to_new_index = {label: idx for idx, label in enumerate(top_N_classes)}\n        other_label_index = N  # Index for 'Other' category\n\n        # Remap labels\n        remapped_true_labels = np.array([\n            label_to_new_index.get(label, other_label_index) for label in true_labels_np\n        ])\n        remapped_preds = np.array([\n            label_to_new_index.get(label, other_label_index) for label in preds_np\n        ])\n\n        # Update class labels for the confusion matrix\n        class_labels_for_cm = [label_names[class_name][label] for label in top_N_classes] + ['Other']\n\n        # Compute confusion matrix\n        cm = confusion_matrix(\n            remapped_true_labels,\n            remapped_preds,\n            labels=list(range(N + 1))\n        )\n\n        # Plot confusion matrix\n        plot_confusion_matrix(\n            cm,\n            class_labels_for_cm,\n            title=f'Training Confusion Matrix for {class_name} (Epoch {epoch+1})'\n        )\n\n    # Validate after each epoch\n    val_loss, val_accuracy = validate(model, val_dataloader, criteria, device, tokenized_texts, selected_class_names, epoch)\n    print(f\"Validation Loss: {val_loss:.4f}\")\n    for class_name in selected_class_names:\n        print(f\"Validation Accuracy for {class_name}: {val_accuracy[class_name]:.4f}\")","metadata":{"execution":{"iopub.status.busy":"2024-10-02T00:44:04.227486Z","iopub.execute_input":"2024-10-02T00:44:04.227827Z","iopub.status.idle":"2024-10-02T03:45:35.537096Z","shell.execute_reply.started":"2024-10-02T00:44:04.227771Z","shell.execute_reply":"2024-10-02T03:45:35.535207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_save_path = 'model.pth'\ntorch.save(model.state_dict(), model_save_path)\nprint(f\"Model saved to {model_save_path}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-10-02T03:45:35.539765Z","iopub.execute_input":"2024-10-02T03:45:35.540193Z","iopub.status.idle":"2024-10-02T03:45:36.561335Z","shell.execute_reply.started":"2024-10-02T03:45:35.540155Z","shell.execute_reply":"2024-10-02T03:45:36.560369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_loss, test_accuracy = validate(model, test_dataloader, criteria, device, text_inputs, selected_class_names, 0)\nprint(f\"Test Loss: {test_loss:.4f}\")\nfor class_name in selected_class_names:\n    print(f\"Test Accuracy for {class_name}: {test_accuracy[class_name]:.4f}\")","metadata":{"execution":{"iopub.status.busy":"2024-10-02T03:45:36.562513Z","iopub.execute_input":"2024-10-02T03:45:36.562877Z","iopub.status.idle":"2024-10-02T03:45:36.926629Z","shell.execute_reply.started":"2024-10-02T03:45:36.562834Z","shell.execute_reply":"2024-10-02T03:45:36.925103Z"},"trusted":true},"execution_count":null,"outputs":[]}]}