Cây quyết định

Machine Learning với PySpark

Andrew Collier

Data Scientist, Fathom Data

Cấu trúc Cây quyết định: Nút gốc

Nút gốc của cây quyết định.

Machine Learning với PySpark

Cấu trúc Cây quyết định: Lần chia đầu

Cây quyết định với một lần chia

Machine Learning với PySpark

Cấu trúc Cây quyết định: Lần chia thứ hai

Cây quyết định với lần chia thứ hai

Machine Learning với PySpark

Cấu trúc Cây quyết định: Lần chia thứ ba

Cây quyết định với lần chia thứ ba

Machine Learning với PySpark

Phân loại xe

Phân loại xe theo quốc gia sản xuất.

+---+----+------+------+----+-----------+----------------------------------+-----+
|cyl|size|mass  |length|rpm |consumption|features                          |label|
+---+----+------+------+----+-----------+----------------------------------+-----+
|6  |3.0 |1451.0|4.775 |5200|9.05       |[6.0,3.0,1451.0,4.775,5200.0,9.05]|1.0  |
|4  |2.2 |1129.0|4.623 |5200|6.53       |[4.0,2.2,1129.0,4.623,5200.0,6.53]|0.0  |
|4  |2.2 |1399.0|4.547 |5600|7.84       |[4.0,2.2,1399.0,4.547,5600.0,7.84]|1.0  |
|4  |1.8 |1147.0|4.343 |6500|7.84       |[4.0,1.8,1147.0,4.343,6500.0,7.84]|0.0  |
|4  |1.6 |1111.0|4.216 |5750|9.05       |[4.0,1.6,1111.0,4.216,5750.0,9.05]|0.0  |
+---+----+------+------+----+-----------+----------------------------------+-----+

label = 0 -> sản xuất tại Hoa Kỳ
      = 1 -> sản xuất ở nơi khác
Machine Learning với PySpark

Chia train/test

Chia dữ liệu thành tập huấn luyện và kiểm tra.

# Specify a seed for reproducibility
cars_train, cars_test = cars.randomSplit([0.8, 0.2], seed=23)

Hai DataFrame: cars_traincars_test.

[cars_train.count(), cars_test.count()]
[79, 13]
Machine Learning với PySpark

Xây dựng mô hình Cây quyết định

from pyspark.ml.classification import DecisionTreeClassifier

Tạo bộ phân loại Cây quyết định.

tree = DecisionTreeClassifier()

Học từ dữ liệu huấn luyện.

tree_model = tree.fit(cars_train)
Machine Learning với PySpark

Đánh giá

Dự đoán trên dữ liệu kiểm tra và so sánh với nhãn đúng.

prediction = tree_model.transform(cars_test)
+-----+----------+---------------------------------------+
|label|prediction|probability                            |
+-----+----------+---------------------------------------+
|1.0  |0.0       |[0.9615384615384616,0.0384615384615385]|
|1.0  |1.0       |[0.2222222222222222,0.7777777777777778]|
|1.0  |1.0       |[0.2222222222222222,0.7777777777777778]|
|0.0  |0.0       |[0.9615384615384616,0.0384615384615385]|
|1.0  |1.0       |[0.2222222222222222,0.7777777777777778]|
+-----+----------+---------------------------------------+
Machine Learning với PySpark

Ma trận nhầm lẫn

Ma trận nhầm lẫn là bảng mô tả hiệu năng mô hình trên dữ liệu kiểm tra.

prediction.groupBy("label", "prediction").count().show()
+-----+----------+-----+
|label|prediction|count|
+-----+----------+-----+
|  1.0|       1.0|    8| <- Dương tính thật (TP)
|  0.0|       1.0|    2| <- Dương tính giả  (FP)
|  1.0|       0.0|    3| <- Âm tính giả     (FN)
|  0.0|       0.0|    6| <- Âm tính thật    (TN)
+-----+----------+-----+

Độ chính xác = (TN + TP) / (TN + TP + FN + FP) — tỉ lệ dự đoán đúng.

Machine Learning với PySpark

Hãy xây mô hình Cây quyết định!

Machine Learning với PySpark

Preparing Video For Download...