Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .github/workflows/basic_test.yml
Original file line number Diff line number Diff line change
Expand Up @@ -27,3 +27,5 @@ jobs:
- name: Run basic test
run: |
tests/basic.sh
- name: Test tree memory assembly
run: python -m unittest discover -s tests/linear -p 'test_*.py'
20 changes: 20 additions & 0 deletions docs/examples/plot_linear_tree_tutorial.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,26 @@
#
# The ``train_tree`` function in this tutorial is based on the work of :cite:t:`SK20a`.
#
# Memory use during training
# ~~~~~~~~~~~~~~~~~~~~~~~~~~
#
# Tree training writes each node's sparse weights to a temporary file and releases
# them before training the next node. It then allocates the final CSC matrix once
# and fills it in chunks of at most 8 MiB for float64 weights. This avoids keeping
# all node matrices and the concatenated model in memory at the same time, while
# preserving weight precision, classifier order, and the saved model format.
#
# Final assembly needs memory for the final sparse model, one chunk, and small
# bookkeeping arrays. Training data, the tree structure, a node's training workspace,
# and previously trained ensemble members still require additional memory.
#
# Allow temporary disk space roughly equal to the total sparse node weights, plus
# array headers. The temporary file is closed automatically on success or failure.
# Set ``TMPDIR`` before starting Python to use a disk with sufficient free space;
# avoid a RAM-backed temporary directory if the goal is to reduce physical RAM use.
# This trades a sequential write/read of the model for lower peak RAM. The returned
# model is fully in memory and does not depend on the temporary file.
#
# ``train_tree`` achieves this speedup by approximating ``train_1vsrest``. To check whether the approximation
# performs well, we'll compute some metrics on the test set.

Expand Down
97 changes: 76 additions & 21 deletions libmultilabel/linear/tree.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

import tempfile
from typing import Callable

import numpy as np
Expand All @@ -16,6 +17,7 @@

DEFAULT_K = 100
DEFAULT_DMAX = 10
_WEIGHT_CHUNK_SIZE = 1 << 20 # At most 8 MiB per chunk for float64 weights or int64 indices.


