Skip to content
Merged
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
99 changes: 64 additions & 35 deletions tests/unit/_utils/test_system.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,10 +2,10 @@

import logging
import sys
from multiprocessing import get_context, synchronize
from multiprocessing import get_context, resource_tracker, synchronize
from multiprocessing.shared_memory import SharedMemory
from types import SimpleNamespace
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, cast
from unittest.mock import Mock

import proclimits
Expand All @@ -19,6 +19,7 @@

if TYPE_CHECKING:
from collections.abc import Callable
from multiprocessing.context import ForkContext, ForkServerContext, SpawnContext

HOST_TOTAL_BYTES = 8 * 1024**3
HOST_AVAILABLE_BYTES = 3 * 1024**3
Expand Down Expand Up @@ -424,11 +425,51 @@ def test_log_resource_limits_lets_a_failing_sensor_surface(monkeypatch: pytest.M
snapshot.assert_called_once()


EXTRA_MEMORY_SIZE = 1024 * 1024 * 100 # 100 MB


# The children of the estimation test below live at module level, so that every start method can pickle them.
def no_extra_memory_child(ready: synchronize.Barrier, measured: synchronize.Barrier) -> None:
ready.wait()
measured.wait()


def extra_memory_child(ready: synchronize.Barrier, measured: synchronize.Barrier) -> None:
memory = SharedMemory(size=EXTRA_MEMORY_SIZE, create=True)
assert memory.buf is not None
fill_buffer(memory.buf, EXTRA_MEMORY_SIZE)
print(f'Using the memory... {memory.buf[-1]}')
ready.wait()
measured.wait()
memory.close()
memory.unlink()


def shared_extra_memory_child(ready: synchronize.Barrier, measured: synchronize.Barrier, memory: SharedMemory) -> None:
assert memory.buf is not None
# Fault every page in: untouched pages never enter the RSS (hiding the overcount this test guards
# against) and are reclaimed first under memory pressure, which drops them from every mapper's PSS.
page_sum = sum(memory.buf[::4096])
print(f'Using the memory... {page_sum}')
ready.wait()
measured.wait()


# The estimation is asserted on absolute memory readings, which hold only as long as nothing else on the machine makes
# the kernel reclaim the pages allocated below. Running alongside the other test workers is enough to break that.
@pytest.mark.run_alone
@pytest.mark.skipif(sys.platform != 'linux', reason='Improved estimation available only on Linux')
def test_memory_estimation_does_not_overestimate_due_to_shared_memory() -> None:
# The start methods differ in how much memory the children share with the rest of the process tree, which is what the
# estimation has to account for. The default is `fork` up to Python 3.13 and `forkserver` from 3.14 on.
@pytest.mark.parametrize(
'start_method',
[
pytest.param('fork', id='fork'),
pytest.param('forkserver', id='forkserver'),
pytest.param('spawn', id='spawn'),
],
)
def test_memory_estimation_does_not_overestimate_due_to_shared_memory(start_method: str) -> None:
"""Test that memory usage estimation is not overestimating memory usage by counting shared memory multiple times.

In this test, the parent process is started and its memory usage is measured in situations where it is running
Expand All @@ -440,41 +481,17 @@ def test_memory_estimation_does_not_overestimate_due_to_shared_memory() -> None:
the same as the unshared memory.
"""

ctx = get_context('fork')
estimated_memory_expectation = ctx.Value('b', False) # noqa: FBT003 # Common usage pattern for multiprocessing.Value
# The measuring process is a closure, which only `fork` can start. The children are started by `start_method`.
fork_ctx = get_context('fork')
estimated_memory_expectation = fork_ctx.Value('b', False) # noqa: FBT003 # Common usage pattern for multiprocessing.Value

def parent_process() -> None:
extra_memory_size = 1024 * 1024 * 100 # 100 MB
ctx = cast('ForkContext | ForkServerContext | SpawnContext', get_context(start_method))
children_count = 4
# Memory calculation is not exact, so allow for some tolerance.
test_tolerance = 0.3
measurement_rounds = 3

def no_extra_memory_child(ready: synchronize.Barrier, measured: synchronize.Barrier) -> None:
ready.wait()
measured.wait()

def extra_memory_child(ready: synchronize.Barrier, measured: synchronize.Barrier) -> None:
memory = SharedMemory(size=extra_memory_size, create=True)
assert memory.buf is not None
fill_buffer(memory.buf, extra_memory_size)
print(f'Using the memory... {memory.buf[-1]}')
ready.wait()
measured.wait()
memory.close()
memory.unlink()

def shared_extra_memory_child(
ready: synchronize.Barrier, measured: synchronize.Barrier, memory: SharedMemory
) -> None:
assert memory.buf is not None
# Fault every page in: untouched pages never enter the RSS (hiding the overcount this test guards
# against) and are reclaimed first under memory pressure, which drops them from every mapper's PSS.
page_sum = sum(memory.buf[::4096])
print(f'Using the memory... {page_sum}')
ready.wait()
measured.wait()

def get_additional_memory_estimation_while_running_processes(
*, target: Callable, count: int = 1, use_shared_memory: bool = False
) -> float:
Expand All @@ -485,9 +502,9 @@ def get_additional_memory_estimation_while_running_processes(
memory_before = get_memory_info().current_size

if use_shared_memory:
shared_memory = SharedMemory(size=extra_memory_size, create=True)
shared_memory = SharedMemory(size=EXTRA_MEMORY_SIZE, create=True)
assert shared_memory.buf is not None
fill_buffer(shared_memory.buf, extra_memory_size)
fill_buffer(shared_memory.buf, EXTRA_MEMORY_SIZE)
extra_args = [shared_memory]
else:
extra_args = []
Expand All @@ -510,6 +527,18 @@ def get_additional_memory_estimation_while_running_processes(

return (memory_during - memory_before).to_mb() / count

# Some start methods launch long-lived helper processes (the fork server, the resource tracker) on first use.
# They belong to the process tree and so to the estimate, but started inside a round they would inflate its
# baseline. Start them ahead of the measurements, keeping the barriers referenced until the child is done. The
# resource tracker is started explicitly: under `fork` no barrier registers with it, and every child creating
# shared memory would start one of its own.
resource_tracker.ensure_running()
ready, measured = ctx.Barrier(parties=1), ctx.Barrier(parties=1)
warm_up = ctx.Process(target=no_extra_memory_child, args=[ready, measured])
warm_up.start()
warm_up.join()
assert warm_up.exitcode == 0

# Under memory pressure the kernel reclaims cold pages, which silently leave the PSS readings and skew a
# round's deltas, so a distorted round is re-measured. A genuine overcount of shared memory misses the
# expectation several times over in every round, so the retries cannot mask it.
Expand Down Expand Up @@ -549,10 +578,10 @@ def get_additional_memory_estimation_while_running_processes(
f'{memory_estimation_difference_ratio=}'
)

process = ctx.Process(target=parent_process)
process = fork_ctx.Process(target=parent_process)
process.start()
process.join()

assert estimated_memory_expectation.value, (
'Estimated memory usage for process with shared memory does not meet the expectation.'
f'Estimated memory usage for process with shared memory does not meet the expectation under {start_method}.'
)
Loading