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
33 changes: 2 additions & 31 deletions atom/entrypoints/openai/api_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,6 @@
import asyncio
import base64
import binascii
import gc
import io
import json
import logging
Expand All @@ -39,7 +38,7 @@
from atom.model_engine.llm_engine import _load_tokenizer
from atom.model_engine.multimodal import build_multimodal_inputs
from atom.model_engine.request import RequestOutput
from atom.utils import envs
from atom.utils import tune_gc
from atom.utils.arg_parser import FlexibleArgumentParser

from .chat_encoders import apply_chat_template, load_custom_message_encoder
Expand Down Expand Up @@ -1106,34 +1105,6 @@ def do_preprocess():
# ============================================================================


def _tune_gc() -> None:
"""Stretch the interval between full (generation-2) collections.

Only gen-2 matters: it is stop-the-world and rescans every tracked
container, so its cost tracks live objects rather than recent garbage.
At c=2048 it was 47% of wall clock while the 23k gen-0/1 passes over the
same window cost 0.1s. Raising the thresholds is close to free because
reference counting, not the collector, reclaims everything acyclic --
peak RSS actually fell, since a larger gen-0 threshold lets more objects
die before a collection can promote them.

gc.freeze() was tried here and removed: once gen-2 stops running there is
nothing left for it to make cheaper, and alone it made collections more
frequent by emptying gen-2, which is the denominator CPython gates full
collections on (long_lived_pending > long_lived_total / 4).
"""
thresholds = envs.ATOM_GC_THRESHOLD
if not thresholds:
return
try:
t = tuple(int(x) for x in thresholds.split(","))
old = gc.get_threshold()
gc.set_threshold(*t)
logger.info("[gc] thresholds %s -> %s", old, t)
except (ValueError, TypeError):
logger.warning("[gc] bad ATOM_GC_THRESHOLD=%r, ignored", thresholds)


async def _refresh_metrics_once() -> None:
if engine is None:
return
Expand Down Expand Up @@ -1161,7 +1132,7 @@ async def lifespan(app: FastAPI):
"""Lifespan context manager for startup and shutdown."""
global _metrics_refresh_task
logger.info("Server started successfully and ready to accept requests")
_tune_gc()
tune_gc()
await _refresh_metrics_once()
_metrics_refresh_task = asyncio.create_task(_metrics_refresh_loop())
try:
Expand Down
8 changes: 8 additions & 0 deletions atom/model_engine/async_proc.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
import zmq
import zmq.asyncio
from aiter.dist.shm_broadcast import MessageQueue

from atom.kv_transfer.disaggregation import KVOutputAggregator
from atom.utils import (
get_mp_context,
Expand All @@ -34,6 +35,7 @@
resolve_obj_by_qualname,
set_process_title,
shutdown_all_processes,
tune_gc,
)
from atom.utils.numa_utils import numa_bind_to_node

Expand Down Expand Up @@ -83,6 +85,12 @@ def __init__(

enable_orphan_reaping()

# Second-order next to the EngineCore's call but not zero: without it
# a freeze still lands at the wave boundary, where the prefill burst
# churns enough objects to trigger gen-2 here. Runs before the model
# is built so the thresholds cover the startup heap.
tune_gc()

# NUMA-local CPU/memory pinning (see atom.utils.numa_utils).
# Auto-detects the GPU's local node by default; gated by
# ATOM_NUMA_BIND. Must run before any large allocation / native
Expand Down
6 changes: 6 additions & 0 deletions atom/model_engine/engine_core.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
init_exit_handler,
make_zmq_socket,
set_process_title,
tune_gc,
)
from atom.utils.distributed.utils import (
stateless_destroy_torch_distributed_process_group,
Expand Down Expand Up @@ -221,6 +222,11 @@ def run_engine(config: Config, input_address: str, output_address: str):
from atom.utils import enable_orphan_reaping

enable_orphan_reaping()

# The process whose GC pauses idle the GPUs: a gen-2 pass here stops
# the scheduler, and every ModelRunner worker then has nothing to run.
tune_gc()

engine: EngineCore = None
try:
if config.pipeline_parallel_size > 1:
Expand Down
37 changes: 37 additions & 0 deletions atom/utils/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
import contextlib
import copy
import dataclasses
import gc
import importlib
import ipaddress
import json
Expand Down Expand Up @@ -293,6 +294,42 @@ def enable_orphan_reaping(sig: int = signal.SIGKILL) -> bool:
return True


def tune_gc() -> None:
"""Stretch the interval between full (generation-2) collections.

Only gen-2 matters: it is stop-the-world and rescans every tracked
container, so its cost tracks live objects rather than recent garbage,
and at high concurrency it dominates while the gen-0/1 passes are noise.
Raising the thresholds is close to free because reference counting, not
the collector, reclaims everything acyclic -- peak RSS falls too, since a
larger gen-0 threshold lets more objects die before a collection can
promote them.

gc.freeze() was tried here and removed: once gen-2 stops running there is
nothing left for it to make cheaper, and alone it made collections more
frequent by emptying gen-2, which is the denominator CPython gates full
collections on (long_lived_pending > long_lived_total / 4).

Thresholds are per-interpreter, so every process has to call this itself:
subprocesses are spawned, and the API server's lifespan runs long after
they exist. Called from the API server, each EngineCore and each
ModelRunner worker; the EngineCore is the one that matters, since its
pauses stall the scheduler and the workers then have nothing to run.
"""
from atom.utils import envs

thresholds = envs.ATOM_GC_THRESHOLD
if not thresholds:
return
try:
t = tuple(int(x) for x in thresholds.split(","))
old = gc.get_threshold()
gc.set_threshold(*t)
logger.info("[gc] thresholds %s -> %s", old, t)
except (ValueError, TypeError):
logger.warning("[gc] bad ATOM_GC_THRESHOLD=%r, ignored", thresholds)


def kill_process_tree(pid: int):
"""
Kills all descendant processes of the given pid by sending SIGKILL.
Expand Down
4 changes: 3 additions & 1 deletion atom/utils/envs.py
Original file line number Diff line number Diff line change
Expand Up @@ -194,7 +194,9 @@
# --- Profiling & Logging ---
"ATOM_TORCH_PROFILER_DIR": lambda: os.getenv("ATOM_TORCH_PROFILER_DIR", None),
# "t0,t1,t2" for gc.set_threshold(); empty keeps CPython's default.
# See _tune_gc in api_server.py.
# Read independently by the API server, each EngineCore and each
# ModelRunner worker -- thresholds are per-interpreter. The EngineCore is
# the one that matters. See tune_gc in atom/utils/__init__.py.
"ATOM_GC_THRESHOLD": lambda: os.getenv("ATOM_GC_THRESHOLD", "").strip(),
"ATOM_PROFILER_MORE": lambda: os.getenv("ATOM_PROFILER_MORE", "0") == "1",
# When profiling is active, append detailed attention aggregates (sqsq, sqsk, sk)
Expand Down
Loading