Skip to content
Merged
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
168 changes: 166 additions & 2 deletions agentlightning/store/memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,13 +6,30 @@
import functools
import hashlib
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 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

from agentlightning.types import (
Attempt,
Expand All @@ -35,6 +52,21 @@
logger = logging.getLogger(__name__)


def estimate_model_size(obj: Any) -> int:
"""Rough recursive size estimate for Pydantic BaseModel instances."""

if isinstance(obj, BaseModel):
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))

Copilot AI Oct 16, 2025

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The mapping size estimation is incomplete. It should include the size of dictionary keys as well as values.

Suggested change
return sum(estimate_model_size(value) for value in mapping.values()) + sys.getsizeof(cast(object, obj))
return sum(estimate_model_size(key) + estimate_model_size(value) for key, value in mapping.items()) + sys.getsizeof(cast(object, obj))

Copilot uses AI. Check for mistakes.
if isinstance(obj, (list, tuple, set)):
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:
"""
Decorator to run the watchdog healthcheck **before** executing the decorated method.
Expand Down Expand Up @@ -82,16 +114,43 @@ 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."""

try:
import psutil

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):
"""
In-memory implementation of LightningStore using Python data structures.
Thread-safe and async-compatible but data is not persistent.

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.

Copilot AI Oct 16, 2025

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Corrected 'VRAM' to 'RAM' as this refers to system memory, not video memory.

Suggested change
By default, it's 70% of the total VRAM available.
By default, it's 70% of the total RAM available.

Copilot uses AI. Check for mistakes.
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__(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
Expand All @@ -105,6 +164,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
Expand Down Expand Up @@ -426,6 +515,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()
Expand All @@ -449,6 +540,77 @@ 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)
else:
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")

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: 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])

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, [])
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]:
"""
Expand Down Expand Up @@ -510,6 +672,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
Expand Down
Loading
Loading