diff --git a/AGENTS.md b/AGENTS.md index 716853ff..ea32117c 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -65,13 +65,38 @@ brings its own way of running a step. `macro_f1` scores the confusion matrix `trainer/cuda.py` is the one exception to that framework independence, and holds everything about a batch-size probe that torch has to answer. `usable_memory_bytes` and `limit_cuda_memory` turn a `--vram-limit-gb` setting into the budget a probe measures against and the cap that holds -the process to it, and capping an allocator has no NVML equivalent. `probe_batch_size` is the -whole probe for a node whose measurement is a single call — it resolves the limit, falls back -without a card, holds the safety margin and runs the search. A node that must build a throwaway -model first reserves the margin before building it, and so composes the same pieces itself: -`reserve_margin`, `measured_fits` and `find_batch_size`. `measured_fits` is where an -out-of-memory failure is told from a bug — both arrive as the same exception types, and a probe -that confuses them reports the smallest batch size as the card's fault. +the process to it, and capping an allocator has no NVML equivalent. + +`measure_batch_size` is what every trainer calls, whatever shape its hyperparameters have: it +takes the requested size as an `int`, so a node parsing into a dataclass enters the same door as +one keeping a dict. `max_batch_size` — the hyperparameter named by the `REQUESTED_BATCH_SIZE` +constant and read through `requested_batch_size`, so every node agrees on the spelling — is the +largest batch the training may use, and it is measured rather than trusted: a size that fits is +used as named, whether or not it is a power of two, and one that does not becomes the largest +power of two below it that does. A training that starts small beats one that runs out of memory +partway through. The settled size is returned and **not** stored anywhere the next measurement +would read it: a trainer reports it as `batch_size`, never under `max_batch_size`, because the +node saves the hyperparameters with the training (`LastTrainingIO`) and a training resumed after a +restart reads them back — it would otherwise take its first run's measurement as its bound instead +of measuring again. The loop does not hand reported values to a later training: it builds each one +from the project configuration and the job's override, taking only `resolution` from a base +training. + +Below it, `probe_batch_size` is the whole probe except the step itself — it resolves the limit, +falls back without a card, holds the safety margin, runs the search and releases what the trials +left behind. `on_out_of_memory` is how a node that builds a throwaway model drops an optimizer's +gradients after a failed trial. `minimum` is for a step that cannot run on a single sample at all +— BatchNorm over a 1x1 feature map, or a training whose validation halves the batch — and +`candidate` is the one way the search returns a size that is not a power of two, and only ever +one that was named and then measured. A node does not compose `reserve_margin`, `measured_fits` +and `find_batch_size` itself; they are the pieces `probe_batch_size` is built from. `measured_fits` +is where an out-of-memory failure is told from a bug — both arrive as the same exception types, +and a probe that confuses them reports the smallest batch size as the card's fault. + +A trainer that probes opts into the budget flag with `node_parser(vram_limit=True)`, which adds +`--vram-limit-gb` / `VRAM_LIMIT_GB`. The cap does not survive a spawn, so a script the trainer +spawns declares the same flag with `add_vram_limit_argument`, and the node passes the value on +explicitly; that script calls `limit_cuda_memory` itself. It imports torch, the package does **not** declare it, and only a trainer imports the module — so the library keeps working where nothing trains. Its unit test installs a stand-in under the name diff --git a/README.md b/README.md index 45ba8693..c695ec8c 100644 --- a/README.md +++ b/README.md @@ -29,6 +29,7 @@ You can configure connection to our Learning Loop by specifying the following en | MAX_UNCERTAIN_THRESHOLD | - | largest confidence (float) at which auto-upload will happen | Detector (opt.) | 0.6 | | EXCLUSIVE_MODEL_BUILD | - | Reject detections during update to save VRAM (set to 1) | Detector (opt.) | 0 | | INFERENCE_BATCH_SIZE | - | Batch size of trainer when calculating detections | Trainer (opt.) | 10 | +| VRAM_LIMIT_GB | - | GPU memory (GB) a training may use; the batch size is probed against it (`--vram-limit-gb`, trainers built with `node_parser(vram_limit=True)`) | Trainer (opt.) | 0 (whole card) | | RESTART_AFTER_TRAINING | - | Restart the trainer after training (set to 1) | Trainer (opt.) | 0 | | KEEP_OLD_TRAININGS | - | Do not delete old trainings (set to 1) | Trainer (opt.) | 0 | | TRAINER_IDLE_TIMEOUT_SEC | - | Automatically shutdown trainer after timeout (in seconds) | Trainer (opt.) | 0 (disabled) | diff --git a/learning_loop_node/helpers/entrypoint.py b/learning_loop_node/helpers/entrypoint.py index 4c42602e..0d56d45d 100644 --- a/learning_loop_node/helpers/entrypoint.py +++ b/learning_loop_node/helpers/entrypoint.py @@ -12,21 +12,27 @@ import configargparse import uvicorn +from ..trainer.batch_size import VRAM_LIMIT_GB_FLAG, VRAM_LIMIT_GB_HELP + logger = logging.getLogger(__name__) -def node_parser(*, description: str, legacy_env_prefix: str = '') -> configargparse.ArgumentParser: +def node_parser(*, description: str, legacy_env_prefix: str = '', + vram_limit: bool = False) -> configargparse.ArgumentParser: """Build the parser for a node, pre-loaded with the settings every node has. :param legacy_env_prefix: A prefix an earlier version of this node required, e.g. ``'MY_DETECTOR_'``. Both spellings keep working, with a warning: the prefix on the current name (``MY_DETECTOR_NODE_HOST``) and the prefix on the flag it was originally applied to (``MY_DETECTOR_HOST``). Leave empty for a node that never used one. + :param vram_limit: Add :data:`~learning_loop_node.trainer.batch_size.VRAM_LIMIT_GB_FLAG`. """ parser = _NodeArgumentParser(description=description, legacy_env_prefix=legacy_env_prefix) parser.add_argument('--host', default='0.0.0.0', env_var='NODE_HOST', help='Host interface to bind to') parser.add_argument('--port', type=int, default=80, env_var='NODE_PORT', help='Port to bind to') + if vram_limit: + parser.add_argument(VRAM_LIMIT_GB_FLAG, type=float, default=0, help=VRAM_LIMIT_GB_HELP) return parser diff --git a/learning_loop_node/tests/unit/test_batch_size.py b/learning_loop_node/tests/unit/test_batch_size.py index 9ad3fd3e..eb503c69 100644 --- a/learning_loop_node/tests/unit/test_batch_size.py +++ b/learning_loop_node/tests/unit/test_batch_size.py @@ -4,13 +4,16 @@ from ...trainer.batch_size import ( MIN_TRAIN_STEPS_PER_EPOCH, + REQUESTED_BATCH_SIZE, batch_count, dataset_limit, find_batch_size, is_out_of_memory, no_gpu_batch_size, + requested_batch_size, smaller_pot, ) +from ...trainer.exceptions import InsufficientMemoryError def test_the_search_doubles_up_to_the_limit(): @@ -30,6 +33,57 @@ def test_a_limit_that_is_not_a_power_of_two_is_rounded_down(): assert find_batch_size(_fits_up_to(1024), limit=1) == 1 +def test_a_minimum_keeps_the_search_off_the_sizes_below_it(): + ran: list[int] = [] + assert find_batch_size(_fits_up_to(16, ran), limit=64, minimum=2) == 16 + assert ran == [2, 4, 8, 16, 32], 'the smallest trial is the minimum, not one' + + +def test_a_minimum_is_rounded_down_to_a_power_of_two(): + ran: list[int] = [] + find_batch_size(_fits_up_to(64, ran), limit=64, minimum=3) + assert ran[0] == 2 + + +def test_a_minimum_that_does_not_fit_reports_its_own_size(): + with pytest.raises(InsufficientMemoryError, match='batch size 4'): + find_batch_size(_fits_up_to(2, []), limit=32, minimum=4) + + +def test_a_minimum_above_the_limit_still_gets_tried(): + ran: list[int] = [] + assert find_batch_size(_fits_up_to(64, ran), limit=2, minimum=8) == 8 + assert ran == [8] + + +def test_a_named_candidate_that_fits_is_used_as_named(): + ran: list[int] = [] + assert find_batch_size(_fits_up_to(24, ran), limit=24, candidate=24) == 24 + assert ran == [1, 2, 4, 8, 16, 24], 'tried only after the doubling reached its ceiling' + + +def test_a_named_candidate_that_does_not_fit_leaves_the_power_of_two(): + assert find_batch_size(_fits_up_to(16, []), limit=24, candidate=24) == 16 + + +def test_a_candidate_is_not_tried_when_memory_stopped_the_doubling_earlier(): + ran: list[int] = [] + assert find_batch_size(_fits_up_to(4, ran), limit=24, candidate=24) == 4 + assert 24 not in ran, 'if 16 does not fit, 24 cannot' + + +def test_a_candidate_above_the_limit_is_ignored(): + ran: list[int] = [] + assert find_batch_size(_fits_up_to(64, ran), limit=16, candidate=24) == 16 + assert 24 not in ran + + +def test_a_candidate_that_is_already_a_power_of_two_changes_nothing(): + ran: list[int] = [] + assert find_batch_size(_fits_up_to(64, ran), limit=16, candidate=16) == 16 + assert ran.count(16) == 1, 'the doubling already tried it' + + def test_a_machine_that_cannot_take_one_sample_is_an_error(): fits, calls = _recording(0) with pytest.raises(RuntimeError, match='batch size 1 does not fit'): @@ -60,6 +114,29 @@ def test_batch_count_covers_the_whole_set_without_overshooting_by_a_batch(): assert covered - sample_count < batch_size +@pytest.mark.parametrize('empty', [{}, {REQUESTED_BATCH_SIZE: None}, {REQUESTED_BATCH_SIZE: ''}, + {REQUESTED_BATCH_SIZE: 0}]) +def test_an_unfilled_batch_size_means_no_bound(empty: dict): + assert requested_batch_size(empty) == 0 + + +def test_the_reported_batch_size_is_never_read_as_a_bound(): + """A trainer reports its settled size as `batch_size`; a resumed training must not read it as its bound.""" + assert REQUESTED_BATCH_SIZE != 'batch_size' + assert requested_batch_size({'batch_size': 16}) == 0 + assert requested_batch_size({'batch_size': 16, REQUESTED_BATCH_SIZE: 64}) == 64 + + +@pytest.mark.parametrize(('value', 'expected'), [('64', 64), (64, 64), (48.0, 48)]) +def test_a_batch_size_is_read_whatever_the_loop_spelled_it_as(value: object, expected: int): + assert requested_batch_size({REQUESTED_BATCH_SIZE: value}) == expected + + +def test_a_batch_size_that_is_not_a_number_is_an_error(): + with pytest.raises(ValueError): + requested_batch_size({REQUESTED_BATCH_SIZE: 'lots'}) + + def test_the_dataset_limit_keeps_enough_steps_per_epoch(): for sample_count in (8, 20, 47, 100, 1000, 118_000): limit = smaller_pot(dataset_limit(sample_count)) @@ -110,6 +187,15 @@ def fits(batch_size: int) -> bool: return fits, calls -def _fits_up_to(capacity: int) -> Callable[[int], bool]: - fits, _ = _recording(capacity) - return fits +def _fits_up_to(capacity: int, ran: list[int] | None = None) -> Callable[[int], bool]: + """A `fits` predicate for a machine of `capacity`, appending every size asked about to `ran`.""" + fits, calls = _recording(capacity) + if ran is None: + return fits + + def recording_fits(batch_size: int) -> bool: + result = fits(batch_size) + ran[:] = calls + return result + + return recording_fits diff --git a/learning_loop_node/tests/unit/test_cuda.py b/learning_loop_node/tests/unit/test_cuda.py index c700f3f2..172ac30e 100644 --- a/learning_loop_node/tests/unit/test_cuda.py +++ b/learning_loop_node/tests/unit/test_cuda.py @@ -6,6 +6,7 @@ """ from __future__ import annotations +import argparse import importlib import logging import sys @@ -59,6 +60,14 @@ def test_nothing_is_capped_without_cuda(load): assert fake.capped == [] +def test_a_spawned_script_takes_the_same_budget_flag_as_its_node(load): + cuda, _ = load() + parser = argparse.ArgumentParser() + cuda.add_vram_limit_argument(parser) + assert parser.parse_args([]).vram_limit_gb == 0 + assert parser.parse_args(['--vram-limit-gb', '6.5']).vram_limit_gb == 6.5 + + def test_a_limit_the_card_cannot_reach_warns_instead_of_capping(load, caplog): cuda, fake = load(total_gb=8.0) with caplog.at_level(logging.WARNING): @@ -133,6 +142,173 @@ def test_a_probe_without_a_gpu_does_not_run_the_step(load): assert not fake.allocated, 'and no margin claimed on a card that is not there' +# --- settling a batch size from the hyperparameters --- + +def test_a_requested_size_that_fits_is_used_as_named(load): + cuda, fake = load() + ran: list[int] = [] + assert cuda.measure_batch_size(_fits_up_to(64, fake, ran), batch_size=16) == 16 + assert ran[-1] == 16, 'and it was measured, not taken on trust' + + +def test_a_requested_size_that_does_not_fit_backs_off_instead_of_running_out_of_memory(load): + cuda, fake = load() + assert cuda.measure_batch_size(_fits_up_to(8, fake, []), batch_size=32) == 8 + + +def test_a_requested_size_is_tried_even_when_it_is_not_a_power_of_two(load): + cuda, fake = load() + ran: list[int] = [] + assert cuda.measure_batch_size(_fits_up_to(24, fake, ran), batch_size=24) == 24 + assert ran == [1, 2, 4, 8, 16, 24], 'the named size only after the search reached its ceiling' + + +def test_a_named_size_that_does_not_fit_falls_back_to_the_power_of_two_below_it(load): + cuda, fake = load() + ran: list[int] = [] + assert cuda.measure_batch_size(_fits_up_to(16, fake, ran), batch_size=24) == 16 + assert ran[-1] == 24, 'it was tried, and it did not fit' + + +def test_a_named_size_is_not_tried_when_memory_stopped_the_search_earlier(load): + cuda, fake = load() + ran: list[int] = [] + assert cuda.measure_batch_size(_fits_up_to(4, fake, ran), batch_size=24) == 4 + assert 24 not in ran, 'if 16 does not fit, 24 cannot' + + +def test_an_absent_batch_size_leaves_the_bound_to_the_library(load): + cuda, fake = load() + assert cuda.measure_batch_size(_fits_up_to(2048, fake, [])) == MAX_BATCH_SIZE + + +def test_a_batch_size_of_zero_means_measure(load): + cuda, fake = load() + ran: list[int] = [] + assert cuda.measure_batch_size(_fits_up_to(16, fake, ran), batch_size=0) == 16 + assert ran + + +def test_the_settled_size_is_only_returned(load): + cuda, fake = load() + settled = cuda.measure_batch_size(_fits_up_to(16, fake, []), batch_size=64) + assert settled == 16, 'the caller reports it; nothing here stores it where a second call would read it' + + +def test_the_dataset_bounds_the_search_as_well(load): + cuda, fake = load() + # 80 samples leave room for 10 per step, rounded down to a power of two + assert cuda.measure_batch_size(_fits_up_to(1024, fake, []), sample_count=80) == 8 + + +def test_a_dataset_bound_is_never_used_as_a_candidate(load): + cuda, fake = load() + ran: list[int] = [] + cuda.measure_batch_size(_fits_up_to(1024, fake, ran), sample_count=80) + assert 10 not in ran, 'samples // 8 is a heuristic, not a size anyone asked for' + + +def test_the_log_says_when_the_dataset_is_what_bounds_the_search(load, caplog): + cuda, fake = load() + with caplog.at_level(logging.INFO): + cuda.measure_batch_size(_fits_up_to(1024, fake, []), sample_count=80) + assert '80 training samples allow at most 10 per batch' in caplog.text + + +def test_the_tighter_of_the_request_and_the_dataset_wins(load): + cuda, fake = load() + assert cuda.measure_batch_size(_fits_up_to(1024, fake, []), batch_size=4, + sample_count=8000) == 4 + assert cuda.measure_batch_size(_fits_up_to(1024, fake, []), batch_size=512, + sample_count=80) == 8 + + +def test_memory_still_decides_below_both_bounds(load): + cuda, fake = load() + assert cuda.measure_batch_size(_fits_up_to(2, fake, []), batch_size=64, + sample_count=8000) == 2 + + +def test_a_negative_batch_size_is_a_mistake_not_a_sentinel(load): + cuda, fake = load() + with pytest.raises(ValueError, match='batch_size'): + cuda.measure_batch_size(_fits_up_to(64, fake, []), batch_size=-1) + + +def test_the_minimum_reaches_the_probe(load): + cuda, fake = load() + ran: list[int] = [] + cuda.measure_batch_size(_fits_up_to(64, fake, ran), minimum=4) + assert ran[0] == 4 + + +# --- a minimum above one --- + +def test_a_minimum_keeps_the_search_off_the_sizes_below_it(load): + cuda, fake = load() + ran: list[int] = [] + assert cuda.probe_batch_size(_fits_up_to(16, fake, ran), limit=64, minimum=2) == 16 + assert ran == [2, 4, 8, 16, 32], 'the smallest trial is the minimum, not one' + + +def test_a_minimum_is_rounded_down_to_a_power_of_two(load): + cuda, fake = load() + ran: list[int] = [] + cuda.probe_batch_size(_fits_up_to(64, fake, ran), limit=64, minimum=3) + assert ran[0] == 2 + + +def test_a_minimum_of_one_searches_exactly_as_before(load): + cuda, fake = load() + ran: list[int] = [] + assert cuda.probe_batch_size(_fits_up_to(16, fake, ran), limit=64, minimum=1) == 16 + assert ran == [1, 2, 4, 8, 16, 32] + + +def test_a_minimum_that_does_not_fit_reports_its_own_size(load): + cuda, fake = load() + ran: list[int] = [] + with pytest.raises(InsufficientMemoryError, match='batch size 4'): + cuda.probe_batch_size(_fits_up_to(2, fake, ran), limit=32, minimum=4) + + +def test_a_minimum_above_the_limit_still_gets_tried(load): + cuda, fake = load() + ran: list[int] = [] + assert cuda.probe_batch_size(_fits_up_to(64, fake, ran), limit=2, minimum=8) == 8 + assert ran == [8] + + +def test_without_a_gpu_the_fallback_respects_the_minimum(load): + cuda, fake = load(cuda_available=False) + ran: list[int] = [] + assert cuda.probe_batch_size(_fits_up_to(1024, fake, ran), limit=64, minimum=16) == 16 + assert not ran + + +# --- cleaning up after a trial that did not fit --- + +def test_the_out_of_memory_hook_runs_after_every_failed_trial(load): + cuda, fake = load() + ran: list[int] = [] + dropped: list[int] = [] + cuda.probe_batch_size(_fits_up_to(4, fake, ran), limit=32, + on_out_of_memory=lambda: dropped.append(len(ran))) + assert dropped == [4], 'once, after the single trial that went over' + + +def test_the_out_of_memory_hook_does_not_run_for_a_bug(load): + cuda, _ = load() + dropped: list[int] = [] + + def run_batch(_: int) -> None: + raise RuntimeError('a real bug') + + with pytest.raises(RuntimeError, match='a real bug'): + cuda.probe_batch_size(run_batch, limit=32, on_out_of_memory=lambda: dropped.append(1)) + assert not dropped + + def test_a_batch_size_of_one_that_does_not_fit_is_an_error(load): cuda, fake = load() ran: list[int] = [] diff --git a/learning_loop_node/tests/unit/test_entrypoint.py b/learning_loop_node/tests/unit/test_entrypoint.py index 9eed399a..21d4dade 100644 --- a/learning_loop_node/tests/unit/test_entrypoint.py +++ b/learning_loop_node/tests/unit/test_entrypoint.py @@ -3,7 +3,8 @@ from ...helpers.entrypoint import node_parser MANAGED = ('WEIGHT_TYPE', 'MY_DETECTOR_WEIGHT_TYPE', 'HOST', 'NODE_HOST', 'NODE_PORT', 'PORT', - 'MY_DETECTOR_HOST', 'MY_DETECTOR_PORT', 'MY_DETECTOR_NODE_HOST') + 'MY_DETECTOR_HOST', 'MY_DETECTOR_PORT', 'MY_DETECTOR_NODE_HOST', 'VRAM_LIMIT_GB', + 'MY_DETECTOR_VRAM_LIMIT_GB') def test_every_node_gets_a_host_and_a_port(): @@ -92,6 +93,30 @@ def test_the_loop_own_host_is_not_adopted_by_a_prefixed_node(monkeypatch: pytest assert _parser(legacy_env_prefix='MY_DETECTOR_').parse_args([]).host == '0.0.0.0' +def test_a_node_that_does_not_probe_has_no_vram_limit(): + assert not hasattr(_parser().parse_args([]), 'vram_limit_gb') + + +def test_the_vram_limit_defaults_to_the_whole_card(): + assert _parser(vram_limit=True).parse_args([]).vram_limit_gb == 0 + + +def test_the_vram_limit_is_read_from_its_variable(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv('VRAM_LIMIT_GB', '6') + assert _parser(vram_limit=True).parse_args([]).vram_limit_gb == 6.0 + + +def test_the_vram_limit_flag_beats_its_variable(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv('VRAM_LIMIT_GB', '6') + assert _parser(vram_limit=True).parse_args(['--vram-limit-gb', '4.5']).vram_limit_gb == 4.5 + + +def test_the_vram_limit_is_still_read_under_a_legacy_prefix(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv('MY_DETECTOR_VRAM_LIMIT_GB', '6') + args = _parser(vram_limit=True, legacy_env_prefix='MY_DETECTOR_').parse_args([]) + assert args.vram_limit_gb == 6.0 + + @pytest.fixture(autouse=True) def clean_env(monkeypatch: pytest.MonkeyPatch): """Every test starts without the variables it is about to set.""" diff --git a/learning_loop_node/trainer/batch_size.py b/learning_loop_node/trainer/batch_size.py index 8fb0418c..c921f27c 100644 --- a/learning_loop_node/trainer/batch_size.py +++ b/learning_loop_node/trainer/batch_size.py @@ -10,12 +10,29 @@ """ import logging -from collections.abc import Callable +from collections.abc import Callable, Mapping +from typing import Any from .exceptions import InsufficientMemoryError logger = logging.getLogger(__name__) +REQUESTED_BATCH_SIZE = 'max_batch_size' +"""The hyperparameter every node reads its `measure_batch_size` argument out of. + +0 or absent means the card decides. It is an input only: a trainer reports the size it settled on +under a different key, conventionally ``batch_size``, and never writes it back here. +""" + +VRAM_LIMIT_GB_FLAG = '--vram-limit-gb' +"""The flag a trainer takes its GPU budget from; ``VRAM_LIMIT_GB`` follows from the name.""" + +VRAM_LIMIT_GB_HELP = ('Gigabytes of GPU memory a training may use. The batch size is probed against this limit ' + 'instead of the whole card, so a lower limit yields a smaller batch size rather than an ' + 'out-of-memory error. Use it to share a GPU or to keep headroom against fragmentation. ' + "The limit is relative to the card's total memory, not to what is currently free. " + '0 (default) means no limit.') + MAX_BATCH_SIZE = 1024 """Where a search stops when its caller sets no bound of its own.""" @@ -26,23 +43,51 @@ """Fewest optimizer steps an epoch must have; :func:`dataset_limit` is derived from it.""" -def find_batch_size(fits: Callable[[int], bool], *, limit: int) -> int: - """Return the largest power-of-two batch size that fits, never exceeding ``limit``. +def find_batch_size(fits: Callable[[int], bool], *, limit: int, minimum: int = 1, + candidate: int = 0) -> int: + """Return the largest batch size that fits, never exceeding ``limit``. + + Powers of two, plus ``candidate``, which is the one size outside that set this will return, and + only when somebody named it and it then measured. :param fits: Runs a representative probe; ``False`` on out-of-memory. - :raises InsufficientMemoryError: If not even a batch size of 1 fits. + :param limit: Upper bound; the doubling stops at the largest power of two within it. + :param minimum: Smallest size to try, rounded down to a power of two. Raise it above one for a + step that cannot run on a single sample at all: BatchNorm over a 1x1 feature map, a + validation pass that halves the batch. + :param candidate: An exact size, tried once the doubling has reached its ceiling, so a size + that was asked for is used as asked for rather than rounded down. Ignored unless it lies + between that ceiling and ``limit``; a bound nobody named — one derived from the dataset, + say — must not be passed here. + :raises InsufficientMemoryError: If not even ``minimum`` fits. """ - limit = smaller_pot(limit) - if not fits(1): - raise InsufficientMemoryError('batch size 1 does not fit in memory') + minimum = smaller_pot(max(1, minimum)) + bound = max(limit, minimum) + ceiling = max(smaller_pot(bound), minimum) + + if not fits(minimum): + raise InsufficientMemoryError(f'batch size {minimum} does not fit in memory') - size = 1 - while size < limit and fits(size * 2): + size = minimum + while size < ceiling and fits(size * 2): size *= 2 + if size == ceiling and ceiling < candidate <= bound and fits(candidate): + size = candidate + return size +def requested_batch_size(hyperparameters: Mapping[str, Any]) -> int: + """The bound a training asked for, read out of the hyperparameters the loop sent. + + An unfilled field arrives as absent, ``None`` or ``''``, and all three mean no bound. + + :raises ValueError: If the value is there but is not a number. + """ + return int(hyperparameters.get(REQUESTED_BATCH_SIZE, 0) or 0) + + def dataset_limit(sample_count: int) -> int: """The batch size ceiling that still leaves ``MIN_TRAIN_STEPS_PER_EPOCH`` steps per epoch.""" if sample_count < 1: diff --git a/learning_loop_node/trainer/cuda.py b/learning_loop_node/trainer/cuda.py index 1b126bb0..f217bcaa 100644 --- a/learning_loop_node/trainer/cuda.py +++ b/learning_loop_node/trainer/cuda.py @@ -11,11 +11,22 @@ import gc import logging +from argparse import ArgumentParser from collections.abc import Callable import torch -from .batch_size import MAX_BATCH_SIZE, find_batch_size, is_out_of_memory, no_gpu_batch_size, smaller_pot +from .batch_size import ( + MAX_BATCH_SIZE, + REQUESTED_BATCH_SIZE, + VRAM_LIMIT_GB_FLAG, + VRAM_LIMIT_GB_HELP, + dataset_limit, + find_batch_size, + is_out_of_memory, + no_gpu_batch_size, + smaller_pot, +) logger = logging.getLogger(__name__) @@ -23,32 +34,80 @@ """Share of the budget held back while probing, against allocator fragmentation later on.""" +def measure_batch_size(run_batch: Callable[[int], str | None], *, batch_size: int = 0, + sample_count: int | None = None, probe: str = 'batch-size probe', + minimum: int = 1, vram_limit_gb: float = 0, + on_out_of_memory: Callable[[], None] | None = None) -> int: + """Settle a training's batch size against what it asked for and what the card allows. + + Every training enters here, whatever shape its hyperparameters have; :func:`probe_batch_size` is + for a probe with no requested size to honour, such as a detection pass bounded only by how many + images there are. + + :param run_batch: Runs the batch; may return a detail to append to the log line. + :param batch_size: What the training asked for, as carried in the + :data:`~.batch_size.REQUESTED_BATCH_SIZE` hyperparameter: the largest batch it may use, + measured rather than trusted. A size that fits is used as asked for, whether or not it is a + power of two; one that does not becomes the largest power of two below it that does. 0 means + the card decides alone. + :param sample_count: Samples in the training split, when the caller knows it. The search is + then bounded by :func:`~.batch_size.dataset_limit`; that bound is never used as the exact + candidate. + :param minimum: Smallest size to try; see :func:`probe_batch_size`. + :param vram_limit_gb: The budget the safety margin is a share of; 0 means the whole card. + :param on_out_of_memory: Runs after a trial ran out of memory, to drop what it left behind. + :raises InsufficientMemoryError: If not even ``minimum`` fits. + :raises ValueError: If the training asked for a negative batch size. + """ + if batch_size < 0: + raise ValueError(f'{REQUESTED_BATCH_SIZE} must be >= 0, got {batch_size}') + + limit = batch_size + if sample_count is not None: + limit = min(limit or MAX_BATCH_SIZE, dataset_limit(sample_count)) + logger.info('%s: %d training samples allow at most %d per batch', probe, sample_count, + dataset_limit(sample_count)) + + return probe_batch_size(run_batch, probe=probe, limit=limit, candidate=batch_size, + minimum=minimum, vram_limit_gb=vram_limit_gb, + on_out_of_memory=on_out_of_memory) + + def probe_batch_size(run_batch: Callable[[int], str | None], *, probe: str = 'batch-size probe', - limit: int = 0, vram_limit_gb: float = 0) -> int: - """Find the largest power-of-two batch size ``run_batch`` fits into. + limit: int = 0, candidate: int = 0, minimum: int = 1, vram_limit_gb: float = 0, + on_out_of_memory: Callable[[], None] | None = None) -> int: + """Run :func:`~learning_loop_node.trainer.batch_size.find_batch_size` against a real card. - For a probe whose measurement is one call. A probe that has to build a throwaway model first - composes :func:`reserve_margin`, :func:`measured_fits` and ``find_batch_size`` itself. + This is the whole of a probe except the step itself: the margin, the search, telling an + out-of-memory failure from a bug, and releasing what the trials left behind. A caller that + builds a throwaway model supplies ``on_out_of_memory`` to drop what a failed trial left on the + card. :param run_batch: Runs the batch; may return a detail to append to the log line. :param probe: Names this probe in the log, so a node running several stays readable. - :param limit: Caps the search, rounded down to a power of two; 0 means + :param limit: Caps the search; 0 means :data:`~learning_loop_node.trainer.batch_size.MAX_BATCH_SIZE`. + :param candidate: An exact size to try once the doubling has reached its ceiling; see + ``find_batch_size``. + :param minimum: Smallest size to try; see ``find_batch_size``. :param vram_limit_gb: The budget the safety margin is a share of; 0 means the whole card. - :raises InsufficientMemoryError: If not even a batch size of 1 fits. + :param on_out_of_memory: Runs after a trial ran out of memory, to drop what it left behind + (an optimizer's gradients, say). + :raises InsufficientMemoryError: If not even ``minimum`` fits. """ - limit = smaller_pot(limit or MAX_BATCH_SIZE) + bound = max(limit or MAX_BATCH_SIZE, minimum) if not torch.cuda.is_available(): - return no_gpu_batch_size(limit, probe) + return max(smaller_pot(max(1, minimum)), no_gpu_batch_size(bound, probe)) margin = reserve_margin(vram_limit_gb, probe=probe) try: - chosen = find_batch_size(measured_fits(run_batch, probe=probe), limit=limit) + fits = measured_fits(run_batch, probe=probe, on_out_of_memory=on_out_of_memory) + chosen = find_batch_size(fits, limit=bound, minimum=minimum, candidate=candidate) finally: del margin free_cuda_memory() - logger.info('%s: selected batch size %d (upper bound %d)', probe, chosen, limit) + logger.info('%s: selected batch size %d (upper bound %d)', probe, chosen, bound) return chosen @@ -108,6 +167,14 @@ def usable_memory_bytes(vram_limit_gb: float) -> int: return min(total_bytes, int(vram_limit_gb * 1024**3)) +def add_vram_limit_argument(parser: ArgumentParser) -> None: + """Give a spawned training script the same GPU budget flag its node has. + + The spawned process still has to call :func:`limit_cuda_memory` with it. + """ + parser.add_argument(VRAM_LIMIT_GB_FLAG, type=float, default=0, help=VRAM_LIMIT_GB_HELP) + + def limit_cuda_memory(vram_limit_gb: float) -> None: """Cap how much of the GPU this process may allocate, to ``vram_limit_gb`` gigabytes.