From cb2fa614ae08eb6797fd5880623c820fbe1701a6 Mon Sep 17 00:00:00 2001 From: AngeloDanducci Date: Wed, 19 Aug 2026 10:30:28 -0400 Subject: [PATCH 1/2] fix: clear ModelOutput fields via mapping to free tensors Signed-off-by: AngeloDanducci --- mellea/backends/huggingface.py | 26 ++++-- test/backends/test_huggingface_unit.py | 113 +++++++++++++++++++++++++ 2 files changed, 130 insertions(+), 9 deletions(-) diff --git a/mellea/backends/huggingface.py b/mellea/backends/huggingface.py index 870e9817d4..249b4db4d6 100644 --- a/mellea/backends/huggingface.py +++ b/mellea/backends/huggingface.py @@ -1519,12 +1519,16 @@ class used during generation, if any. self.cache_put(cache_key, cache_info) # Clear KV cache and scores from HF output; retained via LRU cache above. - hf_output.past_key_values = None - hf_output.scores = None + # ModelOutput mirrors fields into an OrderedDict mapping; assigning the + # attribute to None only clears the __dict__ slot and leaves the mapping + # entry (and its tensor) alive, so clear through the mapping instead. + hf_output["past_key_values"] = None + hf_output["scores"] = None # Clear the raw logits tensor (scores already cleared above if cached). + # Route through the mapping for the same reason as above. if isinstance(hf_output, GenerateDecoderOnlyOutput): - hf_output.logits = None + hf_output["logits"] = None # Only scan for tools if we are not doing structured output and tool calls were provided to the model. if _format is None and tool_calls: @@ -1605,12 +1609,16 @@ class used during generation, if any. import gc hf_out = mot.raw.response - if hasattr(hf_out, "sequences") and hf_out.sequences is not None: - del hf_out.sequences - if hasattr(hf_out, "scores") and hf_out.scores is not None: - del hf_out.scores - if hasattr(hf_out, "logits") and hf_out.logits is not None: - del hf_out.logits + # ModelOutput has no __delattr__ and its __setattr__ skips the mapping + # write for None, so `del hf_out.f` / `hf_out.f = None` leave the mapping + # entry (and its tensor) alive. Clear through the mapping to actually + # release the tensors before dropping the container. + if hf_out.sequences is not None: + hf_out["sequences"] = None + if hf_out.scores is not None: + hf_out["scores"] = None + if hf_out.logits is not None: + hf_out["logits"] = None mot.raw.response = None # Force Python GC and return CUDA memory to device diff --git a/test/backends/test_huggingface_unit.py b/test/backends/test_huggingface_unit.py index beaf521871..12d2b02d79 100644 --- a/test/backends/test_huggingface_unit.py +++ b/test/backends/test_huggingface_unit.py @@ -4,6 +4,8 @@ """Unit tests for HuggingFace backend pure-logic helpers — no model load required.""" import asyncio +import gc +import weakref from types import SimpleNamespace from unittest.mock import MagicMock, patch @@ -1393,3 +1395,114 @@ def _capture_grammar(schema, overrides=None): assert captured[0].get("whitespace_pattern") == r"[\x20\x0A\x0D\x09]{0,20}", ( f"Expected bounded whitespace_pattern to override False in {path_name}" ) + + +async def _post_process_holding_only_weakrefs( + backend: LocalHFBackend, n_steps: int +) -> tuple[ModelOutputThunk, dict[str, list["weakref.ref"]]]: + """Run post_processing on a MOT and return weakrefs to the raw tensors it cleared. + + Builds a `GenerateDecoderOnlyOutput` whose `scores`, `logits`, and + `past_key_values` hold fresh tensors, takes weakrefs to those tensors, then + drops every strong local reference so the only remaining strong references + are the ones the `ModelOutput` mapping keeps. After `post_processing` clears + the fields, a correct clear drops the mapping entries and the weakrefs die on + the next GC; the buggy attribute-set/`del` leaves them alive. + + Args: + backend: The backend whose `post_processing` is exercised. + n_steps: Number of decode steps (length of the scores/logits tuples). + + Returns: + The finalized thunk and a dict mapping field name to the list of + weakrefs (one per step) to the tensors that field held. + """ + input_ids = torch.tensor([[1]]) + sequences = torch.tensor([[0, 0]]) + scores = tuple(torch.zeros(1, 32000) for _ in range(n_steps)) + logits = tuple(torch.ones(1, 32000) for _ in range(n_steps)) + kv = tuple(torch.zeros(2, 4) for _ in range(n_steps)) + + refs: dict[str, list[weakref.ref]] = { + "scores": [weakref.ref(t) for t in scores], + "logits": [weakref.ref(t) for t in logits], + "past_key_values": [weakref.ref(t) for t in kv], + } + + mot = ModelOutputThunk(value="hi") + mot._call.action = Message("user", "noop") + mot._call.model_options = {} + mot.raw.response = GenerateDecoderOnlyOutput( + sequences=sequences, + scores=scores, + logits=logits, + attentions=None, + hidden_states=None, + past_key_values=kv, + ) + + # Drop every strong local reference to the tensors and their tuples. After + # this, the ModelOutput mapping is the only thing keeping them alive. + del scores, logits, kv + + await backend.post_processing(mot, [], None, False, {}, None, input_ids) + + return mot, refs + + +@pytest.mark.asyncio +async def test_post_processing_clearing_raw_logits_actually_releases_them(): + """Clearing `hf_output.logits` must drop the tensors, not just the attribute. + + `GenerateDecoderOnlyOutput` is a `ModelOutput`, i.e. an `OrderedDict` subclass + that mirrors every field into the mapping. `ModelOutput.__setattr__` skips the + mapping write when the value is `None`, and `ModelOutput` defines no + `__delattr__`, so `out.logits = None` and `del out.logits` both leave the + mapping entry — and therefore the tensors — in place. Any code that nulls a + field to free memory while keeping the container has to clear the mapping too. + """ + backend = _make_backend(1) + backend._use_caches = True # keeps raw.response, so the container survives + + mot, refs = await _post_process_holding_only_weakrefs(backend, n_steps=2) + gc.collect() + gc.collect() + + assert mot.raw.response is not None, "test setup: raw.response should be retained" + assert mot.raw.response.logits is None, "test setup: logits attribute was cleared" + for step, ref in enumerate(refs["logits"]): + assert ref() is None, ( + f"raw logits tensor for step {step} is still alive after hf_output.logits " + "was set to None — the ModelOutput mapping entry still references it" + ) + + +@pytest.mark.asyncio +async def test_post_processing_clearing_scores_and_kv_actually_releases_them(): + """The caching branch clears `scores`/`past_key_values` off the MOT. + + Those tensors are retained in the LRU cache, so nulling the fields on the + `ModelOutput` must actually release the container's strong reference (via the + mapping), otherwise the KV cache is held twice — once in the LRU and once on + the MOT — defeating the cleanup. + """ + backend = _make_backend(1) + backend._use_caches = True + + mot, refs = await _post_process_holding_only_weakrefs(backend, n_steps=2) + gc.collect() + gc.collect() + + assert mot.raw.response is not None, "test setup: raw.response should be retained" + assert mot.raw.response.scores is None, "test setup: scores attribute was cleared" + assert mot.raw.response.past_key_values is None, ( + "test setup: past_key_values attribute was cleared" + ) + # With the default zero-capacity LRU, cache_put evicts immediately, so the + # MOT's ModelOutput mapping is the only thing that could keep the scores + # tensors alive; clearing the field through the mapping must release them. + for step, ref in enumerate(refs["scores"]): + assert ref() is None, ( + f"raw scores tensor for step {step} is still alive after hf_output.scores " + "was set to None — the ModelOutput mapping entry still references it" + ) From 2a19d3210f2e9e04e3beaa6082e1e083fd859115 Mon Sep 17 00:00:00 2001 From: AngeloDanducci Date: Wed, 19 Aug 2026 10:54:56 -0400 Subject: [PATCH 2/2] use ignore to satisfy mypy in new test Signed-off-by: AngeloDanducci --- test/backends/test_huggingface_unit.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/test/backends/test_huggingface_unit.py b/test/backends/test_huggingface_unit.py index 12d2b02d79..b4a6eeac6a 100644 --- a/test/backends/test_huggingface_unit.py +++ b/test/backends/test_huggingface_unit.py @@ -1438,7 +1438,10 @@ async def _post_process_holding_only_weakrefs( logits=logits, attentions=None, hidden_states=None, - past_key_values=kv, + # past_key_values is typed Cache | None, but ModelOutput stores whatever + # it is handed in its mapping; a tuple of tensors is all this fixture needs + # to exercise the clearing path. + past_key_values=kv, # type: ignore[arg-type] ) # Drop every strong local reference to the tensors and their tuples. After