From 2d41212e9e8ecdb62115f6661db154b1287e04c2 Mon Sep 17 00:00:00 2001 From: Yuge Zhang Date: Wed, 15 Oct 2025 23:15:52 -0700 Subject: [PATCH 1/6] Streamline in-memory span eviction thresholds --- agentlightning/store/memory.py | 145 ++++++++++++++++++++++++++- tests/store/test_memory.py | 175 ++++++++++++++++++++++++++++++++- 2 files changed, 317 insertions(+), 3 deletions(-) diff --git a/agentlightning/store/memory.py b/agentlightning/store/memory.py index ef06de01f..04e38b583 100644 --- a/agentlightning/store/memory.py +++ b/agentlightning/store/memory.py @@ -5,14 +5,17 @@ import asyncio import functools import hashlib +import importlib import logging +import sys import threading import time import uuid from collections import deque -from typing import Any, Callable, Counter, Dict, List, Literal, Optional, Sequence, TypeVar, cast +from typing import Any, Callable, Counter, Dict, List, Literal, Optional, Sequence, Set, TypeVar, cast from opentelemetry.sdk.trace import ReadableSpan +from pydantic import BaseModel from agentlightning.types import ( Attempt, @@ -35,6 +38,18 @@ logger = logging.getLogger(__name__) +def estimate_model_size(obj) -> int: + """Rough recursive size estimate for Pydantic BaseModel instances.""" + + if isinstance(obj, BaseModel): + return sum(estimate_model_size(v) for v in obj.__dict__.values()) + sys.getsizeof(obj) + if isinstance(obj, dict): + return sum(estimate_model_size(v) for v in obj.values()) + sys.getsizeof(obj) + if isinstance(obj, (list, tuple, set)): + return sum(estimate_model_size(v) for v in obj) + sys.getsizeof(obj) + return sys.getsizeof(obj) + + def _healthcheck_wrapper(func: T_callable) -> T_callable: """ Decorator to run the watchdog healthcheck **before** executing the decorated method. @@ -82,6 +97,18 @@ def _generate_attempt_id() -> str: return "at-" + short_id +def _detect_total_memory_bytes() -> int: + """Best-effort detection of the total available system memory in bytes.""" + + psutil_spec = importlib.util.find_spec("psutil") + if psutil_spec is not None: + psutil = importlib.import_module("psutil") + return int(psutil.virtual_memory().total) + + # Fallback to 8GB if memory cannot be detected. + return 8 * 1024**3 + + class InMemoryLightningStore(LightningStore): """ In-memory implementation of LightningStore using Python data structures. @@ -91,7 +118,13 @@ class InMemoryLightningStore(LightningStore): especially those that are locked. """ - def __init__(self): + def __init__( + self, + *, + eviction_memory_threshold: float | int | None = None, + safe_memory_threshold: float | int | None = None, + span_size_estimator: Callable[[Span], int] | None = None, + ): self._lock = asyncio.Lock() # Task queue and rollouts storage @@ -105,6 +138,36 @@ def __init__(self): # Spans storage self._spans: Dict[str, List[Span]] = {} # rollout_id -> list of spans self._span_sequence_ids: Dict[str, int] = Counter() # rollout_id -> sequence_id + self._span_bytes_by_rollout: Dict[str, int] = Counter() + self._total_span_bytes: int = 0 + self._evicted_rollout_span_sets: Set[str] = set() + + self._memory_capacity_bytes = _detect_total_memory_bytes() + if self._memory_capacity_bytes <= 0: + raise ValueError("Detected memory capacity must be positive") + + self._eviction_threshold_bytes = self._resolve_memory_threshold( + eviction_memory_threshold, + default_ratio=0.7, + capacity_bytes=self._memory_capacity_bytes, + name="eviction_memory_threshold", + minimum=1, + ) + + if safe_memory_threshold is None: + safe_memory_threshold = max(int(self._eviction_threshold_bytes * 0.8), 0) + + self._safe_threshold_bytes = self._resolve_memory_threshold( + safe_memory_threshold, + default_ratio=self._eviction_threshold_bytes / self._memory_capacity_bytes, + capacity_bytes=self._memory_capacity_bytes, + name="safe_memory_threshold", + minimum=0, + ) + + if not (0 <= self._safe_threshold_bytes < self._eviction_threshold_bytes): + raise ValueError("safe_memory_threshold must be smaller than eviction_memory_threshold") + self._custom_span_size_estimator = span_size_estimator # Attempt tracking self._attempts: Dict[str, List[Attempt]] = {} # rollout_id -> list of attempts @@ -416,6 +479,8 @@ async def _add_span_unlocked(self, span: Span) -> Span: if span.rollout_id not in self._spans: self._spans[span.rollout_id] = [] self._spans[span.rollout_id].append(span) + self._account_span_size(span) + self._maybe_evict_spans() # Update attempt heartbeat current_attempt.last_heartbeat_time = time.time() @@ -439,6 +504,80 @@ async def _add_span_unlocked(self, span: Span) -> Span: return span + @staticmethod + def _resolve_memory_threshold( + value: float | int | None, + *, + default_ratio: float, + capacity_bytes: int, + name: str, + minimum: int, + ) -> int: + if value is None: + resolved = int(capacity_bytes * default_ratio) + elif isinstance(value, float): + if minimum == 0: + if not (0 <= value <= 1): + raise ValueError(f"{name} ratio must be between 0 and 1 inclusive") + else: + if not (0 < value <= 1): + raise ValueError(f"{name} ratio must be greater than 0 and at most 1") + resolved = int(capacity_bytes * value) + elif isinstance(value, int): + if value < 0: + raise ValueError(f"{name} must be non-negative") + resolved = value + else: + raise TypeError(f"{name} must be a float ratio or int bytes") + + if resolved < minimum: + raise ValueError(f"{name} must be at least {minimum} bytes") + + return resolved + + def _account_span_size(self, span: Span) -> int: + if self._custom_span_size_estimator is not None: + size = max(int(self._custom_span_size_estimator(span)), 0) + else: + size = estimate_model_size(span) + + self._span_bytes_by_rollout[span.rollout_id] += size + self._total_span_bytes += size + return size + + def _maybe_evict_spans(self) -> None: + if self._total_span_bytes <= self._eviction_threshold_bytes: + return + + candidates = sorted( + ( + ( + ( + self._rollouts.get(rollout_id).start_time + if self._rollouts.get(rollout_id) + else spans[0].start_time or 0.0 + ), + rollout_id, + ) + for rollout_id, spans in self._spans.items() + if spans + ), + key=lambda item: item[0], + ) + + for _, rollout_id in candidates: + if self._total_span_bytes <= self._safe_threshold_bytes: + break + self._evict_spans_for_rollout(rollout_id) + + def _evict_spans_for_rollout(self, rollout_id: str) -> None: + spans = self._spans.pop(rollout_id, []) + if not spans: + return + removed_bytes = self._span_bytes_by_rollout.pop(rollout_id, 0) + self._total_span_bytes = max(self._total_span_bytes - removed_bytes, 0) + self._evicted_rollout_span_sets.add(rollout_id) + @_healthcheck_wrapper async def wait_for_rollouts(self, *, rollout_ids: List[str], timeout: Optional[float] = None) -> List[Rollout]: """ @@ -500,6 +639,8 @@ async def query_spans(self, rollout_id: str, attempt_id: str | Literal["latest"] Returns an empty list if no spans are found. """ async with self._lock: + if rollout_id in self._evicted_rollout_span_sets: + raise RuntimeError(f"Spans for rollout {rollout_id} have been evicted") spans = self._spans.get(rollout_id, []) if attempt_id is None: return spans diff --git a/tests/store/test_memory.py b/tests/store/test_memory.py index 137949731..c99b2403d 100644 --- a/tests/store/test_memory.py +++ b/tests/store/test_memory.py @@ -14,20 +14,27 @@ """ import asyncio +import sys import time from typing import List from unittest.mock import Mock import pytest +from pydantic import BaseModel -from agentlightning.store.memory import InMemoryLightningStore +from agentlightning.store.memory import InMemoryLightningStore, estimate_model_size from agentlightning.types import ( LLM, + Event, + Link, PromptTemplate, + Resource, ResourcesUpdate, Rollout, RolloutConfig, Span, + SpanContext, + TraceStatus, ) # Core CRUD Operations Tests @@ -572,6 +579,172 @@ async def test_query_spans_by_attempt(inmemory_store: InMemoryLightningStore, mo assert len(no_spans) == 0 +@pytest.mark.asyncio +async def test_span_eviction_removes_oldest_rollouts(mock_readable_span: Mock, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr("agentlightning.store.memory._detect_total_memory_bytes", lambda: 100) + store = InMemoryLightningStore( + eviction_memory_threshold=0.5, + safe_memory_threshold=0.05, + span_size_estimator=lambda span: 20, + ) + + attempted_rollouts = [] + for index in range(4): + attempted = await store.start_rollout(input={"index": index}) + attempted_rollouts.append(attempted) + await store.add_otel_span(attempted.rollout_id, attempted.attempt.attempt_id, mock_readable_span) + + for attempted in attempted_rollouts[:3]: + with pytest.raises(RuntimeError): + await store.query_spans(attempted.rollout_id) + + remaining_spans = await store.query_spans(attempted_rollouts[3].rollout_id) + assert len(remaining_spans) == 1 + assert remaining_spans[0].rollout_id == attempted_rollouts[3].rollout_id + + +def test_memory_threshold_accepts_byte_values() -> None: + store = InMemoryLightningStore( + eviction_memory_threshold=150, + safe_memory_threshold=20, + ) + + assert store._eviction_threshold_bytes == 150 + assert store._safe_threshold_bytes == 20 + + +def test_memory_threshold_accepts_ratios_with_zero_safe(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr("agentlightning.store.memory._detect_total_memory_bytes", lambda: 200) + store = InMemoryLightningStore( + eviction_memory_threshold=0.6, + safe_memory_threshold=0.0, + ) + + assert store._eviction_threshold_bytes == int(200 * 0.6) + assert store._safe_threshold_bytes == 0 + + +def test_invalid_safe_threshold_raises_value_error() -> None: + with pytest.raises(ValueError): + InMemoryLightningStore( + eviction_memory_threshold=50, + safe_memory_threshold=100, + ) + + +def test_estimate_model_size_counts_nested_models() -> None: + class Inner(BaseModel): + value: int + data: List[int] + + class Outer(BaseModel): + inner: Inner + mapping: dict[str, str] + tags: List[str] + + inner = Inner(value=7, data=[1, 2, 3]) + outer = Outer(inner=inner, mapping={"alpha": "beta"}, tags=["x", "yz"]) + + inner_expected = ( + sys.getsizeof(inner) + + sys.getsizeof(inner.value) + + sys.getsizeof(inner.data) + + sum(sys.getsizeof(item) for item in inner.data) + ) + assert estimate_model_size(inner) == inner_expected + + mapping_expected = sys.getsizeof(outer.mapping) + sum(sys.getsizeof(v) for v in outer.mapping.values()) + tags_expected = sys.getsizeof(outer.tags) + sum(sys.getsizeof(tag) for tag in outer.tags) + outer_expected = sys.getsizeof(outer) + inner_expected + mapping_expected + tags_expected + assert estimate_model_size(outer) == outer_expected + + +def test_estimate_model_size_handles_span_objects() -> None: + status = TraceStatus(status_code="OK", description="fine") + context = SpanContext(trace_id="trace", span_id="parent", is_remote=False, trace_state={}) + event = Event(name="step", attributes={"detail": "value"}, timestamp=1.0) + link = Link(context=context, attributes=None) + resource = Resource(attributes={"service.name": "unit"}, schema_url="schema") + + span = Span( + rollout_id="ro-1", + attempt_id="at-1", + sequence_id=1, + trace_id="trace", + span_id="span", + parent_id=None, + name="operation", + status=status, + attributes={"foo": "bar", "answer": 42}, + events=[event], + links=[link], + start_time=1.0, + end_time=2.0, + context=None, + parent=None, + resource=resource, + ) + + status_expected = sys.getsizeof(status) + sys.getsizeof(status.status_code) + sys.getsizeof(status.description) + + context_expected = ( + sys.getsizeof(context) + + sys.getsizeof(context.trace_id) + + sys.getsizeof(context.span_id) + + sys.getsizeof(context.is_remote) + + sys.getsizeof(context.trace_state) + + sum(sys.getsizeof(v) for v in context.trace_state.values()) + ) + + event_attributes_expected = sys.getsizeof(event.attributes) + sys.getsizeof("value") + event_expected = ( + sys.getsizeof(event) + sys.getsizeof(event.name) + event_attributes_expected + sys.getsizeof(event.timestamp) + ) + events_expected = sys.getsizeof(span.events) + event_expected + + if link.attributes is None: + link_attributes_expected = sys.getsizeof(None) + else: + link_attributes_expected = sys.getsizeof(link.attributes) + sum( + sys.getsizeof(v) for v in link.attributes.values() + ) + link_expected = sys.getsizeof(link) + context_expected + link_attributes_expected + links_expected = sys.getsizeof(span.links) + link_expected + + attributes_expected = ( + sys.getsizeof(span.attributes) + sys.getsizeof("bar") + sys.getsizeof(span.attributes["answer"]) + ) + + resource_expected = ( + sys.getsizeof(resource) + + sys.getsizeof(resource.attributes) + + sum(sys.getsizeof(v) for v in resource.attributes.values()) + + sys.getsizeof(resource.schema_url) + ) + + expected_size = ( + sys.getsizeof(span) + + sys.getsizeof(span.rollout_id) + + sys.getsizeof(span.attempt_id) + + sys.getsizeof(span.sequence_id) + + sys.getsizeof(span.trace_id) + + sys.getsizeof(span.span_id) + + sys.getsizeof(span.parent_id) + + sys.getsizeof(span.name) + + status_expected + + attributes_expected + + events_expected + + links_expected + + sys.getsizeof(span.start_time) + + sys.getsizeof(span.end_time) + + sys.getsizeof(span.context) + + sys.getsizeof(span.parent) + + resource_expected + ) + + assert estimate_model_size(span) == expected_size + + @pytest.mark.asyncio async def test_span_triggers_status_transition( inmemory_store: InMemoryLightningStore, mock_readable_span: Mock From 1f89907be4b857007a053ae0b449633104f21ef1 Mon Sep 17 00:00:00 2001 From: Yuge Zhang Date: Wed, 15 Oct 2025 23:39:38 -0700 Subject: [PATCH 2/6] Refine span memory accounting and typing --- agentlightning/store/memory.py | 66 ++++++++++++++++++++-------------- tests/store/test_memory.py | 27 +++++++------- 2 files changed, 53 insertions(+), 40 deletions(-) diff --git a/agentlightning/store/memory.py b/agentlightning/store/memory.py index 04e38b583..274e4a62c 100644 --- a/agentlightning/store/memory.py +++ b/agentlightning/store/memory.py @@ -6,13 +6,29 @@ import functools import hashlib import importlib +import importlib.util import logging import sys import threading import time import uuid from collections import deque -from typing import Any, Callable, Counter, Dict, List, Literal, Optional, Sequence, Set, TypeVar, cast +from collections.abc import Iterable +from collections.abc import Mapping as MappingABC +from typing import ( + Any, + Callable, + Counter, + Dict, + List, + Literal, + Mapping, + Optional, + Sequence, + Set, + TypeVar, + cast, +) from opentelemetry.sdk.trace import ReadableSpan from pydantic import BaseModel @@ -38,16 +54,19 @@ logger = logging.getLogger(__name__) -def estimate_model_size(obj) -> int: +def estimate_model_size(obj: Any) -> int: """Rough recursive size estimate for Pydantic BaseModel instances.""" if isinstance(obj, BaseModel): - return sum(estimate_model_size(v) for v in obj.__dict__.values()) + sys.getsizeof(obj) - if isinstance(obj, dict): - return sum(estimate_model_size(v) for v in obj.values()) + sys.getsizeof(obj) + values = cast(Iterable[Any], obj.__dict__.values()) + return sum(estimate_model_size(value) for value in values) + sys.getsizeof(cast(object, obj)) + if isinstance(obj, MappingABC): + mapping = cast(Mapping[Any, Any], obj) + return sum(estimate_model_size(value) for value in mapping.values()) + sys.getsizeof(cast(object, obj)) if isinstance(obj, (list, tuple, set)): - return sum(estimate_model_size(v) for v in obj) + sys.getsizeof(obj) - return sys.getsizeof(obj) + iterable = cast(Iterable[Any], obj) + return sum(estimate_model_size(value) for value in iterable) + sys.getsizeof(cast(object, obj)) + return sys.getsizeof(cast(object, obj)) def _healthcheck_wrapper(func: T_callable) -> T_callable: @@ -523,12 +542,11 @@ def _resolve_memory_threshold( if not (0 < value <= 1): raise ValueError(f"{name} ratio must be greater than 0 and at most 1") resolved = int(capacity_bytes * value) - elif isinstance(value, int): - if value < 0: - raise ValueError(f"{name} must be non-negative") - resolved = value else: - raise TypeError(f"{name} must be a float ratio or int bytes") + value_int = value + if value_int < 0: + raise ValueError(f"{name} must be non-negative") + resolved = value_int if resolved < minimum: raise ValueError(f"{name} must be at least {minimum} bytes") @@ -549,21 +567,15 @@ def _maybe_evict_spans(self) -> None: if self._total_span_bytes <= self._eviction_threshold_bytes: return - candidates = sorted( - ( - ( - ( - self._rollouts.get(rollout_id).start_time - if self._rollouts.get(rollout_id) - else spans[0].start_time or 0.0 - ), - rollout_id, - ) - for rollout_id, spans in self._spans.items() - if spans - ), - key=lambda item: item[0], - ) + candidates: List[tuple[float, str]] = [] + for rollout_id, spans in self._spans.items(): + if not spans: + continue + rollout = self._rollouts.get(rollout_id) + start_time = rollout.start_time if rollout is not None else (spans[0].start_time or 0.0) + candidates.append((start_time, rollout_id)) + + candidates.sort(key=lambda item: item[0]) for _, rollout_id in candidates: if self._total_span_bytes <= self._safe_threshold_bytes: diff --git a/tests/store/test_memory.py b/tests/store/test_memory.py index c99b2403d..f80a8b896 100644 --- a/tests/store/test_memory.py +++ b/tests/store/test_memory.py @@ -16,7 +16,7 @@ import asyncio import sys import time -from typing import List +from typing import List, Optional, cast from unittest.mock import Mock import pytest @@ -25,6 +25,7 @@ from agentlightning.store.memory import InMemoryLightningStore, estimate_model_size from agentlightning.types import ( LLM, + AttemptedRollout, Event, Link, PromptTemplate, @@ -588,7 +589,7 @@ async def test_span_eviction_removes_oldest_rollouts(mock_readable_span: Mock, m span_size_estimator=lambda span: 20, ) - attempted_rollouts = [] + attempted_rollouts: List[AttemptedRollout] = [] for index in range(4): attempted = await store.start_rollout(input={"index": index}) attempted_rollouts.append(attempted) @@ -609,8 +610,8 @@ def test_memory_threshold_accepts_byte_values() -> None: safe_memory_threshold=20, ) - assert store._eviction_threshold_bytes == 150 - assert store._safe_threshold_bytes == 20 + assert store._eviction_threshold_bytes == 150 # pyright: ignore[reportPrivateUsage] + assert store._safe_threshold_bytes == 20 # pyright: ignore[reportPrivateUsage] def test_memory_threshold_accepts_ratios_with_zero_safe(monkeypatch: pytest.MonkeyPatch) -> None: @@ -620,8 +621,8 @@ def test_memory_threshold_accepts_ratios_with_zero_safe(monkeypatch: pytest.Monk safe_memory_threshold=0.0, ) - assert store._eviction_threshold_bytes == int(200 * 0.6) - assert store._safe_threshold_bytes == 0 + assert store._eviction_threshold_bytes == int(200 * 0.6) # pyright: ignore[reportPrivateUsage] + assert store._safe_threshold_bytes == 0 # pyright: ignore[reportPrivateUsage] def test_invalid_safe_threshold_raises_value_error() -> None: @@ -687,13 +688,14 @@ def test_estimate_model_size_handles_span_objects() -> None: status_expected = sys.getsizeof(status) + sys.getsizeof(status.status_code) + sys.getsizeof(status.description) + trace_state_values = context.trace_state.values() if context.trace_state is not None else () context_expected = ( sys.getsizeof(context) + sys.getsizeof(context.trace_id) + sys.getsizeof(context.span_id) + sys.getsizeof(context.is_remote) + sys.getsizeof(context.trace_state) - + sum(sys.getsizeof(v) for v in context.trace_state.values()) + + sum(sys.getsizeof(v) for v in trace_state_values) ) event_attributes_expected = sys.getsizeof(event.attributes) + sys.getsizeof("value") @@ -702,12 +704,11 @@ def test_estimate_model_size_handles_span_objects() -> None: ) events_expected = sys.getsizeof(span.events) + event_expected - if link.attributes is None: - link_attributes_expected = sys.getsizeof(None) - else: - link_attributes_expected = sys.getsizeof(link.attributes) + sum( - sys.getsizeof(v) for v in link.attributes.values() - ) + link_attributes = cast(Optional[dict[str, str]], link.attributes) + link_attribute_values = link_attributes.values() if link_attributes is not None else () + link_attributes_expected = sys.getsizeof(link_attributes if link_attributes is not None else None) + sum( + sys.getsizeof(v) for v in link_attribute_values + ) link_expected = sys.getsizeof(link) + context_expected + link_attributes_expected links_expected = sys.getsizeof(span.links) + link_expected From 3acb8ff2e62f55bb2f6cd58670180e38f8e8319d Mon Sep 17 00:00:00 2001 From: Yuge Zhang Date: Thu, 16 Oct 2025 14:41:11 +0800 Subject: [PATCH 3/6] Fix psutil import --- agentlightning/store/memory.py | 15 +++++++-------- 1 file changed, 7 insertions(+), 8 deletions(-) diff --git a/agentlightning/store/memory.py b/agentlightning/store/memory.py index 274e4a62c..17b555643 100644 --- a/agentlightning/store/memory.py +++ b/agentlightning/store/memory.py @@ -5,8 +5,6 @@ import asyncio import functools import hashlib -import importlib -import importlib.util import logging import sys import threading @@ -119,13 +117,14 @@ def _generate_attempt_id() -> str: def _detect_total_memory_bytes() -> int: """Best-effort detection of the total available system memory in bytes.""" - psutil_spec = importlib.util.find_spec("psutil") - if psutil_spec is not None: - psutil = importlib.import_module("psutil") - return int(psutil.virtual_memory().total) + try: + import psutil - # Fallback to 8GB if memory cannot be detected. - return 8 * 1024**3 + return int(psutil.virtual_memory().total) + except ImportError: + # Fallback to 8GB if memory cannot be detected. + logger.error("psutil is not installed. Falling back to 8GB of memory in total.") + return 8 * 1024**3 class InMemoryLightningStore(LightningStore): From f0ee52cb49e06123e577e8e25ffdf5d1716a782d Mon Sep 17 00:00:00 2001 From: Yuge Zhang Date: Thu, 16 Oct 2025 14:48:15 +0800 Subject: [PATCH 4/6] fix pyright issues --- tests/store/test_memory.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/store/test_memory.py b/tests/store/test_memory.py index f80a8b896..0daa9de2b 100644 --- a/tests/store/test_memory.py +++ b/tests/store/test_memory.py @@ -662,7 +662,7 @@ class Outer(BaseModel): def test_estimate_model_size_handles_span_objects() -> None: status = TraceStatus(status_code="OK", description="fine") - context = SpanContext(trace_id="trace", span_id="parent", is_remote=False, trace_state={}) + context = SpanContext(trace_id="trace", span_id="parent", is_remote=False, trace_state={"foo": "bar"}) event = Event(name="step", attributes={"detail": "value"}, timestamp=1.0) link = Link(context=context, attributes=None) resource = Resource(attributes={"service.name": "unit"}, schema_url="schema") @@ -688,7 +688,7 @@ def test_estimate_model_size_handles_span_objects() -> None: status_expected = sys.getsizeof(status) + sys.getsizeof(status.status_code) + sys.getsizeof(status.description) - trace_state_values = context.trace_state.values() if context.trace_state is not None else () + trace_state_values = context.trace_state.values() context_expected = ( sys.getsizeof(context) + sys.getsizeof(context.trace_id) From 3ed56023ac88887b2d3c77355b8f1c46bf454b12 Mon Sep 17 00:00:00 2001 From: Yuge Zhang Date: Thu, 16 Oct 2025 14:49:48 +0800 Subject: [PATCH 5/6] add docs --- agentlightning/store/memory.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/agentlightning/store/memory.py b/agentlightning/store/memory.py index 17b555643..75696e41f 100644 --- a/agentlightning/store/memory.py +++ b/agentlightning/store/memory.py @@ -134,6 +134,14 @@ class InMemoryLightningStore(LightningStore): The methods in this class should generally not call each other, especially those that are locked. + + Args: + eviction_memory_threshold: The threshold for evicting spans in bytes. + By default, it's 70% of the total VRAM available. + safe_memory_threshold: The threshold for safe memory usage in bytes. + By default, it's 80% of the eviction threshold. + span_size_estimator: A function to estimate the size of a span in bytes. + By default, it's a simple size estimator that uses sys.getsizeof. """ def __init__( From 1686a07633f253f8df7d24d3e341cf9e3910e762 Mon Sep 17 00:00:00 2001 From: Yuge Zhang Date: Thu, 16 Oct 2025 14:53:10 +0800 Subject: [PATCH 6/6] add logs --- agentlightning/store/memory.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/agentlightning/store/memory.py b/agentlightning/store/memory.py index 75696e41f..30d89a705 100644 --- a/agentlightning/store/memory.py +++ b/agentlightning/store/memory.py @@ -584,10 +584,14 @@ def _maybe_evict_spans(self) -> None: candidates.sort(key=lambda item: item[0]) + logger.info(f"Evicting spans for {len(candidates)} rollouts to free up memory...") + memory_consumed_before = self._total_span_bytes for _, rollout_id in candidates: if self._total_span_bytes <= self._safe_threshold_bytes: break + logger.debug(f"Evicting spans for rollout {rollout_id} to free up memory...") self._evict_spans_for_rollout(rollout_id) + logger.info(f"Freed up {memory_consumed_before - self._total_span_bytes} bytes of memory") def _evict_spans_for_rollout(self, rollout_id: str) -> None: spans = self._spans.pop(rollout_id, [])