vllm-vast-node / agent.py
cloud19's picture
обновление agent.py
5f92610 verified
Raw
History Blame Contribute Delete
29.7 kB
#!/usr/bin/env python3
"""Агент LLM-ноды: поднимает vLLM по инстансу на карту и держит их
зарегистрированными в vLLM-менеджере.
Работает внутри образа vllm/vllm-openai, только на стандартной библиотеке
плюс huggingface_hub (он уже в образе). Ничего не устанавливает.
Что делает:
1. читает карты через nvidia-smi и раскладывает по ним инстансы —
контекст и размер батча считаются от VRAM, а не берутся с потолка;
2. один раз скачивает веса, чтобы инстансы не тянули их наперегонки;
3. запускает vLLM и следит за процессами;
4. как только инстанс отвечает на /health — регистрирует его в менеджере
по публичному адресу Vast, а когда он умирает — снимает.
Регистрация идёт по публичному адресу вида http://<ip>:<проброшенный порт>:
менеджер ходит к ноде сам, поэтому внутренний порт ему бесполезен. Соответствие
внутренних портов внешним знает только Vast — его и спрашиваем контейнерным
ключом, который лежит в /root/.vast_api_key.
"""
import json
import os
import re
import signal
import subprocess
import sys
import threading
import time
import urllib.error
import urllib.request
LOG_DIR = "/var/log/vllm-node"
STATE_DIR = "/opt/vllm-node"
VAST_API = "https://console.vast.ai/api/v0"
# Пулы менеджера. Список закрытый: на неизвестный тип /admin/servers/add
# отвечает 422, поэтому лучше отсеять опечатку до отправки.
KNOWN_POOLS = ("regular", "premium", "summary", "critic", "test", "extractor")
def manual_mode() -> bool:
"""Ручной режим: агент держит vLLM, но не трогает менеджер.
Нужен на отладке. Автопилот иначе мешает: снимешь ноду с пула руками —
через минуту сверка вернёт её обратно. Переключается файлом, а не
переменной, чтобы не перезапускать ноду ради смены режима.
"""
return os.path.exists(os.path.join(STATE_DIR, "manual"))
def log(msg: str) -> None:
print(f"[{time.strftime('%Y-%m-%d %H:%M:%S')}] {msg}", flush=True)
def env(name: str, default: str = "") -> str:
return (os.getenv(name) or default).strip()
def env_int(name: str, default: int) -> int:
raw = env(name)
return int(raw) if raw else default
def env_float(name: str, default: float) -> float:
raw = env(name)
return float(raw) if raw else default
# --- Конфигурация -----------------------------------------------------------
MODEL_ID = env("MODEL_ID", "cloud19/G4-MeroMero-26B-FP8-Dynamic-Uncensored")
SERVED_NAME = env("SERVED_MODEL_NAME", "gemma-rp-uncensored")
VLLM_API_KEY = env("VLLM_API_KEY", "dj9hj342fhc4cinj4hj092Hd38D")
HF_TOKEN = env("HF_TOKEN")
MANAGER_URL = env("MANAGER_URL", "http://185.70.186.164:5555").rstrip("/")
MANAGER_API_KEY = env("MANAGER_API_KEY")
NODE_POOLS = env("NODE_POOLS", "test")
GPU_LAYOUT = env("GPU_LAYOUT")
REGISTER = env("REGISTER", "1") != "0"
PORTS = [int(p) for p in env("PORTS", "8080,1111,8082,8083,8084,8085,8086,8087").split(",") if p.strip()]
TP_SIZE = env_int("TENSOR_PARALLEL_SIZE", 1)
VLLM_EXTRA_ARGS = env("VLLM_EXTRA_ARGS")
# Ручные подпорки: если заданы, перебивают всё, что посчитано от железа.
FORCE_CTX = env_int("MAX_MODEL_LEN", 0)
FORCE_SEQS = env_int("MAX_NUM_SEQS", 0)
FORCE_UTIL = env_float("GPU_MEMORY_UTILIZATION", 0.0)
FORCE_KV = env("KV_CACHE_DTYPE")
HEALTH_TIMEOUT = env_int("HEALTH_TIMEOUT", 3600) # сколько ждём первого /health
POLL_INTERVAL = env_int("POLL_INTERVAL", 10)
SYNC_INTERVAL = env_int("SYNC_INTERVAL", 60) # сверка с менеджером
UNHEALTHY_GRACE = env_int("UNHEALTHY_GRACE", 120) # сколько терпим молчание перед снятием
WEIGHTS_GIB = 26.7 # G4-MeroMero-26B FP8; для другой модели см. WEIGHTS_GIB в env
WEIGHTS_GIB = env_float("WEIGHTS_GIB", WEIGHTS_GIB)
# --- Железо -----------------------------------------------------------------
def discover_gpus():
out = subprocess.run(
["nvidia-smi", "--query-gpu=index,name,memory.total,compute_cap",
"--format=csv,noheader,nounits"],
capture_output=True, text=True, timeout=60)
if out.returncode != 0:
raise RuntimeError(f"nvidia-smi не отвечает: {out.stderr.strip()[:200]}")
gpus = []
for line in out.stdout.strip().splitlines():
idx, name, mib, cap = [p.strip() for p in line.split(",")]
gpus.append({"index": int(idx), "name": name, "vram_gib": int(mib) / 1024,
"cap": float(cap)})
return gpus
def gpu_processes(indices):
"""Чужие процессы на наших картах: (pid, занято МБ)."""
found = []
for idx in indices:
out = subprocess.run(
["nvidia-smi", "--query-compute-apps=pid,used_gpu_memory",
"--format=csv,noheader,nounits", "--id", str(idx)],
capture_output=True, text=True, timeout=60)
for line in out.stdout.strip().splitlines():
if not line.strip():
continue
pid, _, mem = line.partition(",")
try:
found.append((int(pid.strip()), int(mem.strip())))
except ValueError:
continue
return found
def shape_for(vram_gib: float, cap: float):
"""Контекст, батч и утилизация под конкретную карту.
Считаем от того, сколько остаётся под KV-кэш после весов. Цифры сверены
с продом: на B200 (183 ГБ) vLLM отдаёт под кэш 132 ГиБ при утилизации 0.90.
"""
util = 0.90
free = vram_gib * util - WEIGHTS_GIB - 4.0 # 4 ГиБ — активации и графы CUDA
if free < 2.0 and vram_gib < 40:
util = 0.94 # на тесных картах выжимаем остаток
free = vram_gib * util - WEIGHTS_GIB - 3.0
if free <= 0.5:
return None # веса не влезают, карту пропускаем
if free >= 100:
ctx, seqs = 16384, 512
elif free >= 40:
ctx, seqs = 16384, 384
elif free >= 12:
ctx, seqs = 8192, 192
elif free >= 4:
ctx, seqs = 8192, 48
else:
ctx, seqs = 4096, 16
# FP8-кэш вдвое дешевле по памяти, но нативно считается с Ada и новее.
kv = "fp8" if cap >= 8.9 else "auto"
return {"ctx": ctx, "seqs": seqs, "util": util, "kv": kv, "kv_free_gib": round(free, 1)}
def parse_layout(spec: str):
"""`0:8192:regular 1:16384:premium,test` — карта, контекст, пулы.
Контекст можно пропустить (`0::regular`) — тогда он считается от VRAM.
"""
plan = {}
for item in re.split(r"[;\s]+", spec.strip()):
if not item:
continue
parts = item.split(":")
if len(parts) != 3:
raise ValueError(f"GPU_LAYOUT: непонятный кусок {item!r}, нужно idx:ctx:пулы")
idx, ctx, pools = parts
plan[int(idx)] = {"ctx": int(ctx) if ctx else 0,
"pools": [p for p in pools.split(",") if p]}
return plan
def build_plan():
gpus = discover_gpus()
log(f"карт найдено: {len(gpus)}")
for g in gpus:
log(f" GPU{g['index']}: {g['name']}, {g['vram_gib']:.0f} ГиБ, compute {g['cap']}")
default_pools = [p.strip() for p in NODE_POOLS.split(",") if p.strip()]
layout = parse_layout(GPU_LAYOUT) if GPU_LAYOUT else {}
groups = []
if TP_SIZE > 1:
for start in range(0, len(gpus) - len(gpus) % TP_SIZE, TP_SIZE):
groups.append(gpus[start:start + TP_SIZE])
else:
groups = [[g] for g in gpus]
plan = []
for slot, group in enumerate(groups):
head = group[0]
shape = shape_for(head["vram_gib"], head["cap"])
if shape is None:
log(f" GPU{head['index']}: {head['vram_gib']:.0f} ГиБ мало под "
f"{WEIGHTS_GIB:.1f} ГиБ весов — карта пропущена")
continue
if slot >= len(PORTS):
log(f" GPU{head['index']}: портов в PORTS меньше, чем карт — карта пропущена")
continue
override = layout.get(head["index"], {})
ctx = FORCE_CTX or override.get("ctx") or shape["ctx"]
pools = override.get("pools") or default_pools
bad = [p for p in pools if p not in KNOWN_POOLS]
if bad:
raise ValueError(f"неизвестные пулы {bad}, допустимы {list(KNOWN_POOLS)}")
plan.append({
"slot": slot,
"gpus": [g["index"] for g in group],
"gpu_name": head["name"],
"port": PORTS[slot],
"ctx": ctx,
"seqs": FORCE_SEQS or shape["seqs"],
"util": FORCE_UTIL or shape["util"],
"kv": FORCE_KV or shape["kv"],
"pools": pools,
"tp": len(group),
})
if not plan:
raise RuntimeError("ни одной пригодной карты — запускать нечего")
return plan
# --- Веса -------------------------------------------------------------------
def fetch_weights():
"""Тянем модель один раз до старта инстансов.
Иначе несколько vLLM лезут в один кэш одновременно: время старта растёт,
а на тесном диске это ещё и лишний риск.
"""
log(f"качаю веса {MODEL_ID}")
started = time.time()
from huggingface_hub import snapshot_download
path = snapshot_download(MODEL_ID, token=HF_TOKEN or None,
max_workers=8, ignore_patterns=["*.pt", "*.bin"])
log(f"веса на месте за {time.time() - started:.0f} с: {path}")
return path
# --- Публичный адрес --------------------------------------------------------
def _read(path: str) -> str:
try:
with open(path) as fh:
return fh.read().strip()
except OSError:
return ""
class PublicAddress:
"""Соответствие внутренних портов внешним.
На Vast его знает только API, поэтому спрашиваем контейнерным ключом.
Вне Vast (голое железо, свой сервер) достаточно PUBLIC_HOST — порты там
и так совпадают.
"""
def __init__(self):
self.host = env("PUBLIC_HOST")
self.ports = {}
self._loaded = False
def load(self) -> None:
if self._loaded:
return
key = env("VAST_CONTAINER_API_KEY") or _read("/root/.vast_api_key")
label = _read("/root/.vast_containerlabel")
instance_id = label.replace("C.", "").strip()
if key and instance_id:
# Здесь нельзя откатываться на «внешний ip + тот же порт»: снаружи
# порт другой, и нода зарегистрируется по адресу, куда менеджер не
# достучится. Лучше долбиться в API, пока не ответит.
for attempt in range(1, 11):
try:
status, data = http_json("GET", f"{VAST_API}/instances/{instance_id}/",
headers={"Authorization": f"Bearer {key}"}, timeout=30)
if status != 200:
raise RuntimeError(f"HTTP {status}: {str(data)[:200]}")
inst = data["instances"]
self.host = self.host or inst.get("public_ipaddr", "").strip()
for spec, binds in (inst.get("ports") or {}).items():
inner = int(spec.split("/")[0])
for b in binds:
if b.get("HostIp") == "0.0.0.0":
self.ports[inner] = int(b["HostPort"])
if not self.host or not self.ports:
raise RuntimeError("Vast ещё не отдал адрес и проброс портов")
log(f"публичный адрес: {self.host}, проброшено портов: {len(self.ports)}")
self._loaded = True
return
except Exception as e:
log(f"Vast API не ответил (попытка {attempt} из 10): {e}")
time.sleep(15)
raise RuntimeError("Vast API так и не отдал проброс портов")
if not self.host:
try:
self.host = http_text("GET", "https://ifconfig.me", timeout=15).strip()
except Exception as e:
raise RuntimeError(f"не удалось определить публичный адрес: {e}")
log(f"публичный адрес: {self.host} (порты как есть)")
self._loaded = True
def url_for(self, port: int) -> str:
self.load()
return f"http://{self.host}:{self.ports.get(port, port)}"
# --- HTTP -------------------------------------------------------------------
def http_raw(method, url, body=None, headers=None, timeout=30):
data = None
headers = dict(headers or {})
if body is not None:
data = json.dumps(body).encode()
headers["Content-Type"] = "application/json"
req = urllib.request.Request(url, data=data, headers=headers, method=method)
try:
with urllib.request.urlopen(req, timeout=timeout) as r:
return r.status, r.read().decode("utf-8", "replace")
except urllib.error.HTTPError as e:
return e.code, e.read().decode("utf-8", "replace")
def http_text(method, url, **kw):
status, text = http_raw(method, url, **kw)
if status >= 400:
raise RuntimeError(f"HTTP {status}: {text[:200]}")
return text
def http_json(method, url, **kw):
status, text = http_raw(method, url, **kw)
try:
return status, json.loads(text)
except json.JSONDecodeError:
return status, {"detail": text[:400]}
# --- Менеджер ---------------------------------------------------------------
class Manager:
def __init__(self, base, key):
self.base = base
self.key = key
self.headers = {"VLLM-MANAGER-API-KEY": key}
def add(self, url, pools):
"""Возвращает (успех, пояснение).
Коды по контракту менеджера: 200 — добавлен хотя бы в один пул,
400 — ни в один (в том числе «уже существует», что нас устраивает),
403 — ключ, 422 — тело, 5xx — беда на той стороне.
"""
status, data = http_json("POST", f"{self.base}/admin/servers/add",
body={"server_url": url, "server_types": list(pools),
"validate_server": True},
headers=self.headers, timeout=60)
detail = data.get("message") or data.get("detail") or ""
if isinstance(detail, list):
detail = json.dumps(detail, ensure_ascii=False)
if status == 200:
return True, str(detail)
if status == 400 and "уже существует" in str(detail):
return True, f"уже был в пулах: {detail}"
return False, f"HTTP {status}: {detail}"
def remove(self, url, pools):
results = []
for pool in pools:
status, data = http_json("POST", f"{self.base}/admin/servers/remove",
body={"server_url": url, "server_type": pool},
headers=self.headers, timeout=30)
results.append(f"{pool}: {status}")
return results
def servers(self):
status, data = http_json("GET", f"{self.base}/admin/servers", headers=self.headers, timeout=30)
return data if status == 200 else {}
# --- Инстанс vLLM -----------------------------------------------------------
class Instance:
def __init__(self, spec, manager, public):
self.spec = spec
self.manager = manager
self.public = public
self.proc = None
self.registered = False
self.healthy = False
self.unhealthy_since = None
self.started_at = 0.0
self.restarts = 0
self.restart_at = None
self.log_path = os.path.join(LOG_DIR, f"gpu{'-'.join(map(str, spec['gpus']))}.log")
@property
def name(self):
return f"GPU{','.join(map(str, self.spec['gpus']))}:{self.spec['port']}"
@property
def url(self):
return self.public.url_for(self.spec["port"])
def command(self):
s = self.spec
cmd = [
sys.executable, "-m", "vllm.entrypoints.openai.api_server",
"--model", MODEL_ID,
"--served-model-name", SERVED_NAME,
"--host", "0.0.0.0",
"--port", str(s["port"]),
"--tensor-parallel-size", str(s["tp"]),
"--dtype", "bfloat16",
"--quantization", "compressed-tensors",
"--load-format", "safetensors",
"--max-model-len", str(s["ctx"]),
"--max-num-batched-tokens", str(s["ctx"]),
"--max-num-seqs", str(s["seqs"]),
"--gpu-memory-utilization", f"{s['util']:.2f}",
"--kv-cache-dtype", s["kv"],
"--no-enable-chunked-prefill",
"--enable-prefix-caching",
# мультимодальность у нас не используется, а VRAM под неё резервируется
"--limit-mm-per-prompt", '{"image": 0, "audio": 0, "video": 0}',
"--trust-remote-code",
"--api-key", VLLM_API_KEY,
"--default-chat-template-kwargs", '{"enable_thinking": false}',
]
if VLLM_EXTRA_ARGS:
cmd += VLLM_EXTRA_ARGS.split()
return cmd
def reap(self):
"""Освободить карту от хвостов прошлого запуска.
vLLM держит веса в отдельном процессе EngineCore. Если умирает только
сервер — а именно так выглядит любое жёсткое падение, — движок остаётся
жив и занимает почти всю VRAM. Следующий запуск тогда падает с
`Free memory on device cuda:0 ... is less than desired GPU memory
utilization`, и нода уходит в вечный цикл перезапусков. Карты этого
инстанса принадлежат только ему, поэтому чистим их целиком.
"""
for _ in range(20):
leftovers = gpu_processes(self.spec["gpus"])
if not leftovers:
return
for pid, mem in leftovers:
log(f"{self.name}: на карте висит процесс {pid} ({mem} МБ) — снимаю")
try:
os.kill(pid, signal.SIGKILL)
except (ProcessLookupError, PermissionError) as e:
log(f"{self.name}: процесс {pid} не снялся: {e}")
time.sleep(3)
log(f"{self.name}: карта так и не освободилась, запускаюсь как есть")
def start(self):
self.reap()
s = self.spec
environ = dict(os.environ)
environ.update({
"CUDA_VISIBLE_DEVICES": ",".join(map(str, s["gpus"])),
"VLLM_USE_MODELSCOPE": "False",
"VLLM_ALLOW_LONG_MAX_MODEL_LEN": "1",
"VLLM_DISABLE_CUSTOM_ALL_REDUCE": "1",
})
environ.setdefault("VLLM_USE_FLASHINFER_SAMPLER", env("VLLM_USE_FLASHINFER_SAMPLER", "1"))
if HF_TOKEN:
environ["HF_TOKEN"] = HF_TOKEN
os.makedirs(LOG_DIR, exist_ok=True)
handle = open(self.log_path, "ab", buffering=0)
self.proc = subprocess.Popen(self.command(), env=environ, stdout=handle,
stderr=subprocess.STDOUT, start_new_session=True)
self.started_at = time.time()
self.healthy = False
self.unhealthy_since = None
self.restart_at = None
log(f"{self.name}: запущен pid={self.proc.pid}, контекст {s['ctx']}, "
f"батч {s['seqs']}, kv={s['kv']}, пулы {','.join(s['pools'])}, лог {self.log_path}")
def alive(self):
return self.proc is not None and self.proc.poll() is None
def probe(self):
try:
status, _ = http_raw("GET", f"http://127.0.0.1:{self.spec['port']}/health", timeout=5)
return status == 200
except Exception:
return False
def stop(self):
if not self.alive():
self.reap()
return
log(f"{self.name}: останавливаю")
# Гасим всю группу процессов: сервер запущен через start_new_session,
# и его потомки (EngineCore с весами на карте) сами не уйдут.
try:
os.killpg(os.getpgid(self.proc.pid), signal.SIGTERM)
except (ProcessLookupError, PermissionError):
self.proc.terminate()
try:
self.proc.wait(timeout=60)
except subprocess.TimeoutExpired:
try:
os.killpg(os.getpgid(self.proc.pid), signal.SIGKILL)
except (ProcessLookupError, PermissionError):
self.proc.kill()
self.reap()
def register(self):
if not REGISTER or self.registered or manual_mode():
return
ok, detail = self.manager.add(self.url, self.spec["pools"])
if ok:
self.registered = True
log(f"{self.name}: зарегистрирован как {self.url}{detail}")
else:
log(f"{self.name}: регистрация не прошла — {detail}")
def unregister(self, why=""):
if not REGISTER or not self.registered:
return
try:
results = self.manager.remove(self.url, self.spec["pools"])
log(f"{self.name}: снят с пулов ({', '.join(results)}){' — ' + why if why else ''}")
except Exception as e:
log(f"{self.name}: снять с пулов не вышло: {e}")
self.registered = False
# --- Главный цикл -----------------------------------------------------------
class Node:
def __init__(self):
self.manager = Manager(MANAGER_URL, MANAGER_API_KEY)
self.public = PublicAddress()
self.instances = []
self.stopping = threading.Event()
def run(self):
if REGISTER and not MANAGER_API_KEY:
log("MANAGER_API_KEY не задан — регистрация выключена, ноду придётся добавить руками")
plan = build_plan()
fetch_weights()
if REGISTER:
self.public.load()
self.instances = [Instance(spec, self.manager, self.public) for spec in plan]
for inst in self.instances:
inst.start()
self.write_state()
last_sync = 0.0
while not self.stopping.is_set():
for inst in self.instances:
self.supervise(inst)
if REGISTER and manual_mode():
for inst in self.instances:
inst.unregister("включён ручной режим")
if REGISTER and not manual_mode() and time.time() - last_sync > SYNC_INTERVAL:
last_sync = time.time()
self.resync()
self.write_state()
self.stopping.wait(POLL_INTERVAL)
def supervise(self, inst):
# Перезапуск отложенный, а не через sleep: соседние карты не должны
# оставаться без присмотра, пока одна ждёт своей паузы.
if not inst.alive():
if inst.restart_at is None:
code = inst.proc.returncode if inst.proc else "?"
inst.unregister(f"процесс умер (код {code})")
inst.restarts += 1
wait = min(30 * inst.restarts, 300)
inst.restart_at = time.time() + wait
log(f"{inst.name}: процесс умер (код {code}), перезапуск через {wait} с "
f"(перезапусков: {inst.restarts}); хвост лога — {inst.log_path}")
elif time.time() >= inst.restart_at:
inst.start()
return
healthy = inst.probe()
if healthy:
if not inst.healthy:
log(f"{inst.name}: поднялся за {time.time() - inst.started_at:.0f} с")
inst.healthy = True
inst.unhealthy_since = None
inst.register()
return
# Пока модель грузится, /health молчит — это нормально и не повод паниковать.
if not inst.healthy:
if time.time() - inst.started_at > HEALTH_TIMEOUT:
log(f"{inst.name}: не поднялся за {HEALTH_TIMEOUT} с, перезапускаю")
inst.stop()
return
inst.unhealthy_since = inst.unhealthy_since or time.time()
if time.time() - inst.unhealthy_since > UNHEALTHY_GRACE:
inst.healthy = False
inst.unregister("перестал отвечать на /health")
def resync(self):
"""Менеджер мог перезапуститься и потерять нас — возвращаемся сами."""
try:
current = self.manager.servers()
except Exception as e:
log(f"сверка с менеджером не удалась: {e}")
return
if not current:
return
for inst in self.instances:
if not inst.healthy:
continue
missing = [p for p in inst.spec["pools"] if inst.url not in (current.get(p) or [])]
if missing:
log(f"{inst.name}: пропал из пулов {missing}, добавляюсь заново")
inst.registered = False
inst.register()
def write_state(self):
os.makedirs(STATE_DIR, exist_ok=True)
state = {
"updated_at": time.time(),
"model": MODEL_ID,
"manager": MANAGER_URL if REGISTER else None,
"mode": "ручной" if manual_mode() else "автоматический",
"instances": [{
"name": i.name, "url": i.url if REGISTER else None,
"port": i.spec["port"], "gpus": i.spec["gpus"], "gpu_name": i.spec["gpu_name"],
"ctx": i.spec["ctx"], "seqs": i.spec["seqs"], "kv": i.spec["kv"],
"pools": i.spec["pools"], "healthy": i.healthy, "registered": i.registered,
"restarts": i.restarts, "log": i.log_path,
} for i in self.instances],
}
tmp = os.path.join(STATE_DIR, "state.json.tmp")
with open(tmp, "w") as fh:
json.dump(state, fh, ensure_ascii=False, indent=2)
os.replace(tmp, os.path.join(STATE_DIR, "state.json"))
def shutdown(self, signum, _frame):
if self.stopping.is_set():
return
log(f"сигнал {signum}: снимаюсь с пулов и глушу инстансы")
self.stopping.set()
for inst in self.instances:
inst.unregister("нода останавливается")
for inst in self.instances:
inst.stop()
def main():
os.makedirs(LOG_DIR, exist_ok=True)
node = Node()
signal.signal(signal.SIGTERM, node.shutdown)
signal.signal(signal.SIGINT, node.shutdown)
try:
node.run()
except Exception as e:
log(f"агент падает: {type(e).__name__}: {e}")
for inst in node.instances:
inst.unregister("агент падает")
inst.stop()
raise
log("агент остановлен")
if __name__ == "__main__":
main()