Trees

A Tree is a recursive node: every node has a name, a branch length and a list of children. Trees are read from and written to Newick, and can be plotted with treeplot().

from matplotlib import pyplot as plt

from picea import Tree, treeplot

Parsing and navigating

tree = Tree.from_newick("((a:1,b:2)ab:1,(c:1,d:1)cd:2)root:0;")
tree.to_newick(branch_lengths=True)
'((a:1.0,b:2.0)ab:1.0,(c:1.0,d:1.0)cd:2.0)root;'
len(tree.nodes), [leaf.name for leaf in tree.leaves]
(7, ['a', 'b', 'c', 'd'])

Nodes can be looked up by name with loc, and every node knows its parent and root.

a = tree.loc["a"]
a.parent.name, a.root.name, a.cumulative_length
('ab', 'root', 2.0)

Traverse the tree depth first (pre- or post-order) or breadth first.

[node.name for node in tree.depth_first(post_order=False)]
['root', 'ab', 'a', 'b', 'cd', 'c', 'd']
[node.name for node in tree.breadth_first()]
['root', 'ab', 'cd', 'a', 'b', 'c', 'd']

Plotting

This gene tree of hydroxycinnamoyl transferase (HCT) homologs has branch lengths and support values as internal node names.

hct = Tree.from_newick(filename="data/tree.newick")
len(hct.leaves)
41
hct.rename_leaves(lambda name: name.removesuffix(".1"))

fig, ax = plt.subplots(figsize=(8, 9))
treeplot(hct, style="square", ax=ax);
../_images/a776e0da83b6af85cf4d575de1d6b995db8cf965d415d39d6d77536ddc168c6f.png

The radial style puts the root in the center.

fig, ax = plt.subplots(figsize=(9, 9))
treeplot(hct, style="radial", node_labels=False, ax=ax);
../_images/b8ad52da378348df194d5cbdcc5d742bf5814d7fb90eb6af87a44450507b9ee7.png

Leaf markers can be styled per leaf with a function, for example to color leaves by species. Without branch lengths the tree is drawn as a cladogram.

from matplotlib.lines import Line2D

species = {
    "AT": ("Arabidopsis thaliana", "tab:blue"),
    "Potri": ("Populus trichocarpa", "tab:orange"),
    "Eucgr": ("Eucalyptus grandis", "tab:green"),
    "Glyma": ("Glycine max", "tab:red"),
    "Medtr": ("Medicago truncatula", "tab:purple"),
    "Fvesca": ("Fragaria vesca", "tab:brown"),
    "PanWU": ("Parasponia andersonii", "tab:pink"),
    "TorRG": ("Trema orientalis", "tab:olive"),
}


def leaf_color(leaf: Tree) -> str:
    return next(color for prefix, (_, color) in species.items() if leaf.name.startswith(prefix))


fig, ax = plt.subplots(figsize=(8, 9))
treeplot(hct, style="square", branchlengths=False, node_labels=False, leaf_marker_fill=leaf_color, ax=ax)
ax.legend(
    handles=[Line2D([], [], marker="o", linestyle="", color=color, label=name) for name, color in species.values()],
    loc="upper left",
    bbox_to_anchor=(1.01, 1),
);
../_images/d8c8161b1fc74369ecc58f35e5e34a37d66ff8ffdc833c01411a831b77fd1a84.png

From hierarchical clustering

Trees can also be created from a fitted scikit-learn AgglomerativeClustering model.

import numpy as np
from sklearn.cluster import AgglomerativeClustering

X = np.array([[1, 2], [1, 4], [1, 0], [4, 2], [4, 4], [4, 0]])
clustering = AgglomerativeClustering().fit(X)
cluster_tree = Tree.from_sklearn(clustering)
cluster_tree.to_newick()
'((2,(0,1)),(4,(3,5)));'
fig, ax = plt.subplots(figsize=(4, 3))
treeplot(cluster_tree, style="square", branchlengths=False, ax=ax);
../_images/8d835300e488c73e995f2edef71f7816709a8c1cbc047399493d0f6c3f476a7b.png