Skip to main content

vrambatch

Run batched GPU work inside a VRAM budget — and when a batch is still too big, halve it and retry instead of crashing.

If you serve inference on a shared or memory-capped GPU, you have written this code already: pick a batch size that fits, and pray no single request is heavier than your estimate. vrambatch is that logic, extracted and tested — size a pass from a memory budget, and self-heal on CUDA out-of-memory so a bad estimate costs an extra pass, not a 500.

pip install vrambatch          # sizing only (pure Python)
pip install vrambatch[torch]   # + auto-sizing and OOM detection on GPU

The problem it removes

A memory-capped process picks a batch size up front. Two things then go wrong:

  • Too small and the GPU sits idle — most of your VRAM cap is unreachable and requests queue instead of fusing.
  • Too big and one unusually heavy item OOMs the whole batch, failing every request in it.

The honest fix is to size for the typical case and recover from the rare overflow. That needs a retry that can make progress — which is exactly what a hand-rolled version usually lacks.

Usage

from vrambatch import run_in_passes

# process() takes a list of items, returns one result per item, in order.
def process(batch):
    feats = preprocess(batch)
    return model.generate(feats)        # your real GPU call

results = run_in_passes(
    items,
    process,
    gb_per_unit=0.25,        # measured VRAM per unit (see below)
    budget_gb=21,            # your cap; omit to read free VRAM automatically
    reserve_gb=3.0,          # model weights + context, subtracted first
    cost=lambda item: n_chunks(item),   # units per item (default 1)
)
  • Items are grouped into passes whose total cost stays within the budget.
  • Each pass is one process() call.
  • On CUDA OOM the offending pass is split in half and retried, recursing to a single item. A lone item that still cannot fit raises OOMError — a real failure, reported as one. Non-OOM errors are never retried.

Pass on_oom=lambda n: log(...) to observe splits.

Just the sizing, if that's all you need

from vrambatch import plan_pass_size, available_vram_gb

n = plan_pass_size(budget_gb=21, gb_per_unit=0.25, reserve_gb=3.0, safety=0.8)
# -> units per pass

free = available_vram_gb()   # device-free + the allocator's reusable pool

Measuring gb_per_unit

Do not guess it high "to be safe" — an over-large figure is what leaves your cap unreachable. Measure it: run increasing batch sizes and fit a line to peak memory.

import torch
for n in (1, 2, 4, 8, 16):
    torch.cuda.reset_peak_memory_stats()
    process(items[:n])
    print(n, torch.cuda.max_memory_allocated() / 2**30)
# peak ≈ reserve_gb + gb_per_unit * n  → the slope is gb_per_unit

Because overflow is recoverable, size for the typical unit and let the halve-and-retry cover the occasional heavy one. The safety factor (default 0.8) absorbs normal variance; the retry covers the rest.

Notes

  • torch is optional. plan_pass_size is pure Python. available_vram_gb and precise OOM typing use torch when present; OOM detection falls back to matching the error message so RuntimeError("CUDA out of memory") is caught either way.
  • Async servers: run_in_passes is synchronous by design — call it from a thread (await asyncio.to_thread(run_in_passes, ...) or a single-worker executor) so it never blocks your event loop.
  • Fragmentation: under a hard set_per_process_memory_fraction cap, also set PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True before torch inits CUDA, or the allocator's cached pool can fragment and fail an allocation that is nominally under budget.

License

MIT

Metadata

Release files for vrambatch 0.1.0

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for vrambatch 0.1.0
File Size Uploaded
vrambatch-0.1.0.tar.gz 9.9 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for vrambatch 0.1.0
File Interpreter ABI Platform
vrambatch-0.1.0-py3-none-any.whl Python 3 none any Details

Total release size: 17.3 kB

Release files / vrambatch-0.1.0.tar.gz

Download URL vrambatch-0.1.0.tar.gz
Size 9.9 kB
Tags Source
SHA-256 checksum
How to use checksums
b4692c438bf9d28392dac137b1e94aac3bd2387a2d292b599f5520378bcfb028
BLAKE2b-256 checksum
How to use checksums
6d16ee46d54b37b3e0be8308c109b65b51fc097414afd87f76bdb9f99e43b328
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.12.3

Release files / vrambatch-0.1.0-py3-none-any.whl

Download URL vrambatch-0.1.0-py3-none-any.whl
Size 7.4 kB
Tags Python 3
SHA-256 checksum
How to use checksums
2bc54e5e721f691e5ae3fb20c1d54d6e3f3e2a43ddfb2e8020ae9a276323508a
BLAKE2b-256 checksum
How to use checksums
0093134ef056beab9234c13205c50a2cd6c03fad3281824eaa6e8d981972a9c4
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.12.3

Release history Release notifications | RSS feed

This release

0.1.0 This release

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page