class Node:
Expand Down Expand Up @@ -235,6 +237,11 @@ def train_tree(
"""Train a linear model for multi-label data using a divide-and-conquer strategy.
The algorithm used is based on https://github.com/xmc-aalto/bonsai.

Node weights are staged in a temporary file during training to avoid retaining
a second copy of the model while assembling the final CSC weight matrix.
The temporary directory (configurable with TMPDIR) needs space for the sparse
node weights. The returned model is held in memory and does not depend on this file.

Args:
y (sparse.csr_matrix): A 0/1 matrix with dimensions number of instances * number of classes.
x (sparse.csr_matrix): A matrix with dimensions number of instances * number of features.
Expand All @@ -252,6 +259,7 @@ def train_tree(
label_representation = sklearn.preprocessing.normalize(label_representation, norm="l2", axis=1)
root = _build_tree(label_representation, np.arange(y.shape[1]), 0, K, dmax)
root.is_root = True
del label_representation

num_nodes = 0
# Both type(x) and type(y) are sparse.csr_matrix
Expand All @@ -265,6 +273,7 @@ def count(node):
node.num_features_used = np.count_nonzero(features_used_perlabel[:, node.label_map].sum(axis=1))

root.dfs(count)
del features_used_perlabel

model_size = get_estimated_model_size(root)
print(f"The estimated tree model size is: {model_size / (1024**3):.3f} GB")
Expand All @@ -286,10 +295,10 @@ def visit(node):
_train_node(y[relevant_instances], x[relevant_instances], options, node)
pbar.update()

root.dfs(visit)
pbar.close()

flat_model, node_ptr = _flatten_model(root)
try:
flat_model, node_ptr = _flatten_model(root, train_node=visit)
finally:
pbar.close()
return TreeModel(root, flat_model, node_ptr)


Expand Down Expand Up @@ -374,7 +383,7 @@ def _train_node(y: sparse.csr_matrix, x: sparse.csr_matrix, options: str, node:
node.model.weights = sparse.csc_matrix(node.model.weights)


def _flatten_model(root: Node) -> tuple[linear.FlatModel, np.ndarray]:
def _flatten_model(root: Node, train_node: Callable[[Node], None] | None = None) -> tuple[linear.FlatModel, np.ndarray]:
"""Flatten tree weight matrices into a single weight matrix. The flattened weight
matrix is used to predict all possible values, which is cached for beam search.
This pessimizes complexity but is faster in practice.
Expand All @@ -386,34 +395,80 @@ def _flatten_model(root: Node) -> tuple[linear.FlatModel, np.ndarray]:

Args:
root (Node): Root of the tree.
train_node (Callable, optional): Train each node immediately before staging
its weights. If omitted, all nodes must already have trained models.

Returns:
tuple[linear.FlatModel, np.ndarray]: The flattened model and the ranges of each node.
"""
index = 0
weights = []
bias = root.model.bias

def visit(node):
assert bias == node.model.bias
nonlocal index
node.index = index
index += 1
weights.append(node.model.__dict__.pop("weights"))

root.dfs(visit)
node_ptr = [0]
node_nnz = []
bias = None
num_features = None
data_dtype = None

# Staging before allocation avoids keeping all node weights and the flattened
# matrix in RAM together. A single file also avoids one open file per node.
with tempfile.TemporaryFile(prefix="libmultilabel-weights-") as weights_file:

def visit(node):
nonlocal bias, num_features, data_dtype
if train_node is not None:
train_node(node)
weights = sparse.csc_matrix(node.model.weights, copy=False)
if node is root:
bias = node.model.bias
num_features = weights.shape[0]
data_dtype = weights.dtype
assert bias == node.model.bias
if weights.shape[0] != num_features:
raise ValueError("Node weight matrices must have the same number of features.")
data_dtype = np.result_type(data_dtype, weights.dtype)
node.index = len(node_nnz)
node_ptr.append(node_ptr[-1] + weights.shape[1])
node_nnz.append(weights.nnz)

for array in (weights.data[: weights.nnz], weights.indices[: weights.nnz], weights.indptr[:-1]):
for start in range(0, array.size, _WEIGHT_CHUNK_SIZE):
np.save(weights_file, array[start : start + _WEIGHT_CHUNK_SIZE], allow_pickle=False)
del node.model.weights

root.dfs(visit)

node_ptr = np.asarray(node_ptr, dtype=np.int64)
total_nnz = sum(node_nnz)
num_classifiers = int(node_ptr[-1])
# Both indices and indptr must use int64 once any dimension or the
# cumulative NNZ exceeds int32, even if every node individually fits.
index_dtype = np.int64 if max(num_features, num_classifiers, total_nnz) > np.iinfo(np.int32).max else np.int32
data = np.empty(total_nnz, dtype=data_dtype)
indices = np.empty(total_nnz, dtype=index_dtype)
indptr = np.empty(num_classifiers + 1, dtype=index_dtype)

weights_file.seek(0)
offset = 0
for i, nnz in enumerate(node_nnz):
end = offset + nnz
columns = slice(node_ptr[i], node_ptr[i + 1])
for array in (data[offset:end], indices[offset:end], indptr[columns]):
for start in range(0, array.size, _WEIGHT_CHUNK_SIZE):
array[start : start + _WEIGHT_CHUNK_SIZE] = np.load(weights_file, allow_pickle=False)
# Offset in the destination dtype to avoid overflowing int32 node pointers.
indptr[columns] += offset
offset = end
indptr[-1] = total_nnz

# Matching index dtypes let SciPy reuse these buffers without a full-model copy.
weights = sparse.csc_matrix((data, indices, indptr), shape=(num_features, num_classifiers), copy=False)

model = linear.FlatModel(
name="flattened-tree",
weights=sparse.hstack(weights, "csc"),
weights=weights,
bias=bias,
thresholds=0,
multiclass=False,
)

# w.shape[1] is the number of labels/metalabels of each node
node_ptr = np.cumsum([0] + list(map(lambda w: w.shape[1], weights)))

return model, node_ptr


Expand Down
113 changes: 113 additions & 0 deletions tests/linear/benchmark_tree_memory.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,113 @@
"""Compare the original and staged tree assembly in separate processes.

Run from the repository root:
PYTHONPATH=. python tests/linear/benchmark_tree_memory.py

The default synthetic model occupies about 387 MiB and needs comparable temporary
disk space. This measures weight construction/assembly, not full dataset training.
"""

import argparse
import hashlib
import json
import resource
import subprocess
import sys
import time

import numpy as np
import psutil
from scipy import sparse

from libmultilabel.linear import linear, tree


def benchmark(args):
children = [tree.Node(np.arange(8 * i, 8 * (i + 1)), []) for i in range(args.nodes)]
root = tree.Node(np.arange(8 * args.nodes), children)
node_number = 0

def train_node(node):
nonlocal node_number
columns = len(node.label_map) if node.isLeaf() else len(node.children)
per_column = args.nnz_per_node // columns
nnz = columns * per_column
weights = sparse.csc_matrix(
(
np.full(nnz, node_number + 0.5),
np.tile(np.arange(per_column, dtype=np.int32), columns),
np.arange(columns + 1, dtype=np.int32) * per_column,
),
shape=(args.nnz_per_node, columns),
)
node.model = linear.FlatModel("node", weights, -1, 0, False)
node_number += 1

baseline_rss = psutil.Process().memory_info().rss
start = time.perf_counter()
if args.mode == "original":
root.dfs(train_node)
blocks = []
root.dfs(lambda node: blocks.append(node.model.__dict__.pop("weights")))
node_ptr = np.cumsum([0] + [block.shape[1] for block in blocks])
weights = sparse.hstack(blocks, format="csc")
del blocks
else:
model, node_ptr = tree._flatten_model(root, train_node)
weights = model.weights
elapsed = time.perf_counter() - start

# Hash buffers directly, without allocating another byte string or dense matrix.
digest = hashlib.sha256()
for array in (weights.data, weights.indices, weights.indptr, node_ptr):
digest.update(memoryview(array))
peak = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
if sys.platform != "darwin":
peak *= 1024
return {
"mode": args.mode,
"nodes": args.nodes + 1,
"nnz": weights.nnz,
"model_mib": sum(a.nbytes for a in (weights.data, weights.indices, weights.indptr)) / 2**20,
"baseline_rss_mib": baseline_rss / 2**20,
"peak_rss_mib": peak / 2**20,
"assembly_seconds": elapsed,
"sha256": digest.hexdigest(),
}


def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--nodes", type=int, default=128, help="Number of leaf nodes.")
parser.add_argument("--nnz-per-node", type=int, default=262144)
parser.add_argument("--mode", choices=("original", "staged"), help=argparse.SUPPRESS)
args = parser.parse_args()
if args.nodes <= 0 or args.nnz_per_node < max(args.nodes, 8):
parser.error("Use positive node counts and at least max(nodes, 8) nonzeros per node.")
if args.mode:
print(json.dumps(benchmark(args)))
return

results = []
for mode in ("original", "staged"):
output = subprocess.check_output(
[
sys.executable,
__file__,
"--mode",
mode,
"--nodes",
str(args.nodes),
"--nnz-per-node",
str(args.nnz_per_node),
],
text=True,
)
results.append(json.loads(output))
if results[0]["sha256"] != results[1]["sha256"]:
raise AssertionError("The original and staged weight buffers differ.")
print(json.dumps({"results": results, "identical_buffers": True}, indent=2))


if __name__ == "__main__":
main()
Loading