You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用含10组二维坐标与类别标签的numpy数组构建决策树分类器

Alright, let's break down exactly how to build a decision tree classifier for your coordinate-based data with 12 direction labels. Your samples are structured as 10 sets of (x,y) pairs plus a string label—here's a step-by-step guide with code examples:

Step 1: Preprocess Your Data (Flatten 2D Coordinates)

Decision tree algorithms (like those in scikit-learn) expect 1D feature vectors, not 2D arrays. So you’ll need to flatten each sample’s 10x2 coordinate array into a single 20-element vector. For example, your sample array becomes [-19.24, -0.69, -17.40, -0.84, ..., 18.94, 6.93].

Here’s how to do this with numpy:

import numpy as np
from sklearn.tree import DecisionTreeClassifier
from sklearn.model_selection import train_test_split
from sklearn.metrics import classification_report

# Assume your raw data is stored in a list like this
data = [
    [np.array([[-19.24, -0.69], [-17.40, -0.84], ...]), 'left'],
    [np.array([[5.67, 2.34], [8.90, 1.23], ...]), 'right'],
    # Add all your other samples here
]

# Split features and labels, flattening the coordinate arrays
X = np.array([sample[0].flatten() for sample in data])
y = np.array([sample[1] for sample in data])
Step 2: Split Data into Training & Test Sets

To make sure your model generalizes well to unseen data, split your dataset into training (for fitting the model) and test (for evaluating performance) subsets. Use stratify=y to preserve the distribution of your 12 labels across both sets—this prevents skewed results if some labels are rare.

X_train, X_test, y_train, y_test = train_test_split(
    X, y, test_size=0.2, random_state=42, stratify=y
)
Step 3: Train the Decision Tree

Start with a basic decision tree using scikit-learn’s DecisionTreeClassifier. You can tweak parameters later to improve performance:

# Initialize the classifier
clf = DecisionTreeClassifier(random_state=42)

# Train on the training data
clf.fit(X_train, y_train)
Step 4: Evaluate Model Performance

Use the test set to measure how well your model predicts unseen labels. The classification_report gives you granular metrics (precision, recall, F1-score) for each of your 12 label classes:

# Generate predictions for the test set
y_pred = clf.predict(X_test)

# Print detailed classification metrics
print(classification_report(y_test, y_pred))

This report will highlight which labels the model struggles with, helping you target improvements.

Decision trees tend to overfit to training data by default. Tune these key parameters to improve generalization:

  • max_depth: Limits how deep the tree can grow (prevents overly complex splits)
  • min_samples_split: Minimum number of samples required to split a node
  • min_samples_leaf: Minimum number of samples required for a leaf node
  • class_weight: Set to 'balanced' if some labels have far fewer samples than others

Use grid search to find the best parameter combination:

from sklearn.model_selection import GridSearchCV

# Define parameter grid to test
param_grid = {
    'max_depth': [3, 5, 7, 10],
    'min_samples_split': [2, 5, 10],
    'min_samples_leaf': [1, 2, 4]
}

# Run grid search with 5-fold cross-validation
grid_search = GridSearchCV(DecisionTreeClassifier(random_state=42), param_grid, cv=5)
grid_search.fit(X_train, y_train)

# Print best parameters and use the optimized model
print("Best hyperparameters:", grid_search.best_params_)
best_clf = grid_search.best_estimator_

# Evaluate the optimized model
y_pred_best = best_clf.predict(X_test)
print(classification_report(y_test, y_pred_best))
Step 6: Visualize the Decision Tree (Optional)

If you want to see how the tree makes decisions, visualize it with scikit-learn’s plot_tree:

from sklearn.tree import plot_tree
import matplotlib.pyplot as plt

# Create a large figure to fit the tree
plt.figure(figsize=(20, 10))
plot_tree(
    best_clf,
    feature_names=[f'coord_{i}' for i in range(20)],  # Name flattened features
    class_names=sorted(np.unique(y)),  # List all 12 label classes
    filled=True,
    rounded=True
)
plt.show()

This plot will show you which coordinate features the tree prioritizes for splitting nodes.


内容的提问来源于stack exchange,提问作者Christian

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.22 08:06:31