PythonMastery
beginner 14 min read · lesson 5 of 6 in Machine Learning Fundamentals

Decision Trees: If-Else on Steroids

1 · The lesson

read

A decision tree is exactly what a thoughtful junior dev would write if you asked them to predict something with if-else. Look at one feature, branch. Look at the next, branch again. Reach a leaf, return the answer.

The clever part isn't the structure — it's that the tree learns the questions and their order from the data. You feed it rows, it works out which feature to split on first, what value to split at, and how deep to go.

For tabular data with messy mixed types and non-linear relationships, trees and their ensemble cousins (random forests, gradient boosting) are the bread and butter of working ML engineers.

Run these right here — scikit-learn, pandas and matplotlib all work in the browser. The first scikit-learn import takes a few seconds while it downloads; after that it's instant. Expected output is also shown in comments below each block.


1. The Intuition — a Worked Example

Predict whether a customer will buy based on three features: age, income, is_existing_customer.

A trained 2-level tree might look like:

python
                  income > 50k ?
                 ╱              ╲
              yes                no
              ╱                    ╲
   is_existing_customer?         age > 35 ?
       ╱        ╲                ╱        ╲
     yes        no             yes        no
      ▼          ▼              ▼          ▼
    BUY      NO BUY           BUY      NO BUY

For a new customer (age 40, income 60k, existing) you walk the tree:
1. income > 50k? → yes, go left.
2. is_existing_customer? → yes, go left.
3. Leaf says BUY.

That's the entire inference algorithm. No matrix multiplications, no gradients — just walking a tree of questions. You can read the tree to a stakeholder out loud, which is a superpower the linear/neural models don't give you.


2. How Does It Pick the Questions?

At each node the tree asks: "of all possible questions I could ask (every feature, every split value), which one separates the classes best?"

"Best" is measured by impurity — how mixed the classes are after the split. The most common measures:

  • Gini impurity (default for classification) — probability of mislabelling a random sample.
  • Entropy — information-theoretic mixedness.
  • Variance reduction — used for regression trees.

You don't need to compute these by hand. scikit-learn picks the best split greedily at every node. The takeaway: the tree automatically discovers which features matter and in what order.


3. Code: A 3-Level Tree

python
import matplotlib.pyplot as plt
from sklearn.datasets import load_iris
from sklearn.tree import DecisionTreeClassifier, plot_tree
from sklearn.model_selection import train_test_split

X, y = load_iris(return_X_y=True)
X_train, X_test, y_train, y_test = train_test_split(
    X, y, test_size=0.2, random_state=42
)

tree = DecisionTreeClassifier(max_depth=3, random_state=42).fit(X_train, y_train)

print("Test accuracy:", tree.score(X_test, y_test))
# → Test accuracy: 1.00

plt.figure(figsize=(12, 6))
plot_tree(tree, feature_names=["sepal_len", "sepal_wid", "petal_len", "petal_wid"],
          class_names=["setosa", "versicolor", "virginica"], filled=True)
plt.show()

You'd see a small tree with three levels of splits, each node labelled with the question (e.g. petal_len <= 2.45), the Gini score, and the class distribution. The leaves are pure (single class) or near-pure.

Same code shape works for regression with DecisionTreeRegressor.


4. The Knobs — Hyperparameters That Matter

Trees overfit easily. Their natural tendency is to grow until every training point is in its own leaf — perfectly fitting the training data, useless on test data. You control growth with three main knobs.

HyperparameterWhat it doesSensible starting value
max_depthHard cap on tree depth.3 – 10
min_samples_splitMin rows in a node before it's allowed to split.10 – 50
min_samples_leafMin rows that must remain in each leaf.5 – 20
criterion"gini" (default) or "entropy". Rarely matters.leave default

Default max_depth=None lets the tree grow forever. Always set this. A max_depth=3 model is often a stronger baseline than a max_depth=None model because it's forced to generalise.


