diff --git a/.github/workflows/basic_test.yml b/.github/workflows/basic_test.yml index c786801f..d59aedd5 100644 --- a/.github/workflows/basic_test.yml +++ b/.github/workflows/basic_test.yml @@ -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' diff --git a/docs/examples/plot_linear_tree_tutorial.py b/docs/examples/plot_linear_tree_tutorial.py index d0c70318..a3846a4e 100644 --- a/docs/examples/plot_linear_tree_tutorial.py +++ b/docs/examples/plot_linear_tree_tutorial.py @@ -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. diff --git a/libmultilabel/linear/tree.py b/libmultilabel/linear/tree.py index f8411b68..222a6233 100644 --- a/libmultilabel/linear/tree.py +++ b/libmultilabel/linear/tree.py @@ -1,5 +1,6 @@ from __future__ import annotations +import tempfile from typing import Callable import numpy as np @@ -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: @@ -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. @@ -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 @@ -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") @@ -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) @@ -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. @@ -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 diff --git a/tests/linear/benchmark_tree_memory.py b/tests/linear/benchmark_tree_memory.py new file mode 100644 index 00000000..e2688235 --- /dev/null +++ b/tests/linear/benchmark_tree_memory.py @@ -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() diff --git a/tests/linear/test_tree_memory.py b/tests/linear/test_tree_memory.py new file mode 100644 index 00000000..dfa56be7 --- /dev/null +++ b/tests/linear/test_tree_memory.py @@ -0,0 +1,202 @@ +"""Regression tests for bounded-memory assembly of tree weights. + +Run with: python -m unittest discover -s tests/linear -p 'test_*.py' +""" + +import copy +import pickle +import tempfile +import unittest +import weakref +from unittest.mock import patch + +import numpy as np +from numpy.testing import assert_array_equal, assert_allclose +from scipy import sparse + +from libmultilabel.linear import linear, tree + + +def make_tree(bias=1, dtype=np.float64): + branch = tree.Node(np.array([2, 0, 3]), [tree.Node(np.array([2, 0]), []), tree.Node(np.array([3]), [])]) + root = tree.Node(np.arange(4), [branch, tree.Node(np.array([1]), [])]) + root.is_root = True + nodes = [] + root.dfs(nodes.append) + rng = np.random.default_rng(42) + for i, node in enumerate(nodes): + columns = len(node.label_map) if node.isLeaf() else len(node.children) + weights = rng.normal(size=(5, columns)).astype(dtype) + weights[weights < 0] = 0 + if i == 2: + weights[:] = 0 + node.model = linear.FlatModel("node", sparse.csc_matrix(weights), bias, 0, False) + return root, nodes + + +def reference_flatten(root): + """The original in-memory hstack path, retained only as a test oracle.""" + nodes = [] + root.dfs(nodes.append) + weights = [] + for i, node in enumerate(nodes): + node.index = i + weights.append(node.model.__dict__.pop("weights")) + flat = linear.FlatModel("flattened-tree", sparse.hstack(weights, format="csc"), root.model.bias, 0, False) + return flat, np.cumsum([0] + [w.shape[1] for w in weights]) + + +class TreeMemoryTests(unittest.TestCase): + def assert_sparse_equal(self, actual, expected): + self.assertEqual(actual.format, "csc") + self.assertEqual(actual.shape, expected.shape) + self.assertEqual(actual.dtype, expected.dtype) + assert_array_equal(actual.data, expected.data) + assert_array_equal(actual.indices, expected.indices) + assert_array_equal(actual.indptr, expected.indptr) + + def test_matches_hstack_and_predictions(self): + for dtype in (np.float32, np.float64): + for bias in (-1, 1): + with self.subTest(dtype=dtype, bias=bias): + root, nodes = make_tree(bias, dtype) + expected_root = copy.deepcopy(root) + expected = tree.TreeModel(expected_root, *reference_flatten(expected_root)) + # Force multiple chunks even for these small matrices. + with patch.object(tree, "_WEIGHT_CHUNK_SIZE", 2): + actual = tree.TreeModel(root, *tree._flatten_model(root)) + self.assert_sparse_equal(actual.flat_model.weights, expected.flat_model.weights) + assert_array_equal(actual.node_ptr, expected.node_ptr) + self.assertEqual([node.index for node in nodes], list(range(len(nodes)))) + self.assertTrue(all(not hasattr(node.model, "weights") for node in nodes)) + x = sparse.csr_matrix(np.random.default_rng(7).normal(size=(6, 4 if bias > 0 else 5))) + for beam_width in (1, 2, 4): + assert_array_equal(actual.predict_values(x, beam_width), expected.predict_values(x, beam_width)) + + def test_mixed_dtypes_and_noncanonical_entries(self): + root, nodes = make_tree() + nodes[1].model.weights = nodes[1].model.weights.astype(np.float32) + # Duplicates, unsorted indices, and an explicit zero must be preserved. + root.model.weights = sparse.csc_matrix(([2.0, 0.0, 3.0], [3, 1, 3], [0, 3, 3]), shape=(5, 2)) + expected, _ = reference_flatten(copy.deepcopy(root)) + actual, _ = tree._flatten_model(root) + self.assert_sparse_equal(actual.weights, expected.weights) + + def test_empty_and_single_node_models(self): + for shape in ((5, 0), (5, 3), (0, 3)): + with self.subTest(shape=shape): + root = tree.Node(np.arange(shape[1]), []) + root.model = linear.FlatModel("node", sparse.csc_matrix(shape), -1, 0, False) + flat, node_ptr = tree._flatten_model(root) + self.assertEqual(flat.weights.shape, shape) + self.assertEqual(flat.weights.nnz, 0) + assert_array_equal(node_ptr, [0, shape[1]]) + + def test_large_row_indices_use_int64(self): + rows = np.iinfo(np.int32).max + 2 + root = tree.Node(np.arange(2), []) + weights = sparse.csc_matrix(([1.5, 2.5], [0, rows - 1], [0, 1, 2]), shape=(rows, 2)) + root.model = linear.FlatModel("node", weights, -1, 0, False) + flat, _ = tree._flatten_model(root) + self.assert_sparse_equal(flat.weights, weights) + self.assertEqual(flat.weights.indices.dtype, np.int64) + self.assertEqual(flat.weights.indptr.dtype, np.int64) + + def test_training_releases_each_node_and_reads_bounded_chunks(self): + root, nodes = make_tree() + for node in nodes: + del node.model + refs = [] + visited = [] + original_load = np.load + + def train_node(node): + self.assertTrue(all(ref() is None for ref in refs)) + visited.append(node) + columns = len(node.label_map) if node.isLeaf() else len(node.children) + weights = sparse.csc_matrix(np.ones((5, columns))) + refs.extend(weakref.ref(array) for array in (weights.data, weights.indices, weights.indptr)) + node.model = linear.FlatModel("node", weights, 1, 0, False) + + def load_chunk(*args, **kwargs): + self.assertEqual(visited, nodes) + self.assertTrue(all(ref() is None for ref in refs)) + chunk = original_load(*args, **kwargs) + self.assertLessEqual(chunk.size, 2) + return chunk + + with patch.object(tree, "_WEIGHT_CHUNK_SIZE", 2), patch.object(tree.np, "load", side_effect=load_chunk): + flat, _ = tree._flatten_model(root, train_node) + assert_array_equal(flat.weights.toarray(), np.ones((5, 8))) + + def test_final_csc_reuses_allocated_buffers(self): + root, _ = make_tree() + original_csc = sparse.csc_matrix + checked = [] + + def create_csc(arg, *args, **kwargs): + result = original_csc(arg, *args, **kwargs) + if isinstance(arg, tuple) and len(arg) == 3: + for original, final in zip(arg, (result.data, result.indices, result.indptr)): + self.assertTrue(np.shares_memory(original, final)) + checked.append(True) + return result + + with patch.object(tree.sparse, "csc_matrix", side_effect=create_csc): + tree._flatten_model(root) + self.assertEqual(checked, [True]) + + def test_temporary_file_closes_on_success_and_failure(self): + for failure in (None, "train", "write", "read"): + with self.subTest(failure=failure): + root, _ = make_tree() + staging_file = tempfile.TemporaryFile() + + def train_node(node): + if failure == "train" and node is not root: + raise RuntimeError("training interrupted") + + with patch.object(tree.tempfile, "TemporaryFile", return_value=staging_file): + if failure in ("write", "read"): + operation = "save" if failure == "write" else "load" + with patch.object(tree.np, operation, side_effect=OSError("disk failure")): + with self.assertRaisesRegex(OSError, "disk failure"): + tree._flatten_model(root, train_node) + elif failure == "train": + with self.assertRaisesRegex(RuntimeError, "training interrupted"): + tree._flatten_model(root, train_node) + else: + model = tree.TreeModel(root, *tree._flatten_model(root, train_node)) + self.assertTrue(staging_file.closed) + if failure is None: + restored = pickle.loads(pickle.dumps(model, protocol=pickle.HIGHEST_PROTOCOL)) + self.assert_sparse_equal(restored.flat_model.weights, model.flat_model.weights) + x = sparse.csr_matrix(np.ones((2, 4))) + assert_array_equal(restored.predict_values(x), model.predict_values(x)) + + def test_real_training_matches_original_assembly(self): + rng = np.random.default_rng(6) + x = sparse.csr_matrix(rng.normal(size=(32, 4))) + y = sparse.csr_matrix(rng.integers(0, 2, size=(32, 4))) + root, nodes = make_tree() + for node in nodes: + del node.model + + def old_assembly(root, train_node): + root.dfs(train_node) + return reference_flatten(root) + + # A primal solver avoids differences from LIBLINEAR's random dual updates. + with patch.object(tree, "_flatten_model", side_effect=old_assembly): + expected = tree.train_tree(y, x, options="-s 2 -B 1 -m 1 -q", root=copy.deepcopy(root), verbose=False) + actual = tree.train_tree(y, x, options="-s 2 -B 1 -m 1 -q", root=root, verbose=False) + assert_allclose( + actual.flat_model.weights.toarray(), expected.flat_model.weights.toarray(), rtol=1e-12, atol=1e-12 + ) + assert_array_equal(actual.node_ptr, expected.node_ptr) + for beam_width in (1, 3): + assert_allclose(actual.predict_values(x, beam_width), expected.predict_values(x, beam_width), rtol=1e-12) + + +if __name__ == "__main__": + unittest.main()