Autograd was on: 3 GB per highlighted chunk and eight OOM kills in three days
Citation highlighting ran a Hugging Face model with autograd enabled. Each chunk kept 3 GB it never gave back. One line of torch.inference_mode() fixed it.
Between 15 and 17 September 2026 our production rag_service container was killed by the kernel for running out of memory eight times, three of them inside twelve minutes while a customer was testing their public chatbox. The cause was a sentence-highlighting model running inference with PyTorch's autograd switched on, so every 6,000-character chunk it scored left about 3 GB of memory behind that was never released.
The fix was one with statement. Getting to it took longer, mostly because the first things I tried were the usual fixes for a Python process that keeps growing, and none of them touched the problem.
What the customer saw
Two symptoms came in together, both from the same customer chatbox. Streaming answers would intermittently die part way through with this in the chat window:
Streaming error: Stream query error: peer closed connection without sending
complete message body (incomplete chunked read)
And a plain search on the same deployment would sometimes return "Internal server error".
Both errors had the same cause. When the kernel kills rag_service, every stream it is serving dies mid-response, which is where "peer closed connection" comes from. The container then restarts and spends about 35 seconds loading knowledgebases before it can answer anything, and the first request that lands in that window gets a 500 with a message ending "not found or not loaded".
The feature doing the damage is semantic highlighting: when Certant cites a chunk behind an answer, a small encoder (zilliz/semantic-highlight-bilingual-v1) scores each sentence in the chunk against the question so the UI can highlight the ones that support it. If you've read the piece on moving that same model to the GPU, this is a different failure in the same few hundred lines of code. That piece was about speed on a contended CPU.
Confirming it was the OOM killer
Everything at this stage was read-only on production. I started with the container's restart count, then the kernel log. Grepping dmesg for cgroup OOM events gives you lines like Memory cgroup out of memory and an oom-kill:constraint line carrying the cgroup ID of the victim:
docker inspect <container> | grep RestartCount
sudo dmesg -T | grep 'Memory cgroup out of memory'
sudo dmesg -T | grep 'oom-kill:constraint'
docker ps --no-trunc # map the cgroup id back to a container name
The cgroup ID mapped back to rag_service, which runs with a 12 GB memory limit. docker logs -t on the container showed each kill from the other side: Killed, followed by a fresh Application startup complete.
The kernel log went back to 9 September and had no kills before a mid-September release went out and the new chatbox started taking traffic.
Reproducing it locally
I had a local copy of the same knowledgebase, so I reproduced it there rather than poke at production any further, using two tools.
The first was a small wrapper that sampled docker stats for the container while firing one query with semantic highlighting enabled. At service level, one highlighted query took rag_service from 6.7 GB to 15.5 GB. Against a 12 GB limit, production could only get through one or two cited queries before the kernel stepped in.
The second was a standalone script run inside the container. It loaded the model through the service's own semantic_highlight.get_model(), called the model's process() method directly on real cited chunks from that knowledgebase, and printed four things after each call: RSS, USS (the memory unique to the process, which is what matters when you're asking "did this process keep something?"), the number of child processes, and the total size of every live torch.Tensor object Python could find.
Per chunk, the numbers were ugly. A 6,000-character chunk added 3.0 GB of RSS and took between 5 and 19 seconds. A 1,700-character chunk added 0.6 GB. None of it came back after the call returned.
The fixes that did nothing
A long-running Python process whose memory only goes up has a well-worn list of suspects, and I went through most of it. Each of these was tried and measured with the same script, and each one was a no-op:
- calling
malloc_trim, to hand freed heap back to the OS MALLOC_ARENA_MAX=2, to stop glibc creating a fresh arena per thread (highlighting runs chunks across a thread pool)- tuning
MALLOC_MMAP_THRESHOLD_, so large allocations go throughmmapand get returned on free - turning oneDNN off
- limiting torch to 2 threads
- setting the model's DataLoader preprocessing workers to 0
That last one was aimed at a different suspect: forked preprocessing workers holding the memory in child processes. It turned out the model's own tuning already sets workers to 0 for anything under 2,000 jobs, so there were no forked children to blame, and the child count in the probe script stayed at zero throughout.
In hindsight the allocator knobs were never going to work. They help when memory has been freed but not yet returned to the OS, and this memory was still in use.
Live tensors stayed flat while memory climbed
The live tensor total from the probe script stayed at 2.12 GB across every run, which is the model's weights. It did not move after a 6,000-character chunk, and it did not move after several. Meanwhile USS kept climbing by gigabytes.
So Python could see 2.12 GB of tensors, and the process was holding a great deal more than that. Whatever was growing was not reachable as a Python tensor object, which is why the garbage collector never saw it. The memory was being held from C++.
In a PyTorch process, the obvious C++-side thing that holds large tensors without Python owning them is autograd. When gradient tracking is on, every operation in a forward pass records itself in a graph and saves the intermediate activations it will need to compute gradients later. For a transformer encoder run over a few thousand tokens, those saved activations are big. The Python code never touches them directly; they hang off the graph.
Our code never asked for gradients or called backward(), but it never switched autograd off either, and neither did the model.
The highlighter is loaded with trust_remote_code=True, which means process() is the model author's own method shipped alongside the weights. It runs its forward passes with autograd enabled. It never wraps them in torch.no_grad() or torch.inference_mode(). PyTorch's default is gradients on, so every cited chunk on every query built a gradient graph.
The one-line fix
The change wraps the call to process() in inference mode. It lives in a configuration file, inside _highlight_one, which is what each thread in the highlighting pool runs for one chunk:
def _inference_context():
"""torch.inference_mode() when torch is importable, else a no-op."""
try:
import torch
return torch.inference_mode()
except Exception: # torch absent (unit tests with a fake model)
return contextlib.nullcontext()
# inside _highlight_one()
with _inference_context():
result = model.process(
question=query,
context=content,
threshold=threshold,
return_sentence_metrics=True,
)
The helper exists only so the unit tests, which use a fake model and don't always have torch installed, still import cleanly. In the service it is torch.inference_mode(), full stop.
inference_mode() goes a step further than no_grad(): as well as not recording a graph, it skips the version-counter and view bookkeeping autograd would otherwise do, and tensors created inside it can't be fed back into autograd later. For a model whose output is a list of sentence scores that nobody will ever differentiate, there is no downside.
I ran the same script on the same chunks before and after. All figures are in-container measurements on the local copy of the knowledgebase, 17 September 2026:
| Measurement | Autograd on (before) | inference_mode() (after) |
|---|---|---|
| Memory added by a 6,000-character chunk | +3.0 GB, never released | +0.3 GB |
| Memory added by repeating that chunk | not recorded | +0 GB |
| Memory added by a 1,700-character chunk | +0.6 GB | not recorded |
| Time for a 6,000-character chunk | 5 to 19 s | 2 s |
rag_service after highlighted queries |
6.7 GB rising to 15.5 GB after one query | steady at 6.0 GB, new chunks add nothing |
The speed-up came free with the memory fix. I didn't break down how much of the old 5 to 19 seconds was graph bookkeeping and how much was the process allocating fresh gigabytes on every call, so I won't guess at the split.
The fix shipped in release v2.21.6.
Keeping it fixed
A single missing context manager is exactly the sort of thing that gets lost in a refactor, so the fix came with a regression test. The fake model in test_semantic_highlight_parallel.py records torch.is_inference_mode_enabled() when its process() is called, and the test checks both sides of the call:
def test_model_runs_under_inference_mode(fake_model):
torch = pytest.importorskip("torch")
assert not torch.is_inference_mode_enabled()
sh.highlight_chunks("q", [{"id": "c1", "content": "q here. other."}], 0.5)
assert fake_model.inference_mode_seen is True
# and the mode does not leak past the call
assert not torch.is_inference_mode_enabled()
The last assertion matters as much as the middle one. Inference mode is thread-local and scoped to the with block, and a version that leaked it into the rest of the request would break anything downstream that did want gradients.
What I'd tell anyone serving a Hugging Face model
Calling a model's forward pass, or a convenience method built on top of it, does not turn autograd off for you. If the code that runs the model doesn't say no_grad or an environment variable, you are building a gradient graph on every call. With trust_remote_code you are running whatever the model author wrote, and in our case the author wrote a method meant for scoring that never disabled gradients.
Whether that graph gets freed promptly depends on what holds references to it, and I would rather not rely on that. In our process it didn't get freed, and a container with a 12 GB limit was dead after a couple of queries.
If you want to check your own service, the diagnostic is cheap. Load the model the way the service does, call it on a realistic input a few times, and after each call print the process's USS alongside the total size of live torch.Tensor objects. If the tensor total sits still at roughly the size of your weights while USS climbs, stop tuning the allocator. Wrap the call in torch.inference_mode() and measure again.
And if you see "peer closed connection" and an occasional 500 arriving together from the same service, check the restart count and dmesg before you go looking for two bugs.