5. Why Trees Beat Linear Models (Sometimes)

Trees win when:

  • The relationship is non-linear — sharp thresholds, interactions, kinks. Trees handle these natively.
  • Features are mixed types — numeric, categorical, boolean. No scaling, no one-hot encoding required by the tree itself (though scikit-learn's implementation still expects numeric input; encode categoricals upstream).
  • You want interpretability — plot_tree renders the actual logic the model uses.
  • Some features have wildly different scales — trees are scale-invariant. A column in millions and a column in [0,1] play together fine. No StandardScaler needed.

Trees lose when:

  • The data is tiny (under ~50 rows). A linear baseline beats trees here.
  • The data is text or images. Trees can't see structure across pixels or token sequences. Deep learning wins.
  • You need a single number summary of a smooth phenomenon. Trees output staircases, not curves.

6. The High-Variance Problem and Its Fix

Single decision trees are high variance — a small change to the training data can grow a completely different tree with a completely different decision boundary. They overfit and they're unstable.

Both problems are solved by the same trick: fit many trees and average them.

Random Forest — fit lots of trees in parallel

python
from sklearn.ensemble import RandomForestClassifier

forest = RandomForestClassifier(n_estimators=100, max_depth=10,
                                random_state=42).fit(X_train, y_train)
print("Random forest:", forest.score(X_test, y_test))
# → Random forest: 1.00
+ setup added so this can run · defines X_train, y_train, X_test, y_test
# Lightweight mock for objects whose attributes/methods aren't critical
class _AutoMock:
    def __init__(self, name='mock'): self._name = name
    def __getattr__(self, k): return _AutoMock(self._name + '.' + k)
    def __call__(self, *a, **kw):
        print('-> ' + self._name + '() called')
        return _AutoMock(self._name + '()')
    def __repr__(self): return '<mock ' + self._name + '>'
    def __str__(self): return '<mock ' + self._name + '>'
    def __bool__(self): return True
    def __iter__(self): return iter([])
    def __len__(self): return 0
    def __getitem__(self, k): return _AutoMock(self._name + '[...]')
    def __setitem__(self, k, v): pass
    def __enter__(self): return self
    def __exit__(self, *a): return False
    async def __aenter__(self): return self
    async def __aexit__(self, *a): return False
    def __add__(self, o): return self
    def __radd__(self, o): return self
    def __sub__(self, o): return self
    def __mul__(self, o): return self
    def __rmul__(self, o): return self
    def __truediv__(self, o): return self
    def __eq__(self, o): return isinstance(o, _AutoMock)
    def __hash__(self): return hash(self._name)
    def __lt__(self, o): return True
    def __le__(self, o): return True
    def __gt__(self, o): return False
    def __ge__(self, o): return False
    def __mro_entries__(self, bases): return (object,)

X_train = _AutoMock('X_train')
y_train = _AutoMock('y_train')
X_test = _AutoMock('X_test')
y_test = _AutoMock('y_test')

Each tree sees a random subset of the rows (bootstrapping) and a random subset of the features at each split. The forest averages all trees' votes. The variance averages out, the boundary smooths, accuracy jumps.

n_estimators=100 is a fine default. More trees rarely hurts; they just take longer to train.

Gradient Boosting — fit trees sequentially, each correcting the last

python
from sklearn.ensemble import GradientBoostingClassifier

gbm = GradientBoostingClassifier(n_estimators=100, max_depth=3,
                                 learning_rate=0.1,
                                 random_state=42).fit(X_train, y_train)
print("Gradient boosting:", gbm.score(X_test, y_test))
# → Gradient boosting: 1.00
+ setup added so this can run · defines X_train, y_train, X_test, y_test
# Lightweight mock for objects whose attributes/methods aren't critical
class _AutoMock:
    def __init__(self, name='mock'): self._name = name
    def __getattr__(self, k): return _AutoMock(self._name + '.' + k)
    def __call__(self, *a, **kw):
        print('-> ' + self._name + '() called')
        return _AutoMock(self._name + '()')
    def __repr__(self): return '<mock ' + self._name + '>'
    def __str__(self): return '<mock ' + self._name + '>'
    def __bool__(self): return True
    def __iter__(self): return iter([])
    def __len__(self): return 0
    def __getitem__(self, k): return _AutoMock(self._name + '[...]')
    def __setitem__(self, k, v): pass
    def __enter__(self): return self
    def __exit__(self, *a): return False
    async def __aenter__(self): return self
    async def __aexit__(self, *a): return False
    def __add__(self, o): return self
    def __radd__(self, o): return self
    def __sub__(self, o): return self
    def __mul__(self, o): return self
    def __rmul__(self, o): return self
    def __truediv__(self, o): return self
    def __eq__(self, o): return isinstance(o, _AutoMock)
    def __hash__(self): return hash(self._name)
    def __lt__(self, o): return True
    def __le__(self, o): return True
    def __gt__(self, o): return False
    def __ge__(self, o): return False
    def __mro_entries__(self, bases): return (object,)

X_train = _AutoMock('X_train')
y_train = _AutoMock('y_train')
X_test = _AutoMock('X_test')
y_test = _AutoMock('y_test')

Boosting trains tree #2 to fix tree #1's mistakes, tree #3 to fix the residual after #1+#2, and so on. The trees are kept shallow (depth 3 is typical) — each is a weak learner, and the sum is strong.

In Kaggle and in industry, gradient boosting (often via XGBoost, LightGBM, or CatBoost) is the most consistently winning approach on tabular data. Random forest is the safer, easier-to-tune cousin.

Forward link: deeper coverage in trees and ensembles.


7. Feature Importance

Trees give you feature_importances_ — a per-feature score for how much that feature contributed to splits across the tree.

python
import pandas as pd

importances = pd.Series(forest.feature_importances_,
                        index=["sepal_len", "sepal_wid", "petal_len", "petal_wid"])
print(importances.sort_values(ascending=False))
# → petal_len     0.43
# → petal_wid     0.41
# → sepal_len     0.11
# → sepal_wid     0.05
+ setup added so this can run · defines forest
# Lightweight mock for objects whose attributes/methods aren't critical
class _AutoMock:
    def __init__(self, name='mock'): self._name = name
    def __getattr__(self, k): return _AutoMock(self._name + '.' + k)
    def __call__(self, *a, **kw):
        print('-> ' + self._name + '() called')
        return _AutoMock(self._name + '()')
    def __repr__(self): return '<mock ' + self._name + '>'
    def __str__(self): return '<mock ' + self._name + '>'
    def __bool__(self): return True
    def __iter__(self): return iter([])
    def __len__(self): return 0
    def __getitem__(self, k): return _AutoMock(self._name + '[...]')
    def __setitem__(self, k, v): pass
    def __enter__(self): return self
    def __exit__(self, *a): return False
    async def __aenter__(self): return self
    async def __aexit__(self, *a): return False
    def __add__(self, o): return self
    def __radd__(self, o): return self
    def __sub__(self, o): return self
    def __mul__(self, o): return self
    def __rmul__(self, o): return self
    def __truediv__(self, o): return self
    def __eq__(self, o): return isinstance(o, _AutoMock)
    def __hash__(self): return hash(self._name)
    def __lt__(self, o): return True
    def __le__(self, o): return True
    def __gt__(self, o): return False
    def __ge__(self, o): return False
    def __mro_entries__(self, bases): return (object,)

forest = _AutoMock('forest')

Two caveats before you ship a "feature importance" chart to your boss:

1. Biased toward high-cardinality features. A column with 1000 unique values has more split opportunities than a binary column, so it scores higher just by surface area.
2. Doesn't say causation, only correlation. A feature could be important because it's a proxy for the real cause.

For more robust importance, use permutation importance (sklearn.inspection.permutation_importance) — shuffle one feature and measure how much accuracy drops.


Common Mistakes

  • Leaving max_depth=None. Default-grown trees memorise their training set. Always cap the depth.
  • Trusting a single tree. Run the same code with a different random_state and the tree may look entirely different. For decisions that matter, use a forest.
  • One-hot encoding everything. Trees handle ordinal categoricals fine if encoded as integers. One-hot blows up the feature count and slows training. (Recent versions of sklearn handle categoricals natively — check your version.)
  • Reading too much into the exact tree structure. "Look, the model says income > £49,237.50!" That number is an artefact of the training sample. The feature is the signal; the threshold is wobbly.
  • Forgetting to set n_jobs=-1 on random forests. Forests parallelise naturally; that one flag uses all your cores.

🎯 Your Turn — Tree on the Wine Dataset

The built-in load_wine dataset has 178 wines with 13 chemical features and 3 cultivar classes.

Your task:

1. Load the dataset.
2. Split 80/20 with random_state=42.
3. Train a DecisionTreeClassifier(max_depth=3, random_state=42).
4. Return the test accuracy rounded to 3 decimals and the top feature by importance, as a tuple.

python
from sklearn.datasets import load_wine
from sklearn.model_selection import train_test_split
from sklearn.tree import DecisionTreeClassifier

def train_wine_tree():
    data = load_wine()
    X, y = data.data, data.target
    feature_names = data.feature_names

    # TODO 1: split 80/20, random_state=42
    # TODO 2: fit DecisionTreeClassifier(max_depth=3, random_state=42)
    # TODO 3: compute test_acc = round(tree.score(X_test, y_test), 3)
    # TODO 4: find the feature name with the highest .feature_importances_
    # TODO 5: return (test_acc, top_feature_name)

    pass

print(train_wine_tree())
# → expected something like (0.944, 'proline')
Hint 1 — feature importances are an array tree.feature_importances_ is a NumPy array, one entry per feature in the same order as feature_names. To find the top one, use feature_names[tree.feature_importances_.argmax()].
Hint 2 — the answer should be a tuple of (float, str) Don't forget to round the accuracy. Final line: return (test_acc, top_feature_name).
Show full solution
python
from sklearn.datasets import load_wine
from sklearn.model_selection import train_test_split
from sklearn.tree import DecisionTreeClassifier

def train_wine_tree():
    data = load_wine()
    X, y = data.data, data.target
    feature_names = data.feature_names

    X_train, X_test, y_train, y_test = train_test_split(
        X, y, test_size=0.2, random_state=42
    )

    tree = DecisionTreeClassifier(max_depth=3, random_state=42).fit(X_train, y_train)

    test_acc = round(tree.score(X_test, y_test), 3)
    top_feature_name = feature_names[tree.feature_importances_.argmax()]

    return (test_acc, top_feature_name)

print(train_wine_tree())
# → (0.944, 'proline')

A depth-3 tree on raw wine chemistry hits 94% test accuracy. Try max_depth=10 — the train accuracy will hit 1.0 but the test accuracy won't move much. That's the overfitting tax. Try swapping to RandomForestClassifier and watch test accuracy creep up further.


What You Learned

  • A decision tree is a learned hierarchy of if-else questions. Each node splits on the feature that best separates the classes.
  • The three knobs that control overfitting: max_depth, min_samples_split, min_samples_leaf. Never leave max_depth=None.
  • Trees handle non-linearity and mixed scales natively — no scaling required.
  • Single trees are high-variance. Fix with ensembles: random forest (parallel, averaged) or gradient boosting (sequential, error-correcting).
  • feature_importances_ is useful but biased — prefer permutation importance for serious analysis.
  • Gradient boosting (XGBoost / LightGBM) is the default winner on tabular data; random forest is the easier baseline.

Next: Your First ML Pipeline — chaining preprocessing and a model into one bullet-proof object.