mirror of
https://github.com/0xTriboulet/T-1
synced 2026-06-06 15:14:27 +00:00
Initial DTC implementation and visual extraction
This commit is contained in:
Binary file not shown.
|
After Width: | Height: | Size: 814 KiB |
@@ -1,8 +1,15 @@
|
||||
Process Count,Process Count/User,User Count,Sandbox Score
|
||||
33,8,4,1
|
||||
34,8.5,4,1
|
||||
157,157,1,0
|
||||
158,158,1,0
|
||||
30,7.5,4,1
|
||||
31,7.75,4,1
|
||||
84,84,1,0
|
||||
85,85,1,0
|
||||
195,195,1,0
|
||||
34,9,4,1
|
||||
150,75,2,1
|
||||
196,196,1,0
|
||||
34,8.5,4,1
|
||||
35,8.75,4,1
|
||||
150,75,2,1
|
||||
151,75.5,2,1
|
||||
|
@@ -1,13 +1,64 @@
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
from sklearn.model_selection import RandomizedSearchCV
|
||||
from sklearn.tree import DecisionTreeClassifier
|
||||
from sklearn.tree import plot_tree
|
||||
|
||||
# Define the file path
|
||||
CSV_FILE = './dataset/process_data.csv'
|
||||
|
||||
MIN_DATASET_SIZE = 50
|
||||
|
||||
|
||||
# BadDecisionTree class contains core functionality
|
||||
class BadDecisionTree:
|
||||
def __init__(self, df: pd.DataFrame):
|
||||
self.data: pd.DataFrame = df
|
||||
self.best_model = None
|
||||
self.model: DecisionTreeClassifier = DecisionTreeClassifier()
|
||||
self.hyperparams: dict = {
|
||||
'max_depth': np.arange(1, 10).astype(np.uint64),
|
||||
'min_samples_leaf': np.arange(1, 50).astype(np.uint64),
|
||||
'min_samples_split': np.arange(2, 10).astype(np.uint64),
|
||||
'criterion': ["gini", "entropy"],
|
||||
'splitter': ["best", "random"],
|
||||
}
|
||||
|
||||
self.data = df
|
||||
|
||||
if self.data.shape[0] < MIN_DATASET_SIZE:
|
||||
self._data_augment()
|
||||
|
||||
self.X = self.data.drop('Sandbox Score', axis=1)
|
||||
self.y = self.data['Sandbox Score']
|
||||
|
||||
def _data_augment(self) -> None:
|
||||
self.data = self.data.sample(n=MIN_DATASET_SIZE, replace=True)
|
||||
|
||||
def tune_hyperparameters(self) -> None:
|
||||
# Number of iteration for RandomizedSearchCV
|
||||
n_iter_search = 20
|
||||
# Setting up the Randomized Search with cross validation
|
||||
self.best_model = RandomizedSearchCV(self.model, param_distributions=self.hyperparams,
|
||||
n_iter=n_iter_search, cv=5)
|
||||
|
||||
self.best_model.fit(self.X, self.y)
|
||||
return self.best_model.best_params_
|
||||
|
||||
def visualize_tree(self, filename: str):
|
||||
# Check if best_model is not None and if it's been fit
|
||||
if self.best_model and hasattr(self.best_model, 'best_estimator_'):
|
||||
# Initialize the figure size
|
||||
plt.figure(figsize=(100, 60), dpi=100) # increase figure size and dpi
|
||||
# Use plot_tree from sklearn.tree to visualize the tree
|
||||
plot_tree(self.best_model.best_estimator_,
|
||||
feature_names=self.X.columns,
|
||||
class_names=True,
|
||||
filled=True)
|
||||
# Save the tree
|
||||
plt.savefig(filename)
|
||||
else:
|
||||
print("The model has not been created or fitted yet")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
@@ -19,3 +70,22 @@ if __name__ == "__main__":
|
||||
|
||||
# Print the contents of the DataFrame to the console
|
||||
print(BadTree.data)
|
||||
print(BadTree.X)
|
||||
print(BadTree.y)
|
||||
|
||||
# Print hyperparameters
|
||||
print("\n\nTraining BadTree:")
|
||||
print(BadTree.tune_hyperparameters())
|
||||
|
||||
# Define the data and the column names
|
||||
data = [[140, 70, 2]] # The data inside two brackets makes it a 2D list which represents one row and three columns
|
||||
columns = ['Process Count', 'Process Count/User', 'User Count']
|
||||
|
||||
test_frame = pd.DataFrame(data, columns=columns)
|
||||
|
||||
# Make a prediction on our novel test frame
|
||||
print(f'BadTree test frame:\n {test_frame}')
|
||||
print(f'BadTree prediction frame: {BadTree.best_model.predict(test_frame)}')
|
||||
|
||||
# Visualize and save the decision tree
|
||||
BadTree.visualize_tree("badtree.png")
|
||||
|
||||
Reference in New Issue
Block a user