Skip to content
Open
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
39 changes: 32 additions & 7 deletions AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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) |
Expand Down
8 changes: 7 additions & 1 deletion learning_loop_node/helpers/entrypoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
92 changes: 89 additions & 3 deletions learning_loop_node/tests/unit/test_batch_size.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand All @@ -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'):
Expand Down Expand Up @@ -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))
Expand Down Expand Up @@ -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
Loading
Loading