{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":4117,"databundleVersionId":46665,"sourceType":"competition"},{"sourceId":4727131,"sourceType":"datasetVersion","datasetId":2734949}],"dockerImageVersionId":30301,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Credit to Dheemanth Bhat for the source code using the XGBoost algorithm: https://www.kaggle.com/code/dheemanthbhat/malware-classification-with-multiprocessing\n\nassume I have not edited code up until section 9 unless a comment is marked with the indicator B.C.","metadata":{}},{"cell_type":"markdown","source":"## 1. Overview","metadata":{}},{"cell_type":"markdown","source":"Dataset Source: https://arxiv.org/abs/1802.10135\n\n\n\n### 1.1 Problem\n\n\n\nWe are provided with a set of known malware files in _train.7z_ and _test.7z_ zip files. Each set representing a mix of Nine different malware families. Class labels of these malware files are provided in the separate _trainLabels.csv_ CSV file.\n\nEach malware file in the train and test set has:\n\n* **_Id_**: A 20 character hash value uniquely identifying the file as its file-name. Refer zip files.\n* **_Class_**: An integer representing one of 9 family names to which the malware may belong. Refer _trainLabels.csv_ CSV file.\n\nTypes of Malwares:\n1. Ramnit\n2. Lollipop\n3. Kelihos_ver3\n4. Vundo\n5. Simda\n6. Tracur\n7. Kelihos_ver1\n8. Obfuscator.ACY\n9. Gatak\n\nFor each file, the raw data contains the **hexadecimal representation** of the file's binary content, **without the PE header (to ensure sterility)**.  You are also provided a metadata manifest (files with _.asm_ extension), which is a log containing various metadata information extracted from the binary, such as function calls, strings, etc. This was generated using the IDA disassembler tool. Our task is to develop the best mechanism for classifying files in the test set into their respective family affiliations.","metadata":{}},{"cell_type":"markdown","source":"### 1.2 Solution\n\n1. Randomly equal number of malware files from each class (except Simda).\n2. Extract hexadecimal strings (of length 2) from the _.bytes_ file.\n3. Build Bag-of-Words using scikit-learn `CountVectorizer`.\n4. Build a model using XGBoost `XGBClassifier`.\n\n> **Note:** \n> 1. Part of this solution is based on the case-study by my beloved mentor [Csk Verma sir][1].\n> 2. We can improve the solution by using the metadata manifest (files with .asm extension).\n> 3. We can further improve the solution by more using advanced featurization techniques.\n\n[1]: https://www.kaggle.com/srikanthvarmachekuri","metadata":{}},{"cell_type":"markdown","source":"### 1.3 Objective\n\nObjective of this kernel is to:\n\n1. Pick a **balanced sample** from the original dataset.\n2. Since the dataset is insanely huge, **improve the solution by parallelizing** computationally intensive tasks and file operations.","metadata":{}},{"cell_type":"markdown","source":"### 1.4 Project Folder Structure\n\nFinal project structure:\n```\n├── code\n│    ├── Microsoft Malware Classification (with Multiprocessing).ipynb\n│    ├── HexVectorizer.py\n│    └── submission.csv\n└── input\n     ├── malware-classification\n     │    ├── trainLabels.csv\n     │    ├── train\n     │    │    ├── ...\n     │    │    ├── 0ACDbR5M3ZhBJajygTuf.bytes        (10,868 byte-file)\n     │    │    ├── ...\n     │    └── test\n     │         ├── ...\n     │         ├── ITSUPtCmh7WdJcsYDwQ5.bytes        (10,873 byte-file)\n     │         ├── ...\n     └── microsoft-malware-sample\n          ├── trainLabels_bal.csv\n          ├── train_vec.csv\n          └── test_vec.csv\n```","metadata":{}},{"cell_type":"markdown","source":"## 2. Setup","metadata":{}},{"cell_type":"markdown","source":"### 2.1 Add custom module\n\nRef: [Custom Packages in Kernels](https://www.kaggle.com/getting-started/58516#364577)","metadata":{}},{"cell_type":"markdown","source":"### 2.2 Import and configure libraries","metadata":{}},{"cell_type":"code","source":"%%writefile HexVectorizer.py\n\nfrom sklearn.feature_extraction.text import CountVectorizer\n\nimport os\nimport re\n\n# Byte file properties\nLINE_LEN = 16\nADDR_LEN = 9\n\n\ndef hex_to_str(hex_line):\n    \"\"\"\n    Function to strip \\r, \\n and remove address\n    of length eight character+space from hex line.\n    \"\"\"\n    return hex_line.decode(errors=\"ignore\").strip()[ADDR_LEN:]\n        # B.C. - small change ignores errors that may be caused be non decodable \n        # characters (not likely necessary since the dataset is self contained,\n        # however it is good practice if the models intended to be used)\n\n\ndef remove_non_hex(hex_str):\n    \"\"\"\n    Function to remove non hex characters.\n    \"\"\"\n    return re.sub(r\"[^0-9A-F]+\", \" \", hex_str, flags=re.IGNORECASE).strip().lower()\n        # B.C. - reduced function to be contained on a single line, sice it \n        # contained unnecessary variables, saving memory\n\n\n\nclass Vectorizer(CountVectorizer):\n    \"\"\"\n    Convert strings to vectors\n    \"\"\"\n\n    def transform_byte_file(self, ts_file):\n        \"\"\"\n        Function to convert a byte-file into a data-point/row in CSV file.\n        Each data-point will have file-name, file-size & byte-string columns.\n        \"\"\"\n\n        # Open the byte-file for reading.\n        with open(ts_file, \"rb\") as byt_f:\n            # Remove memory address in the beginning of each line and concatenate\n            # all lines in the byte-file into a single string separated by space.\n            byt_str = \" \".join(hex_to_str(line) for line in byt_f)\n                # B.C. - previous function would consume a lot of memory so instead process \n                # line by line\n            byt_str = remove_non_hex(byt_str)\n\n            # Get the byte-file name.\n            f_path, _ = os.path.splitext(byt_f.name)  # Full path, extension.\n            _, f_name = f_path.rsplit(\"/\", 1)  # Relative path, file-name.\n\n            # Get the byte-file size.\n            file_info = os.stat(byt_f.name)\n            f_size = file_info.st_size\n\n            bow = self.transform([byt_str]).toarray()[0].tolist()\n            return [f_name, f_size] + bow\n\n    def get_feature_names(self):\n        \"\"\"\n        Function to return vocabulary as feature names.\n        \"\"\"\n        return self.get_feature_names_out().tolist()","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.021322Z","iopub.execute_input":"2024-12-06T15:14:14.022685Z","iopub.status.idle":"2024-12-06T15:14:14.032504Z","shell.execute_reply.started":"2024-12-06T15:14:14.022632Z","shell.execute_reply":"2024-12-06T15:14:14.031091Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Data manipulation libraries.\nimport numpy as np\nimport pandas as pd\n\n# Data visualization libraries.\nimport matplotlib\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom prettytable import PrettyTable\n\n# Data modeling libraries.\nimport sklearn\nfrom sklearn.feature_extraction.text import CountVectorizer\nfrom sklearn.preprocessing import LabelEncoder, StandardScaler, MinMaxScaler\nfrom sklearn.manifold import TSNE\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.utils.class_weight import compute_sample_weight\nfrom xgboost import XGBClassifier\nfrom sklearn.metrics import log_loss\nfrom sklearn.model_selection import GridSearchCV\n\n# General Imports\nimport os\nimport csv\nimport re\nimport time\nimport math\nfrom tqdm import tqdm\nimport multiprocessing\nfrom multiprocessing import Pool\n\n# Custom modules\nfrom HexVectorizer import Vectorizer\n\n\n# Library versions.\nprint(\"NumPy version:\", np.__version__)\nprint(\"Pandas version:\", pd.__version__)\nprint(\"Matplotlib version:\", matplotlib.__version__)\nprint(\"Seaborn version:\", sns.__version__)\nprint(\"Scikit-learn version:\", sklearn.__version__)\n\n# Configure NumPy.\n# Set `Line width` to Maximum 130 characters in the output, post which it will continue in next line.\nnp.set_printoptions(linewidth=130)\n\n# Configure Pandas.\n# Set display width to maximum 130 characters in the output, post which it will continue in next line.\npd.options.display.width = 130\n# pd.options.display.max_rows = None  # Very dangerous! if dataset is large.\n\n# Configure Seaborn.\nsns.set_style(\"whitegrid\")  # Set white background with grid.\nsns.set_palette(\"deep\")  # Set color palette.\nsns.set_context(\"paper\", font_scale=1.5)  # Set font to scale 1.5 more than normal.","metadata":{"_kg_hide-input":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.035528Z","iopub.execute_input":"2024-12-06T15:14:14.036036Z","iopub.status.idle":"2024-12-06T15:14:14.054229Z","shell.execute_reply.started":"2024-12-06T15:14:14.035992Z","shell.execute_reply":"2024-12-06T15:14:14.052847Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2.3 Constants and helper functions","metadata":{}},{"cell_type":"code","source":"TRAIN_DIR = \"../input/malware-classification/train/\"\nLABELS_CSV = \"../input/malware-classification/trainLabels.csv\"\nTEST_DIR = \"../input/malware-classification/test/\"\nPREPS_DIR = \"../input/microsoft-malware-sample/\"\nBYTE_EXT = \".bytes\"\n\nBAL_CL_CSV = os.path.join(PREPS_DIR, \"trainLabels_bal.csv\")\nTR_VEC_CSV = os.path.join(PREPS_DIR, \"train_vec.csv\")\nTS_VEC_CSV = os.path.join(PREPS_DIR, \"test_vec.csv\")\n\n# Byte file properties\nLINE_LEN = 16\nADDR_LEN = 9\n\n\n# Malware class-labels.\nMALWARE_CLS = {\n    1: \"Ramnit\",\n    2: \"Lollipop\",\n    3: \"Kelihos_ver3\",\n    4: \"Vundo\",\n    5: \"Simda\",\n    6: \"Tracur\",\n    7: \"Kelihos_ver1\",\n    8: \"Obfuscator.ACY\",\n    9: \"Gatak\",\n}\nmw_codes = lambda: list(MALWARE_CLS.keys())\nmw_names = lambda: list(MALWARE_CLS.values())\n\n\nSAMPLE_SIZE = 200  # Number of samples taken from each class.\n# Function to pick random sample of size SAMPLE_SIZE from each class.\nget_sample = lambda grp: grp.sample(min(SAMPLE_SIZE, len(grp)))\n\n\ndef get_weights(cls):\n    class_weights = {\n        0: 1,\n        1: 1,\n        2: 1,\n        3: 1,\n        4: 4,\n        5: 1,\n        6: 1,\n        7: 1,\n        8: 1,\n    }\n\n    return [class_weights[cl] for cl in cls]","metadata":{"_kg_hide-input":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.055703Z","iopub.execute_input":"2024-12-06T15:14:14.056022Z","iopub.status.idle":"2024-12-06T15:14:14.069974Z","shell.execute_reply.started":"2024-12-06T15:14:14.055994Z","shell.execute_reply":"2024-12-06T15:14:14.068521Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2.4 Multiprocessing","metadata":{}},{"cell_type":"code","source":"# Multiprocessing on Jupyter Notebook on windows:\n# https://stackoverflow.com/a/47374811/5070460\ndef parallelize(task, data):\n    \"\"\"\n    Function to parallelize `task()` for list of items passed as `data`.\n    \"\"\"\n    # Confirm that the code is under main function.\n    if __name__ == \"__main__\":\n        # Set pool size to number of logical-processors available.\n        POOL_SIZE = multiprocessing.cpu_count()\n        pool = Pool(processes=POOL_SIZE)\n        print(\"Pool size:\", POOL_SIZE)\n\n        outputs = []\n        pbar = tqdm(total=len(data))\n        for output in pool.imap_unordered(task, data):\n            outputs.append(output)\n            pbar.update()\n        pbar.close()\n\n        return outputs","metadata":{"_kg_hide-input":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.071779Z","iopub.execute_input":"2024-12-06T15:14:14.072206Z","iopub.status.idle":"2024-12-06T15:14:14.088285Z","shell.execute_reply.started":"2024-12-06T15:14:14.072156Z","shell.execute_reply":"2024-12-06T15:14:14.086988Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. EDA","metadata":{}},{"cell_type":"markdown","source":"### 3.1 Byte-file class-labels: _trainLabels.csv_","metadata":{}},{"cell_type":"code","source":"# Class-labels DataFrame.\ncl_df = pd.read_csv(LABELS_CSV)\ncl_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.091603Z","iopub.execute_input":"2024-12-06T15:14:14.092065Z","iopub.status.idle":"2024-12-06T15:14:14.142902Z","shell.execute_reply.started":"2024-12-06T15:14:14.092023Z","shell.execute_reply":"2024-12-06T15:14:14.141481Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"rows, cols = cl_df.shape\nprint(f\"There are around {rows} malware files available in the Train dataset.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.144427Z","iopub.execute_input":"2024-12-06T15:14:14.144867Z","iopub.status.idle":"2024-12-06T15:14:14.152171Z","shell.execute_reply.started":"2024-12-06T15:14:14.144832Z","shell.execute_reply":"2024-12-06T15:14:14.150762Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cls_count = cl_df[\"Class\"].value_counts().sort_index()\ncls_count","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.153993Z","iopub.execute_input":"2024-12-06T15:14:14.154465Z","iopub.status.idle":"2024-12-06T15:14:14.178414Z","shell.execute_reply.started":"2024-12-06T15:14:14.154401Z","shell.execute_reply":"2024-12-06T15:14:14.177162Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Observations**\n\n1. Dataset is **highly imbalanced**.\n2. **Only 42 byte-files** are available of malware class **5** (_Simda_).","metadata":{}},{"cell_type":"markdown","source":"### 3.2 Balanced subset from Original dataset","metadata":{}},{"cell_type":"code","source":"if os.path.exists(BAL_CL_CSV) and os.stat(BAL_CL_CSV).st_size > 0:\n    cl_df_b = pd.read_csv(BAL_CL_CSV)\nelse:\n    cl_df_b = cl_df.groupby(\"Class\", group_keys=False).apply(get_sample)\n    # Save DataFrame as CSV for future use.\n    cl_df_b.to_csv(BAL_CL_CSV, index=False)\n\nrows, cols = cl_df_b.shape\nprint(f\"Balanced subset will contain data of {rows} byte-files.\")\n\ncls_count = cl_df_b[\"Class\"].value_counts().sort_index()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.181118Z","iopub.execute_input":"2024-12-06T15:14:14.181526Z","iopub.status.idle":"2024-12-06T15:14:14.203249Z","shell.execute_reply.started":"2024-12-06T15:14:14.181476Z","shell.execute_reply":"2024-12-06T15:14:14.201819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(6, 6))\n\nplt.pie(x=cls_count, labels=mw_names(), autopct=\"%1.0f%%\")\nplt.title(\"Malware Class-labels\")\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.204882Z","iopub.execute_input":"2024-12-06T15:14:14.205219Z","iopub.status.idle":"2024-12-06T15:14:14.435131Z","shell.execute_reply.started":"2024-12-06T15:14:14.205188Z","shell.execute_reply":"2024-12-06T15:14:14.433654Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Data wrangling","metadata":{}},{"cell_type":"markdown","source":"### 4.1 The Problem\n\nData is provided as _.bytes_ and _.asm_ files. We have to convert this raw data to features/vectors before building a classification model. For example a _.bytes_ file contain data as show below. Sample taken from _0ACDbR5M3ZhBJajygTuf.bytes_ file:\n```\n...\n0046F4B0 4E 01 82 01 00 01 AE 01 0E 11 E2 11 0A 10 2A 01\n0046F4C0 60 01 C8 00 E8 10 28 01 04 00 82 00 62 10 E4 01\n0046F4D0 EA 00 CE 01 A6 01 46 11 0C 00 00 00 ?? ?? ?? ??\n0046F4E0 ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ??\n0046F4F0 ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ??\n0046F500 ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ?? ??\n...\n```\n\nWhere,\n\n1. First eight characters are **memory address**.\n2. A hexadecimal representation of the command (byte), always in 2 characters.\n3. Missing values with `??` placeholder.","metadata":{}},{"cell_type":"markdown","source":"> **Note:**  \n> The double question mark (??) is a placeholder that is used in a byte file to represent a single byte of information that is unknown or unspecified. This placeholder is often used in situations where the exact value of the byte is not important or relevant, such as when a byte file is being used as a template or when a byte file is being created using a program that automatically fills in the placeholder values with actual data. The double question mark is a common convention used in byte files, but it is not a language-specific construct, so its exact meaning and usage may vary depending on the context in which it is used.","metadata":{}},{"cell_type":"markdown","source":"### 4.2 Introducing `HexVectorizer`\n\nBuild a custom Vectorizer called `HexVectorizer` to:\n\n1. Remove memory-address and `??` placeholder.\n2. Add file-size as a feature.\n3. Replace string with its BoW vector representation.\n\n#### Vocabulary\n\nSince we already know that byte-files contain only hexadecimal characters and hex characters range b/w `0-9` followed by `A-F`, we can build the vocabulary with 256 (16X16) unique words in it.","metadata":{"heading_collapsed":true}},{"cell_type":"code","source":"hex_chars = [\"0\", \"1\", \"2\", \"3\", \"4\", \"5\", \"6\", \"7\", \"8\", \"9\", \"a\", \"b\", \"c\", \"d\", \"e\", \"f\"]\n\nvectorizer = Vectorizer()\nvectorizer.vocabulary = [f\"{i}{j}\" for i in hex_chars for j in hex_chars]","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.437017Z","iopub.execute_input":"2024-12-06T15:14:14.437692Z","iopub.status.idle":"2024-12-06T15:14:14.451137Z","shell.execute_reply.started":"2024-12-06T15:14:14.437632Z","shell.execute_reply":"2024-12-06T15:14:14.449314Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 4.3 Parallelized File Processing\n\nImport `HexVectorizer` module and pass it to `multiprocessing` `Pool` class for processing multiple files in parallel.","metadata":{}},{"cell_type":"code","source":"final_ftrs = [\"Id\", \"BytFSize\"] + vectorizer.get_feature_names()\nfile_paths = [os.path.join(TRAIN_DIR, f\"{Id}{BYTE_EXT}\") for Id in cl_df_b[\"Id\"]]\n    # B.C. - as the dataset is large we need to optimise small things like string \n    # concatenation, here we optimise by using list comprehension more efficiently\n\nfile_count = len(file_paths)\n\nprint(f\"Select {file_count} byte-files from unzipped files for processing.\")\n\nif os.path.exists(TR_VEC_CSV) and os.stat(TR_VEC_CSV).st_size > 0:\n    print(\"Sample is already vectorized in:\", TR_VEC_CSV)\n    train_df = pd.read_csv(TR_VEC_CSV)\nelse:\n    print(\"Vectorization of balanced sample from original dataset...\")\n\n    # Parallelized File Processing.\n    tr_ftr_vecs = parallelize(vectorizer.transform_byte_file, file_paths)\n\n    # Convert vectors to DataFrame.\n    train_df = pd.DataFrame(tr_ftr_vecs, columns=final_ftrs)\n\n    def get_class(Id):\n        fltr = cl_df_b[\"Id\"] == Id\n        return cl_df_b.loc[fltr, \"Class\"].item()\n\n    # Append class labels.\n    train_df[\"Class\"] = train_df[\"Id\"].apply(get_class)\n\n    # Save DataFrame as CSV for future use.\n    train_df.to_csv(TR_VEC_CSV, index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.452964Z","iopub.execute_input":"2024-12-06T15:14:14.453616Z","iopub.status.idle":"2024-12-06T15:14:14.554456Z","shell.execute_reply.started":"2024-12-06T15:14:14.453567Z","shell.execute_reply":"2024-12-06T15:14:14.553142Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Feature Engineering\n\nThere is nothing much left to engineer any features since non-hex characters and unwanted data are already removed during data wrangling.","metadata":{}},{"cell_type":"code","source":"train_df.sample(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.555937Z","iopub.execute_input":"2024-12-06T15:14:14.556238Z","iopub.status.idle":"2024-12-06T15:14:14.579804Z","shell.execute_reply.started":"2024-12-06T15:14:14.556210Z","shell.execute_reply":"2024-12-06T15:14:14.578481Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df.info()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.581968Z","iopub.execute_input":"2024-12-06T15:14:14.582573Z","iopub.status.idle":"2024-12-06T15:14:14.611143Z","shell.execute_reply.started":"2024-12-06T15:14:14.582522Z","shell.execute_reply":"2024-12-06T15:14:14.609732Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"rows, cols = train_df.shape\nprint(f\"Dataset contains {rows} rows and {cols} columns\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.615030Z","iopub.execute_input":"2024-12-06T15:14:14.615400Z","iopub.status.idle":"2024-12-06T15:14:14.627564Z","shell.execute_reply.started":"2024-12-06T15:14:14.615367Z","shell.execute_reply":"2024-12-06T15:14:14.625906Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 5.1 Validate class-labels","metadata":{}},{"cell_type":"code","source":"for uid, cl in train_df[[\"Id\", \"Class\"]].values:\n    fltr = (cl_df[\"Id\"] == uid) & (cl_df[\"Class\"] == cl)\n    f_id = cl_df.loc[fltr, \"Id\"]\n    if f_id.empty:\n        raise ValueError(f\"Id: {uid} with class: {cl} not found in cl_df.\")\n\nprint(\"All file-ids and class-labels are matching.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:14.629258Z","iopub.execute_input":"2024-12-06T15:14:14.629638Z","iopub.status.idle":"2024-12-06T15:14:16.838341Z","shell.execute_reply.started":"2024-12-06T15:14:14.629605Z","shell.execute_reply":"2024-12-06T15:14:16.837031Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 5.2 Feature Scaling","metadata":{}},{"cell_type":"code","source":"X = train_df.drop([\"Id\", \"Class\"], axis=1)\ny = train_df[\"Class\"]\n\n# Centering data-points.\nscaler = StandardScaler()\nX_scaled = scaler.fit_transform(X)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:16.840572Z","iopub.execute_input":"2024-12-06T15:14:16.840900Z","iopub.status.idle":"2024-12-06T15:14:16.876720Z","shell.execute_reply.started":"2024-12-06T15:14:16.840869Z","shell.execute_reply":"2024-12-06T15:14:16.875590Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 5.3 High dimensional data visualization","metadata":{}},{"cell_type":"markdown","source":"#### T-SNE","metadata":{}},{"cell_type":"code","source":"# import warnings\n# warnings.filterwarnings(\"ignore\")\n\ntsne = TSNE(\n    n_components=2,\n    perplexity=10,\n    init=\"random\",\n    learning_rate=\"auto\",\n    verbose=1,\n    random_state=42,\n    n_jobs=-1,\n)\nz = tsne.fit_transform(X_scaled)\nkld = np.round(tsne.kl_divergence_, 4)\n\nprint(\"KL-Divergence:\", kld)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:16.878412Z","iopub.execute_input":"2024-12-06T15:14:16.878907Z","iopub.status.idle":"2024-12-06T15:14:23.153217Z","shell.execute_reply.started":"2024-12-06T15:14:16.878865Z","shell.execute_reply":"2024-12-06T15:14:23.152312Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.DataFrame()\n\ndf[\"Class\"] = y.astype(\"category\")\ndf[\"Comp-1\"] = z[:, 0]\ndf[\"Comp-2\"] = z[:, 1]\n\nplt.figure(figsize=(12, 7))\n\nsns.scatterplot(data=df, x=\"Comp-1\", y=\"Comp-2\", hue=\"Class\")\nplt.title(f\"T-SNE projection. KL-Divergence score: {kld}\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:23.154282Z","iopub.execute_input":"2024-12-06T15:14:23.154617Z","iopub.status.idle":"2024-12-06T15:14:23.680490Z","shell.execute_reply.started":"2024-12-06T15:14:23.154584Z","shell.execute_reply":"2024-12-06T15:14:23.679194Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Model Training","metadata":{"heading_collapsed":true}},{"cell_type":"markdown","source":"Replace class-labels in the range [1-9] with range [0-8] because of the [class-label constraints XGBoost][1].\n\n[1]: https://stackoverflow.com/a/72132612/5070460","metadata":{"hidden":true}},{"cell_type":"code","source":"le = LabelEncoder()\ny = le.fit_transform(y)","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:23.682273Z","iopub.execute_input":"2024-12-06T15:14:23.683220Z","iopub.status.idle":"2024-12-06T15:14:23.689387Z","shell.execute_reply.started":"2024-12-06T15:14:23.683175Z","shell.execute_reply":"2024-12-06T15:14:23.688123Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 6.1 Train, Test split","metadata":{"hidden":true}},{"cell_type":"code","source":"X_train, X_cval, y_train, y_cval = train_test_split(\n    X_scaled,\n    y,\n    stratify=y,\n    shuffle=True,\n    test_size=0.3,\n    random_state=42,\n)\nscaler = StandardScaler()\nX_train = scaler.fit_transform(X_train)\nX_cval = scaler.transform(X_cval)\n    # B.C. - after splitting the data into training and testing portions, scale the\n    # data to prevent overfitting\n\n\nprint(\"Train dataset shape:\", X_train.shape)\nprint(\"Cross-Val dataset shape:\", X_cval.shape)","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:23.691266Z","iopub.execute_input":"2024-12-06T15:14:23.691762Z","iopub.status.idle":"2024-12-06T15:14:23.716845Z","shell.execute_reply.started":"2024-12-06T15:14:23.691716Z","shell.execute_reply":"2024-12-06T15:14:23.715294Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 6.2 Baseline Model","metadata":{"hidden":true}},{"cell_type":"code","source":"base_m = XGBClassifier(random_state=42, n_jobs=-1)\nbase_m.fit(X_train, y_train, sample_weight=get_weights(y_train))","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:23.718289Z","iopub.execute_input":"2024-12-06T15:14:23.718640Z","iopub.status.idle":"2024-12-06T15:14:29.606051Z","shell.execute_reply.started":"2024-12-06T15:14:23.718609Z","shell.execute_reply":"2024-12-06T15:14:29.604952Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 6.3 Accuracy","metadata":{"hidden":true}},{"cell_type":"code","source":"train_acc = base_m.score(X_train, y_train)\ntest_acc = base_m.score(X_cval, y_cval)\n\nprint(\"Train accuracy:\", round(train_acc, 4))\nprint(\"Test accuracy:\", round(test_acc, 4))","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:29.607561Z","iopub.execute_input":"2024-12-06T15:14:29.607984Z","iopub.status.idle":"2024-12-06T15:14:29.633325Z","shell.execute_reply.started":"2024-12-06T15:14:29.607934Z","shell.execute_reply":"2024-12-06T15:14:29.630375Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 6.4 Log Loss","metadata":{"hidden":true}},{"cell_type":"code","source":"y_probs = base_m.predict_proba(X_cval)\nll = log_loss(y_cval, y_probs)\n\nprint(\"Log loss:\", round(ll, 4))","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:29.634285Z","iopub.execute_input":"2024-12-06T15:14:29.634601Z","iopub.status.idle":"2024-12-06T15:14:29.655594Z","shell.execute_reply.started":"2024-12-06T15:14:29.634571Z","shell.execute_reply":"2024-12-06T15:14:29.654738Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Hyperparameter Tuning","metadata":{"heading_collapsed":true}},{"cell_type":"markdown","source":"### 7.1 Helper functions","metadata":{"hidden":true}},{"cell_type":"code","source":"def get_log_loss(kwargs):\n    xgb_clf = XGBClassifier(random_state=42, n_jobs=-1, **kwargs)\n    xgb_clf.fit(X_train, y_train)\n\n    y_train_probs = xgb_clf.predict_proba(X_train)\n    train_ll = log_loss(y_train, y_train_probs)\n\n    y_test_probs = xgb_clf.predict_proba(X_cval)\n    test_ll = log_loss(y_cval, y_test_probs)\n\n    return train_ll, test_ll\n\n\ndef hyperparameter_tuning(**kwargs):\n    train_errs = []\n    test_errs = []\n\n    params = list(kwargs.items())\n    hypr_params = dict(params[0:-1])\n    # Get the last item with range.\n    param_key, range_vals = params[-1]\n\n    for param_val in range_vals:\n        hypr_params[param_key] = param_val\n        train_err, test_err = get_log_loss(hypr_params)\n        train_errs.append(train_err)\n        test_errs.append(test_err)\n\n    plt.figure(figsize=(5, 4))\n\n    sns.lineplot(x=range_vals, y=train_errs, label=\"Train error\")\n    sns.lineplot(x=range_vals, y=test_errs, label=\"Test error\")\n    plt.title(f\"Param: `{param_key}`\")\n    plt.xlabel(\"Value\")\n    plt.ylabel(\"Log Loss\")\n    plt.xticks(range_vals)\n\n    plt.show()","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:29.659441Z","iopub.execute_input":"2024-12-06T15:14:29.659780Z","iopub.status.idle":"2024-12-06T15:14:29.669754Z","shell.execute_reply.started":"2024-12-06T15:14:29.659751Z","shell.execute_reply":"2024-12-06T15:14:29.668612Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 7.2 Tuning\n\nTo decide what range to use for above hyperparameters, plot Train vs Validation (or test) error and pick the sweet spot where:\n1. Test-error is least \n2. Gap b/w Train-error and Train-error is least.","metadata":{"hidden":true}},{"cell_type":"markdown","source":"#### 1. Param: `n_estimators`","metadata":{"hidden":true}},{"cell_type":"code","source":"hyperparameter_tuning(n_estimators=range(1, 70, 10))","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:29.671129Z","iopub.execute_input":"2024-12-06T15:14:29.671607Z","iopub.status.idle":"2024-12-06T15:14:53.101863Z","shell.execute_reply.started":"2024-12-06T15:14:29.671567Z","shell.execute_reply":"2024-12-06T15:14:53.100496Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### 2. Param: `max_depth`","metadata":{"hidden":true}},{"cell_type":"code","source":"hyperparameter_tuning(n_estimators=40, max_depth=range(1, 10))","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:14:53.104022Z","iopub.execute_input":"2024-12-06T15:14:53.104333Z","iopub.status.idle":"2024-12-06T15:15:26.351084Z","shell.execute_reply.started":"2024-12-06T15:14:53.104305Z","shell.execute_reply":"2024-12-06T15:15:26.349601Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### 3. Param: `learning_rate`","metadata":{"hidden":true}},{"cell_type":"code","source":"hyperparameter_tuning(max_depth=4, n_estimators=40, learning_rate=np.arange(0.1, 0.55, 0.05))","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:15:26.352871Z","iopub.execute_input":"2024-12-06T15:15:26.353223Z","iopub.status.idle":"2024-12-06T15:16:03.006890Z","shell.execute_reply.started":"2024-12-06T15:15:26.353194Z","shell.execute_reply":"2024-12-06T15:16:03.005357Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### 4. Param: `reg_alpha`","metadata":{"hidden":true}},{"cell_type":"code","source":"hyperparameter_tuning(\n    max_depth=4,\n    n_estimators=40,\n    learning_rate=0.3,\n    reg_alpha=np.arange(0.1, 1, 0.2),\n)","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:16:03.008831Z","iopub.execute_input":"2024-12-06T15:16:03.009323Z","iopub.status.idle":"2024-12-06T15:16:24.993717Z","shell.execute_reply.started":"2024-12-06T15:16:03.009277Z","shell.execute_reply":"2024-12-06T15:16:24.992511Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### 5. Param: `reg_lambda`","metadata":{"hidden":true}},{"cell_type":"code","source":"hyperparameter_tuning(\n    max_depth=3,\n    n_estimators=40,\n    learning_racate=0.3,\n    reg_lambda=np.arange(0.1, 1, 0.2),\n)","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:16:24.995122Z","iopub.execute_input":"2024-12-06T15:16:24.995454Z","iopub.status.idle":"2024-12-06T15:16:43.298495Z","shell.execute_reply.started":"2024-12-06T15:16:24.995399Z","shell.execute_reply":"2024-12-06T15:16:43.297069Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 7.3 Grid Search CV\n\nPerform Grid Search CV using XGBClassifier model to further fine tune parameters.","metadata":{"hidden":true}},{"cell_type":"code","source":"grid = {\n    \"n_estimators\": [40, 45, 50],\n    \"max_depth\": [3, 4, 5],\n    \"learning_rate\": [0.35, 0.4, 0.45],\n}\n\nstart = time.time()\n\nxgb_clf1 = XGBClassifier(random_state=42, n_jobs=-1)\nxgb_cvm = GridSearchCV(estimator=xgb_clf1, param_grid=grid, n_jobs=-1, cv=5)\nxgb_cvm.fit(X, y, sample_weight=get_weights(y))","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:16:43.300859Z","iopub.execute_input":"2024-12-06T15:16:43.301256Z","iopub.status.idle":"2024-12-06T15:24:45.921353Z","shell.execute_reply.started":"2024-12-06T15:16:43.301223Z","shell.execute_reply":"2024-12-06T15:24:45.919743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"end = time.time()\nprint(\"Time taken in hyperparameter tuning:\", math.ceil((end - start) / 60), \"mins.\")\n\ntuned_model = xgb_cvm.best_estimator_\n\nprint(\"\\nBest parameters:\")\nprint(xgb_cvm.best_params_)","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:24:45.923208Z","iopub.execute_input":"2024-12-06T15:24:45.923563Z","iopub.status.idle":"2024-12-06T15:24:45.930821Z","shell.execute_reply.started":"2024-12-06T15:24:45.923528Z","shell.execute_reply":"2024-12-06T15:24:45.929725Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Model Accuracy","metadata":{"hidden":true}},{"cell_type":"code","source":"print(\"Accuracy:\", round(xgb_cvm.best_score_, 4))","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:24:45.936391Z","iopub.execute_input":"2024-12-06T15:24:45.936794Z","iopub.status.idle":"2024-12-06T15:24:45.951319Z","shell.execute_reply.started":"2024-12-06T15:24:45.936765Z","shell.execute_reply":"2024-12-06T15:24:45.950020Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 7.4 Final Model","metadata":{"hidden":true}},{"cell_type":"code","source":"best_model = XGBClassifier(**xgb_cvm.best_params_, random_state=42, n_jobs=-1)\nbest_model.fit(X, y, sample_weight=get_weights(y))","metadata":{"hidden":true,"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:24:45.952903Z","iopub.execute_input":"2024-12-06T15:24:45.953231Z","iopub.status.idle":"2024-12-06T15:24:52.262838Z","shell.execute_reply.started":"2024-12-06T15:24:45.953204Z","shell.execute_reply":"2024-12-06T15:24:52.261892Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Model Testing","metadata":{}},{"cell_type":"markdown","source":"### 8.1 Feature Engineering\n\nLoad vectorized Test dataset from a CSV file if present otherwise generate feature vectors from byte file.","metadata":{}},{"cell_type":"code","source":"if os.path.exists(TS_VEC_CSV) and os.stat(TS_VEC_CSV).st_size > 0:\n    print(\"Loading vectors from\", TS_VEC_CSV)\n    test_df = pd.read_csv(TS_VEC_CSV, index_col=\"Id\")\nelse:\n    # Get complete list of byte-files in input directory.\n    _, _, ts_file_names = next(os.walk(TEST_DIR))\n\n    # Generate list for file paths.\n    ts_file_paths = [os.path.join(TEST_DIR, fn) for fn in ts_file_names]\n\n    print(f\"Test dataset contains {len(ts_file_paths)} files.\")\n    print(\"Vectorization of test dataset...\")\n\n    # Parallelized File Processing.\n    ts_ftr_vecs = parallelize(vectorizer.transform_byte_file, ts_file_paths)\n\n    # Convert vectors to DataFrame.\n    test_df = pd.DataFrame(ts_ftr_vecs, columns=final_ftrs)\n    test_df.set_index([\"Id\"], inplace=True)\n\n    # Save DataFrame as CSV for future use.\n    test_df.to_csv(TS_VEC_CSV, index_label=\"Id\")\n\ntest_df.sample(3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:24:52.264169Z","iopub.execute_input":"2024-12-06T15:24:52.264510Z","iopub.status.idle":"2024-12-06T15:24:52.694511Z","shell.execute_reply.started":"2024-12-06T15:24:52.264480Z","shell.execute_reply":"2024-12-06T15:24:52.693158Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"rows, cols = test_df.shape\nprint(f\"Test dataset contains {rows} rows and {cols} columns.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:24:52.695918Z","iopub.execute_input":"2024-12-06T15:24:52.696209Z","iopub.status.idle":"2024-12-06T15:24:52.703139Z","shell.execute_reply.started":"2024-12-06T15:24:52.696183Z","shell.execute_reply":"2024-12-06T15:24:52.701650Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 8.2 Predictions","metadata":{}},{"cell_type":"code","source":"predictions = np.round(best_model.predict_proba(test_df).tolist(), 1).tolist()\n\nprint(\"Sample rows from predictions:\")\npredictions[:5]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:24:52.704842Z","iopub.execute_input":"2024-12-06T15:24:52.705272Z","iopub.status.idle":"2024-12-06T15:24:52.809590Z","shell.execute_reply.started":"2024-12-06T15:24:52.705232Z","shell.execute_reply":"2024-12-06T15:24:52.808221Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data = np.column_stack((test_df.index, predictions))\ncol_names = [\"Id\"] + [f\"Prediction{i}\" for i in range(1, 10)]\n\noutput = pd.DataFrame(data, columns=col_names)\noutput.to_csv(\"submission.csv\", index=False)\nprint(\"Your submission was successfully saved!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:24:52.810912Z","iopub.execute_input":"2024-12-06T15:24:52.811205Z","iopub.status.idle":"2024-12-06T15:24:52.893915Z","shell.execute_reply.started":"2024-12-06T15:24:52.811181Z","shell.execute_reply":"2024-12-06T15:24:52.892510Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 9. CNN Expansion -B.C.","metadata":{}},{"cell_type":"markdown","source":"This section uses the same dataset, but using a CNN model in place of XGBoost, to investigate how the CNN model can interact with this form of Malware classification","metadata":{}},{"cell_type":"code","source":"# Import TensorFlow/Keras for CNN\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras.layers import Conv2D, MaxPooling2D, Flatten, Dense\nfrom tensorflow.keras.utils import to_categorical\nimport numpy as np\nfrom sklearn.metrics import confusion_matrix\n\n# Malware class names\nmalware_classes = [\"Ramnit\", \"Lollipop\", \"Kelihos_ver3\", \"Vundo\", \"Simda\", \"Tracur\", \"Kelihos_ver1\", \"Obfuscator.ACY\", \"Gatak\"]\n\n# Prepare the data for CNN\ncnn_X = train_df.drop([\"Id\", \"Class\"], axis=1).values\nNUM_FEATURES = cnn_X.shape[1]\n\n# Define approximate dimensions for reshaping\nIMG_HEIGHT = 16  # Set desired height (adjust based on your preference)\nIMG_WIDTH = (NUM_FEATURES + IMG_HEIGHT - 1) // IMG_HEIGHT  # Calculate width to fit all features\n\n# Pad or truncate feature data to match reshape dimensions\nNEW_NUM_FEATURES = IMG_HEIGHT * IMG_WIDTH\nif NUM_FEATURES < NEW_NUM_FEATURES:\n    # Pad with zeros if there are fewer features than needed\n    padding = np.zeros((cnn_X.shape[0], NEW_NUM_FEATURES - NUM_FEATURES))\n    cnn_X = np.hstack((cnn_X, padding))\n    print(\"Padding has occurred\")\n\nelif NUM_FEATURES > NEW_NUM_FEATURES:\n    # Truncate excess features if there are too many features\n    cnn_X = cnn_X[:, :NEW_NUM_FEATURES]\n    print(\"Truncation has occurred\")\n\n# Reshape data to 4D tensor (batch_size, IMG_HEIGHT, IMG_WIDTH, 1)\ncnn_X_reshaped = cnn_X.reshape(-1, IMG_HEIGHT, IMG_WIDTH, 1)  # Add channel dimension\ncnn_y = to_categorical(train_df[\"Class\"] - 1)  # Convert class labels to one-hot encoding\n\n# Split data into training and validation sets\ncnn_X_train, cnn_X_cval, cnn_y_train, cnn_y_cval = train_test_split(\n    cnn_X_reshaped,\n    cnn_y,\n    stratify=train_df[\"Class\"],\n    test_size=0.3,\n    random_state=42,\n)\n\n# Define a simple CNN model\ncnn_model = Sequential([\n    Conv2D(32, (3, 3), activation=\"relu\", input_shape=(IMG_HEIGHT, IMG_WIDTH, 1)),\n    MaxPooling2D(pool_size=(2, 2)),\n    Flatten(),\n    Dense(128, activation=\"relu\"),\n    Dense(cnn_y.shape[1], activation=\"softmax\"),\n])\n\n# Compile the model\ncnn_model.compile(optimizer=\"adam\", loss=\"categorical_crossentropy\", metrics=[\"accuracy\"])\n\n# Train the model\ncnn_model.fit(cnn_X_train, cnn_y_train, validation_data=(cnn_X_cval, cnn_y_cval), epochs=25, batch_size=32)\n\n# Evaluate the CNN\ncnn_loss, cnn_accuracy = cnn_model.evaluate(cnn_X_cval, cnn_y_cval)\nprint(f\"CNN Validation Loss: {cnn_loss:.4f}, Accuracy: {cnn_accuracy:.4f}\")\n\n# Predict using CNN\ncnn_predictions = cnn_model.predict(cnn_X_cval)\n\n# Convert predictions from one-hot encoding to class labels (i.e., the index with highest probability)\npredicted_classes = np.argmax(cnn_predictions, axis=1)\n\n# Convert true labels from one-hot encoding to class labels\ntrue_classes = np.argmax(cnn_y_cval, axis=1)\n\n# Calculate per-class accuracy\nconf_matrix = confusion_matrix(true_classes, predicted_classes)\n\n# Get accuracy per class\nper_class_accuracy = conf_matrix.diagonal() / conf_matrix.sum(axis=1)\n\n# Print per-class accuracy\nfor class_idx, accuracy in enumerate(per_class_accuracy):\n    print(f\"{malware_classes[class_idx]} Accuracy: {accuracy:.4f}\")\n\n\n# Sample CNN predictions\nprint(\"Sample CNN Predictions:\", cnn_predictions[:5])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T15:24:52.895631Z","iopub.execute_input":"2024-12-06T15:24:52.895993Z","iopub.status.idle":"2024-12-06T15:25:07.885328Z","shell.execute_reply.started":"2024-12-06T15:24:52.895961Z","shell.execute_reply":"2024-12-06T15:25:07.883988Z"}},"outputs":[],"execution_count":null}]}