Initial import: grid-bot — grid trading bot for BTC-USDT on Cifra Markets

This commit is contained in:
Kolp
2026-09-24 13:22:23 +07:00
commit 642cc11a9f
18968 changed files with 5683248 additions and 0 deletions
@@ -0,0 +1,147 @@
import logging
import multiprocessing
import socket
from typing import Literal, TYPE_CHECKING
# import for registration side effect
import torch.distributed.debug._handlers # noqa: F401
from torch._C._distributed_c10d import _WorkerServer
from torch.distributed.debug._store import get_rank, tcpstore_client
if TYPE_CHECKING:
from torch.distributed.debug._frontend import DebugHandler
__all__ = [
"start_debug_server",
"stop_debug_server",
]
logger: logging.Logger = logging.getLogger(__name__)
_WORKER_SERVER: _WorkerServer | None = None
_DEBUG_SERVER_PROC: multiprocessing.Process | None = None
def start_debug_server(
port: int = 25999,
worker_port: int = 0,
start_method: Literal["fork", "spawn", "forkserver"] | None = None,
dump_dir: str | None = None,
dump_interval: float = 60.0,
enabled_dumps: set[str] | None = None,
handlers: list["DebugHandler"] | None = None,
fetch_timeout: float = 60.0,
) -> None:
"""
Start the debug server stack on all workers. The frontend debug server is
only started on rank0 while the per rank worker servers are started on all
ranks.
This server provides an HTTP frontend that allows for debugging slow and
deadlocked distributed jobs across all ranks simultaneously. This collects
data such as stack traces, FlightRecorder events, and performance profiles.
This depends on dependencies which are not installed by default.
Dependencies:
- Jinja2
- aiohttp
WARNING: This is intended to only be used in trusted network environments.
The debug server is not designed to be secure and should not be exposed to
the public internet. See SECURITY.md for more details.
WARNING: This is an experimental feature and may change at any time.
Args:
port (int): The port to start the frontend debug server on.
worker_port (int): The port to start the worker server on. Defaults to 0, which
will cause the worker server to bind to an ephemeral port.
start_method (str | None): The multiprocessing start method to use for the
frontend server process. One of "fork", "spawn", or "forkserver".
If None, uses the default start method. Using "spawn" is recommended
when using CUDA or when fork safety is a concern.
dump_dir (str | None): Directory to write periodic debug dumps to. If None,
periodic dumping is disabled.
dump_interval (float): Seconds between periodic dumps. Defaults to 60.
enabled_dumps (set[str] | None): Set of handler dump filenames to enable
(e.g. {"stacks", "fr_trace", "tcpstore"}). If None, all handlers that
implement dump() are enabled.
handlers (list[DebugHandler] | None): List of debug handlers to use. If None,
uses the default handlers. See torch.distributed.debug._handlers for
the default handlers.
fetch_timeout (float): Timeout in seconds for fetching data from individual
workers. Defaults to 60. Workers that don't respond within this time
will be reported as unavailable.
"""
global _WORKER_SERVER, _DEBUG_SERVER_PROC
if _WORKER_SERVER is not None:
raise AssertionError("debug server already started")
if _DEBUG_SERVER_PROC is not None:
raise AssertionError("debug server already started")
logger.info("Starting debug server on port %d", port)
store = tcpstore_client()
_WORKER_SERVER = _WorkerServer("::", worker_port)
RANK = get_rank()
store.set(f"rank{RANK}", f"http://{socket.gethostname()}:{_WORKER_SERVER.port}")
if RANK == 0:
from torch.distributed.debug._debug_handlers import default_handlers
from torch.distributed.debug._frontend import main
if handlers is None:
handlers = default_handlers()
if enabled_dumps is None:
enabled_dumps = {
"stacks",
"fr_trace",
}
main_kwargs = {
"port": port,
"dump_dir": dump_dir,
"dump_interval": dump_interval,
"enabled_dumps": enabled_dumps,
"handlers": handlers,
"fetch_timeout": fetch_timeout,
}
if start_method is not None:
ctx = multiprocessing.get_context(start_method)
# pyre-ignore[16]: BaseContext has Process attribute at runtime
_DEBUG_SERVER_PROC = ctx.Process(
target=main, kwargs=main_kwargs, daemon=True
)
else:
_DEBUG_SERVER_PROC = multiprocessing.Process(
target=main, kwargs=main_kwargs, daemon=True
)
_DEBUG_SERVER_PROC.start()
def stop_debug_server() -> None:
"""
Shutdown the debug server and stop the frontend debug server process.
"""
global _WORKER_SERVER, _DEBUG_SERVER_PROC
if _DEBUG_SERVER_PROC is None:
raise AssertionError
if _WORKER_SERVER is None:
raise AssertionError
logger.info("Stopping debug server")
_DEBUG_SERVER_PROC.terminate()
_WORKER_SERVER.shutdown()
_DEBUG_SERVER_PROC.join()
_WORKER_SERVER = None
_DEBUG_SERVER_PROC = None
@@ -0,0 +1,588 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from tabulate import tabulate
from torch.distributed.debug._frontend import (
DebugHandler,
fetch_all,
format_fetch_summary,
format_json,
NavLink,
Response,
Route,
)
from torch.distributed.debug._store import tcpstore_client
from torch.distributed.flight_recorder.components.builder import build_db
from torch.distributed.flight_recorder.components.config_manager import JobConfig
from torch.distributed.flight_recorder.components.types import (
Collective,
Database,
Group,
Membership,
NCCLCall,
)
if TYPE_CHECKING:
from torch.distributed.debug._frontend import FrontendServer, HTTPRequestHandler
# ---------------------------------------------------------------------------
# Handler-specific templates
# ---------------------------------------------------------------------------
INDEX_TEMPLATE = """
{% extends "base.html" %}
{% block header %}
<h1>{% block title %}Index{% endblock %}</h1>
{% endblock %}
{% block content %}
Hi
{% endblock %}
"""
PYSPY_DUMP_TEMPLATE = """
{% extends "base.html" %}
{% block header %}
<h1>{% block title %}py-spy Stack Traces{% endblock %}</h1>
{% endblock %}
{% block content %}
<form action="" method="get">
<input type="checkbox" id="native" name="native" value="1"/>
<label for="native">Native</label>
<input type="checkbox" id="subprocesses" name="subprocesses" value="1"/>
<label for="subprocesses">Subprocesses</label>
<input type="submit" value="Submit">
</form>
{% for i, (addr, resp) in enumerate(zip(addrs, resps)) %}
<h2>Rank {{ i }}: {{ addr }}</h2>
{% if resp.status_code != 200 %}
<p>Failed to fetch: status={{ resp.status_code }}</p>
<pre>{{ resp.text }}</pre>
{% else %}
<pre>{{ resp.text }}</pre>
{% endif %}
{% endfor %}
{% endblock %}
"""
FR_TRACE_TEMPLATE = """
{% extends "base.html" %}
{% block header %}
<h1>{% block title %}{{ title }}{% endblock %}</h1>
{% endblock %}
{% block content %}
{% if fetch_summary %}<pre>{{ fetch_summary }}</pre>{% endif %}
<h2>Groups</h2>
{{ groups | safe }}
<h2>Memberships</h2>
{{ memberships | safe }}
<h2>Collectives</h2>
{{ collectives | safe }}
<h2>NCCL Calls</h2>
{{ ncclcalls | safe }}
{% endblock %}
"""
PROFILE_TEMPLATE = """
{% extends "base.html" %}
{% block header %}
<h1>{% block title %}torch.profiler{% endblock %}</h1>
{% endblock %}
{% block content %}
<form action="" method="get">
<label for="duration">Duration (seconds):</label>
<input type="number" id="duration" name="duration" value="{{ duration }}" min="1" max="60">
<input type="submit" value="Submit">
</form>
<script>
function stringToArrayBuffer(str) {
const encoder = new TextEncoder();
return encoder.encode(str).buffer;
}
async function openPerfetto(data) {
const ui = window.open('https://ui.perfetto.dev/#!/');
if (!ui) { alert('Popup blocked. Allow popups for this page and click again.'); return; }
// Perfetto readiness handshake: PING until we receive PONG
await new Promise((resolve, reject) => {
const onMsg = (e) => {
if (e.source === ui && e.data === 'PONG') {
window.removeEventListener('message', onMsg);
clearInterval(pinger);
resolve();
}
};
window.addEventListener('message', onMsg);
const pinger = setInterval(() => { try { ui.postMessage('PING', '*'); } catch (_e) {} }, 250);
setTimeout(() => { clearInterval(pinger); window.removeEventListener('message', onMsg); reject(); }, 20000);
}).catch(() => { alert('Perfetto UI did not respond. Try again.'); return; });
ui.postMessage({
perfetto: {
buffer: stringToArrayBuffer(JSON.stringify(data)),
title: "torch profiler",
fileName: "trace.json",
}
}, '*');
}
</script>
{% for i, (addr, resp) in enumerate(zip(addrs, resps)) %}
<h2>Rank {{ i }}: {{ addr }}</h2>
{% if resp.status_code != 200 %}
<p>Failed to fetch: status={{ resp.status_code }}</p>
<pre>{{ resp.text }}</pre>
{% else %}
<script>
function run{{ i }}() {
var data = {{ resp.text | safe }};
openPerfetto(data);
}
</script>
<button onclick="run{{ i }}()">View {{ i }}</button>
{% endif %}
{% endfor %}
{% endblock %}
"""
TCPSTORE_TEMPLATE = """
{% extends "base.html" %}
{% block header %}
<h1>{% block title %}TCPStore Keys{% endblock %}</h1>
{% endblock %}
{% block content %}
<pre>
{% for k, v in zip(keys, values) -%}
{{ k }}: {{ v | truncate(100) }}
{% endfor %}
</pre>
{% endblock %}
"""
# ---------------------------------------------------------------------------
# Handler classes
# ---------------------------------------------------------------------------
class IndexHandler(DebugHandler):
def routes(self) -> list[Route]:
return [Route("/", self._handle)]
def nav_links(self) -> list[NavLink]:
return [NavLink("/", "Home")]
def templates(self) -> dict[str, str]:
return {"index.html": INDEX_TEMPLATE}
def _handle(self, req: HTTPRequestHandler) -> bytes:
return req.frontend.render_template("index.html")
class StacksHandler(DebugHandler):
def routes(self) -> list[Route]:
return [Route("/stacks", self._handle)]
def nav_links(self) -> list[NavLink]:
return [NavLink("/stacks", "Python Stack Traces")]
def _handle(self, req: HTTPRequestHandler) -> bytes:
addrs, resps = fetch_all("dump_traceback", timeout=self.fetch_timeout)
return req.frontend.render_template(
"raw_resp.html", title="Stacks", addrs=addrs, resps=resps
)
def dump(self) -> str | None:
addrs, resps = fetch_all("dump_traceback", timeout=self.fetch_timeout)
parts: list[str] = []
summary = format_fetch_summary(addrs, resps)
if summary:
parts.append(summary)
parts.append("")
for i, (addr, resp) in enumerate(zip(addrs, resps)):
parts.append(f"=== Rank {i}: {addr} ===")
parts.append(
resp.text if resp.status_code == 200 else f"Error: {resp.status_code}"
)
return "\n".join(parts)
def dump_filename(self) -> str:
return "stacks"
class PySpyHandler(DebugHandler):
def routes(self) -> list[Route]:
return [Route("/pyspy_dump", self._handle)]
def nav_links(self) -> list[NavLink]:
return [NavLink("/pyspy_dump", "py-spy Stacks")]
def templates(self) -> dict[str, str]:
return {"pyspy_dump.html": PYSPY_DUMP_TEMPLATE}
def _handle(self, req: HTTPRequestHandler) -> bytes:
query = req.get_raw_query()
if "nonblocking" not in query:
query = f"nonblocking=1&{query}" if query else "nonblocking=1"
addrs, resps = fetch_all("pyspy_dump", query, timeout=self.fetch_timeout)
return req.frontend.render_template(
"pyspy_dump.html",
addrs=addrs,
resps=resps,
)
def dump(self) -> str | None:
addrs, resps = fetch_all(
"pyspy_dump", "nonblocking=1", timeout=self.fetch_timeout
)
parts: list[str] = []
summary = format_fetch_summary(addrs, resps)
if summary:
parts.append(summary)
parts.append("")
for i, (addr, resp) in enumerate(zip(addrs, resps)):
parts.append(f"=== Rank {i}: {addr} ===")
parts.append(
resp.text if resp.status_code == 200 else f"Error: {resp.status_code}"
)
return "\n".join(parts)
def dump_filename(self) -> str:
return "pyspy_dump"
class FlightRecorderHandler(DebugHandler):
def routes(self) -> list[Route]:
return [
Route("/fr_trace", self._handle_fr_trace),
Route("/fr_trace_json", self._handle_fr_trace_json),
Route("/fr_trace_nccl", self._handle_fr_trace_nccl),
Route("/fr_trace_nccl_json", self._handle_fr_trace_nccl_json),
]
def nav_links(self) -> list[NavLink]:
return [
NavLink("/fr_trace", "FlightRecorder CPU"),
NavLink("/fr_trace_json", "(JSON)"),
NavLink("/fr_trace_nccl", "FlightRecorder NCCL"),
NavLink("/fr_trace_nccl_json", "(JSON)"),
]
def templates(self) -> dict[str, str]:
return {"fr_trace.html": FR_TRACE_TEMPLATE}
@staticmethod
def _build_db(addrs: list[str], resps: list[Response]) -> Database:
config = JobConfig()
args = config.parse_args(args=[])
args.allow_incomplete_ranks = True
args.verbose = True
details = {}
for rank, resp in enumerate(resps):
if resp.status_code != 200:
continue
dump = {
"rank": rank,
"host_name": addrs[rank],
**resp.json(),
}
if "entries" not in dump:
dump["entries"] = []
details[f"rank{rank}.json"] = dump
if not details:
raise RuntimeError("All workers failed to respond")
version = next(iter(details.values()))["version"]
# pyrefly: ignore [bad-argument-type]
return build_db(details, args, version)
def _render_tables(
self, server: FrontendServer, addrs: list[str], resps: list[Response]
) -> bytes:
db = self._build_db(addrs, resps)
return server.render_template(
"fr_trace.html",
title="FlightRecorder",
fetch_summary=format_fetch_summary(addrs, resps),
groups=tabulate(db.groups, headers=Group._fields, tablefmt="html"),
memberships=tabulate(
db.memberships, headers=Membership._fields, tablefmt="html"
),
collectives=tabulate(
db.collectives, headers=Collective._fields, tablefmt="html"
),
ncclcalls=tabulate(db.ncclcalls, headers=NCCLCall._fields, tablefmt="html"),
)
def _handle_fr_trace(self, req: HTTPRequestHandler) -> bytes:
addrs, resps = fetch_all("fr_trace_json", timeout=self.fetch_timeout)
return self._render_tables(req.frontend, addrs, list(resps))
def _handle_fr_trace_json(self, req: HTTPRequestHandler) -> bytes:
addrs, resps = fetch_all("fr_trace_json", timeout=self.fetch_timeout)
return req.frontend.render_template(
"json_resp.html",
title="FlightRecorder",
addrs=addrs,
resps=resps,
)
def _handle_fr_trace_nccl(self, req: HTTPRequestHandler) -> bytes:
addrs, resps = fetch_all(
"dump_nccl_trace_json", "onlyactive=true", timeout=self.fetch_timeout
)
return self._render_tables(req.frontend, addrs, list(resps))
def _handle_fr_trace_nccl_json(self, req: HTTPRequestHandler) -> bytes:
addrs, resps = fetch_all(
"dump_nccl_trace_json", "onlyactive=true", timeout=self.fetch_timeout
)
return req.frontend.render_template(
"json_resp.html",
title="FlightRecorder NCCL",
addrs=addrs,
resps=resps,
)
def dump(self) -> str | None:
parts = []
addrs, resps = fetch_all("fr_trace_json", timeout=self.fetch_timeout)
summary = format_fetch_summary(addrs, resps)
if summary:
parts.append(summary)
parts.append("")
db = self._build_db(addrs, resps)
parts.extend(
[
"=== FR Trace ===",
"--- Groups ---",
tabulate(db.groups, headers=Group._fields, tablefmt="plain"),
"--- Memberships ---",
tabulate(db.memberships, headers=Membership._fields, tablefmt="plain"),
"--- Collectives ---",
tabulate(db.collectives, headers=Collective._fields, tablefmt="plain"),
"--- NCCL Calls ---",
tabulate(db.ncclcalls, headers=NCCLCall._fields, tablefmt="plain"),
]
)
try:
nccl_addrs, nccl_resps = fetch_all(
"dump_nccl_trace_json", "onlyactive=true", timeout=self.fetch_timeout
)
nccl_db = self._build_db(nccl_addrs, nccl_resps)
parts.extend(
[
"",
"=== FR Trace NCCL ===",
"--- Groups ---",
tabulate(nccl_db.groups, headers=Group._fields, tablefmt="plain"),
"--- Memberships ---",
tabulate(
nccl_db.memberships,
headers=Membership._fields,
tablefmt="plain",
),
"--- Collectives ---",
tabulate(
nccl_db.collectives,
headers=Collective._fields,
tablefmt="plain",
),
"--- NCCL Calls ---",
tabulate(
nccl_db.ncclcalls,
headers=NCCLCall._fields,
tablefmt="plain",
),
]
)
except Exception:
parts.append("\n=== FR Trace NCCL ===\nFailed to fetch NCCL trace")
return "\n".join(parts)
def dump_filename(self) -> str:
return "fr_trace"
class ProfilerHandler(DebugHandler):
def routes(self) -> list[Route]:
return [Route("/profile", self._handle)]
def nav_links(self) -> list[NavLink]:
return [NavLink("/profile", "torch profiler")]
def templates(self) -> dict[str, str]:
return {"profile.html": PROFILE_TEMPLATE}
def _handle(self, req: HTTPRequestHandler) -> bytes:
duration = req.get_query_arg("duration", default=1.0, type=float)
addrs, resps = fetch_all(
"torch_profile", f"duration={duration}", timeout=self.fetch_timeout
)
return req.frontend.render_template("profile.html", addrs=addrs, resps=resps)
class WaitCountersHandler(DebugHandler):
def routes(self) -> list[Route]:
return [Route("/wait_counters", self._handle)]
def nav_links(self) -> list[NavLink]:
return [NavLink("/wait_counters", "Wait Counters")]
def _handle(self, req: HTTPRequestHandler) -> bytes:
addrs, resps = fetch_all("wait_counter_values", timeout=self.fetch_timeout)
return req.frontend.render_template(
"json_resp.html", title="Wait Counters", addrs=addrs, resps=resps
)
def dump(self) -> str | None:
addrs, resps = fetch_all("wait_counter_values", timeout=self.fetch_timeout)
parts: list[str] = []
summary = format_fetch_summary(addrs, resps)
if summary:
parts.append(summary)
parts.append("")
for i, (addr, resp) in enumerate(zip(addrs, resps)):
parts.append(f"=== Rank {i}: {addr} ===")
if resp.status_code == 200:
parts.append(format_json(resp.text))
else:
parts.append(f"Error: {resp.status_code}")
return "\n".join(parts)
def dump_filename(self) -> str:
return "wait_counters"
class TCPStoreHandler(DebugHandler):
def routes(self) -> list[Route]:
return [Route("/tcpstore", self._handle)]
def nav_links(self) -> list[NavLink]:
return [NavLink("/tcpstore", "TCPStore")]
def templates(self) -> dict[str, str]:
return {"tcpstore.html": TCPSTORE_TEMPLATE}
def _handle(self, req: HTTPRequestHandler) -> bytes:
store = tcpstore_client(prefix="")
keys = store.list_keys()
keys.sort()
values = [repr(v) for v in store.multi_get(keys)]
return req.frontend.render_template("tcpstore.html", keys=keys, values=values)
def dump(self) -> str | None:
store = tcpstore_client(prefix="")
keys = store.list_keys()
keys.sort()
values = [repr(v) for v in store.multi_get(keys)]
parts = [f"{k}: {v}" for k, v in zip(keys, values)]
return "\n".join(parts)
def dump_filename(self) -> str:
return "tcpstore"
class TorchCommsFlightRecorderHandler(DebugHandler):
"""Handler for TorchComms FlightRecorder trace data."""
def routes(self) -> list[Route]:
return [
Route("/torchcomms_fr_trace", self._handle_torchcomms_fr_trace),
Route("/torchcomms_fr_trace_json", self._handle_torchcomms_fr_trace_json),
]
def nav_links(self) -> list[NavLink]:
return [
NavLink("/torchcomms_fr_trace", "TorchComms FR"),
NavLink("/torchcomms_fr_trace_json", "(JSON)"),
]
def templates(self) -> dict[str, str]:
return {"fr_trace.html": FR_TRACE_TEMPLATE}
def _handle_torchcomms_fr_trace(self, req: HTTPRequestHandler) -> bytes:
addrs, resps = fetch_all(
"torchcomms_fr_trace_json", "onlyactive=true", timeout=self.fetch_timeout
)
return self._render_tables(req.frontend, addrs, list(resps))
def _handle_torchcomms_fr_trace_json(self, req: HTTPRequestHandler) -> bytes:
addrs, resps = fetch_all(
"torchcomms_fr_trace_json", "onlyactive=true", timeout=self.fetch_timeout
)
return req.frontend.render_template(
"json_resp.html",
title="TorchComms FlightRecorder",
addrs=addrs,
resps=resps,
)
def _render_tables(
self, server: FrontendServer, addrs: list[str], resps: list[Response]
) -> bytes:
db = FlightRecorderHandler._build_db(addrs, resps)
return server.render_template(
"fr_trace.html",
title="TorchComms FlightRecorder",
fetch_summary=format_fetch_summary(addrs, resps),
groups=tabulate(db.groups, headers=Group._fields, tablefmt="html"),
memberships=tabulate(
db.memberships, headers=Membership._fields, tablefmt="html"
),
collectives=tabulate(
db.collectives, headers=Collective._fields, tablefmt="html"
),
ncclcalls=tabulate(db.ncclcalls, headers=NCCLCall._fields, tablefmt="html"),
)
def dump(self) -> str | None:
addrs, resps = fetch_all("torchcomms_fr_trace_json", timeout=self.fetch_timeout)
parts: list[str] = []
summary = format_fetch_summary(addrs, resps)
if summary:
parts.append(summary)
parts.append("")
db = FlightRecorderHandler._build_db(addrs, resps)
parts.extend(
[
"=== TorchComms FR Trace ===",
"--- Groups ---",
tabulate(db.groups, headers=Group._fields, tablefmt="plain"),
"--- Memberships ---",
tabulate(db.memberships, headers=Membership._fields, tablefmt="plain"),
"--- Collectives ---",
tabulate(db.collectives, headers=Collective._fields, tablefmt="plain"),
"--- NCCL Calls ---",
tabulate(db.ncclcalls, headers=NCCLCall._fields, tablefmt="plain"),
]
)
return "\n".join(parts)
def dump_filename(self) -> str:
return "torchcomms_fr_trace"
def default_handlers() -> list[DebugHandler]:
return [
IndexHandler(),
StacksHandler(),
PySpyHandler(),
FlightRecorderHandler(),
TorchCommsFlightRecorderHandler(),
ProfilerHandler(),
WaitCountersHandler(),
TCPStoreHandler(),
]
@@ -0,0 +1,497 @@
import asyncio
import json
import logging
import os
import socket
import threading
import time
from abc import ABC, abstractmethod
from collections.abc import Callable
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from urllib.parse import parse_qs, urlparse
from jinja2 import DictLoader, Environment
from torch.distributed.debug._store import get_world_size, tcpstore_client
logger: logging.Logger = logging.getLogger(__name__)
_DEFAULT_FETCH_TIMEOUT: float = 60.0
# ---------------------------------------------------------------------------
# Base types
# ---------------------------------------------------------------------------
@dataclass(slots=True)
class Response:
status_code: int
text: str
def raise_for_status(self):
if self.status_code != 200:
raise RuntimeError(f"HTTP {self.status_code}: {self.text}")
def json(self):
return json.loads(self.text)
@dataclass(slots=True)
class NavLink:
path: str
label: str
@dataclass(slots=True)
class Route:
path: str
handler: Callable[["HTTPRequestHandler"], bytes]
class DebugHandler(ABC):
fetch_timeout: float = _DEFAULT_FETCH_TIMEOUT
@abstractmethod
def routes(self) -> list[Route]: ...
@abstractmethod
def nav_links(self) -> list[NavLink]: ...
def templates(self) -> dict[str, str]:
return {}
def dump(self) -> str | None:
return None
def dump_filename(self) -> str:
return type(self).__name__.lower()
# ---------------------------------------------------------------------------
# Network helpers
# ---------------------------------------------------------------------------
def fetch_thread_pool(urls: list[str], timeout: float) -> list[Response]:
# late import for optional dependency
import requests
max_workers = 20
def get(url: str) -> Response:
try:
resp = requests.post(url, timeout=timeout)
return Response(resp.status_code, resp.text)
except requests.exceptions.Timeout as e:
return Response(408, f"Timeout: {e}")
except requests.exceptions.ConnectionError as e:
return Response(503, f"ConnectionError: {e}")
except Exception as e:
return Response(502, f"{type(e).__name__}: {e}")
with ThreadPoolExecutor(max_workers=max_workers) as executor:
resps = list(executor.map(get, urls))
return resps
def fetch_aiohttp(urls: list[str], timeout: float) -> list[Response]:
# late import for optional dependency
# pyrefly: ignore [missing-import]
import aiohttp
async def fetch(session: aiohttp.ClientSession, url: str) -> Response:
try:
async with session.post(url) as resp:
text = await resp.text()
return Response(resp.status, text)
except asyncio.TimeoutError as e:
return Response(408, f"TimeoutError: {e}")
except aiohttp.ClientError as e:
return Response(503, f"{type(e).__name__}: {e}")
except Exception as e:
return Response(502, f"{type(e).__name__}: {e}")
async def gather(urls: list[str]) -> list[Response]:
client_timeout = aiohttp.ClientTimeout(total=timeout)
async with aiohttp.ClientSession(timeout=client_timeout) as session:
return list(await asyncio.gather(*[fetch(session, url) for url in urls]))
return asyncio.run(gather(urls))
def fetch_all(
endpoint: str, args: str = "", *, timeout: float
) -> tuple[list[str], list[Response]]:
store = tcpstore_client()
keys = [f"rank{r}" for r in range(get_world_size())]
addrs = store.multi_get(keys)
addrs = [f"{addr.decode()}/handler/{endpoint}?{args}" for addr in addrs]
try:
resps = fetch_aiohttp(addrs, timeout=timeout)
except ImportError:
resps = fetch_thread_pool(addrs, timeout=timeout)
return addrs, resps
def format_json(blob: str):
parsed = json.loads(blob)
return json.dumps(parsed, indent=2)
def format_fetch_summary(addrs: list[str], resps: list[Response]) -> str | None:
"""Return a summary string if any workers failed, or None if all succeeded."""
failed = [(i, r) for i, r in enumerate(resps) if r.status_code != 200]
if not failed:
return None
total = len(addrs)
ok = total - len(failed)
lines = [f"PARTIAL DATA: {ok}/{total} workers responded"]
for rank, resp in failed:
lines.append(f" Rank {rank}: {resp.text}")
return "\n".join(lines)
# ---------------------------------------------------------------------------
# Template constants
# ---------------------------------------------------------------------------
BASE_TEMPLATE = """
<!doctype html>
<head>
<title>{% block title %}{% endblock %} - PyTorch Distributed</title>
<link rel="shortcut icon" type="image/x-icon" href="https://pytorch.org/favicon.ico?">
<style>
body {
margin: 0;
font-family:
-apple-system,BlinkMacSystemFont,"Segoe UI",Roboto,
"Helvetica Neue",Arial,"Noto Sans",sans-serif,"Apple Color Emoji",
"Segoe UI Emoji","Segoe UI Symbol","Noto Color Emoji";
font-size: 1rem;
font-weight: 400;
line-height: 1.5;
color: #212529;
text-align: left;
background-color: #fff;
}
h1, h2, h2, h4, h5, h6, .h1, .h2, .h2, .h4, .h5, .h6 {
margin-bottom: .5rem;
font-weight: 500;
line-height: 1.2;
}
nav {
background-color: rgba(0, 0, 0, 0.17);
padding: 10px;
display: flex;
align-items: center;
padding: 16px;
justify-content: flex-start;
}
nav h1 {
display: inline-block;
margin: 0;
}
nav a {
margin: 0 8px;
}
section {
max-width: 1280px;
padding: 16px;
margin: 0 auto;
}
pre {
white-space: pre-wrap;
max-width: 100%;
}
</style>
</head>
<nav>
<h1>Torch Distributed Debug Server</h1>
{{ nav_links | safe }}
</nav>
<section class="content">
{% block header %}{% endblock %}
{% block content %}{% endblock %}
</section>
"""
RAW_RESP_TEMPLATE = """
{% extends "base.html" %}
{% block header %}
<h1>{% block title %}{{title}}{% endblock %}</h1>
{% endblock %}
{% block content %}
{% for i, (addr, resp) in enumerate(zip(addrs, resps)) %}
<h2>Rank {{ i }}: {{ addr }}</h2>
{% if resp.status_code != 200 %}
<p>Failed to fetch: status={{ resp.status_code }}</p>
<pre>{{ resp.text }}</pre>
{% else %}
<pre>{{ resp.text }}</pre>
{% endif %}
{% endfor %}
{% endblock %}
"""
JSON_RESP_TEMPLATE = """
{% extends "base.html" %}
{% block header %}
<h1>{% block title %}{{ title }}{% endblock %}</h1>
{% endblock %}
{% block content %}
{% for i, (addr, resp) in enumerate(zip(addrs, resps)) %}
<h2>Rank {{ i }}: {{ addr }}</h2>
{% if resp.status_code != 200 %}
<p>Failed to fetch: status={{ resp.status_code }}</p>
<pre>{{ resp.text }}</pre>
{% else %}
<pre>{{ format_json(resp.text) }}</pre>
{% endif %}
{% endfor %}
{% endblock %}
"""
# ---------------------------------------------------------------------------
# PeriodicDumper
# ---------------------------------------------------------------------------
class PeriodicDumper:
def __init__(
self,
handlers: list[DebugHandler],
output_dir: str,
interval_seconds: float = 60.0,
) -> None:
self._handlers = handlers
self._output_dir = output_dir
self._interval_seconds = interval_seconds
self._stop_event = threading.Event()
self._thread: threading.Thread | None = None
def start(self) -> None:
os.makedirs(self._output_dir, exist_ok=True)
self._thread = threading.Thread(
target=self._run,
daemon=True,
name="distributed.debug.PeriodicDumper",
)
self._thread.start()
def stop(self) -> None:
self._stop_event.set()
if self._thread is not None:
self._thread.join()
def _run(self) -> None:
while not self._stop_event.is_set():
for handler in self._handlers:
try:
content = handler.dump()
except Exception:
logger.exception("Failed to dump %s", handler.dump_filename())
continue
if content is None:
continue
timestamp = time.strftime("%Y%m%d_%H%M%S")
filename = f"{handler.dump_filename()}_{timestamp}.txt"
path = os.path.join(self._output_dir, filename)
try:
with open(path, "w") as f:
f.write(content)
except Exception:
logger.exception("Failed to write dump to %s", path)
self._stop_event.wait(self._interval_seconds)
# ---------------------------------------------------------------------------
# HTTP server
# ---------------------------------------------------------------------------
class _IPv6HTTPServer(ThreadingHTTPServer):
address_family: socket.AddressFamily = socket.AF_INET6 # pyre-ignore
request_queue_size: int = 1024
class HTTPRequestHandler(BaseHTTPRequestHandler):
frontend: "FrontendServer"
def log_message(self, format, *args):
logger.info(
"%s %s",
self.client_address[0],
format % args,
)
def do_GET(self):
self.frontend._handle_request(self)
def get_path(self) -> str:
return urlparse(self.path).path
def get_query(self) -> dict[str, list[str]]:
return parse_qs(self.get_raw_query())
def get_raw_query(self) -> str:
return urlparse(self.path).query
def get_query_arg(
self, name: str, default: object = None, type: type = str
) -> object:
query = self.get_query()
if name not in query:
return default
return type(query[name][0])
class FrontendServer:
def __init__(
self,
port: int,
handlers: list[DebugHandler] | None = None,
):
if handlers is None:
from torch.distributed.debug._debug_handlers import default_handlers
handlers = default_handlers()
# Build nav HTML from handlers
nav_html = "\n".join(
f' <a href="{link.path}">{link.label}</a> <!--@lint-ignore-->'
for handler in handlers
for link in handler.nav_links()
)
# Merge all handler templates + shared templates
all_templates: dict[str, str] = {
"base.html": BASE_TEMPLATE,
"raw_resp.html": RAW_RESP_TEMPLATE,
"json_resp.html": JSON_RESP_TEMPLATE,
}
for handler in handlers:
all_templates.update(handler.templates())
loader = DictLoader(all_templates)
self._jinja_env = Environment(loader=loader, enable_async=True)
self._jinja_env.globals.update(
zip=zip,
format_json=format_json,
enumerate=enumerate,
nav_links=nav_html,
)
# Build route table from handlers
self._routes: dict[str, Callable[[HTTPRequestHandler], bytes]] = {}
for handler in handlers:
for route in handler.routes():
self._routes[route.path] = route.handler
self._handlers = handlers
# Create HTTP server
RequestHandlerClass = type(
"HTTPRequestHandler",
(HTTPRequestHandler,),
{"frontend": self},
)
server_address = ("", port)
self._server = _IPv6HTTPServer(server_address, RequestHandlerClass)
self._thread = threading.Thread(
target=self._serve,
args=(),
daemon=True,
name="distributed.debug.FrontendServer",
)
self._thread.start()
def _serve(self) -> None:
try:
self._server.serve_forever()
except Exception:
logger.exception("got exception in frontend server")
def join(self) -> None:
self._thread.join()
def _handle_request(self, req: HTTPRequestHandler) -> None:
path = req.get_path()
if path not in self._routes:
req.send_error(404, f"Handler not found: {path}")
return
handler = self._routes[path]
try:
resp = handler(req)
# Catch SystemExit to not crash when FlightRecorder errors.
except (Exception, SystemExit) as e:
logger.exception(
"Exception in frontend server when handling %s",
path,
)
req.send_error(500, f"Exception: {repr(e)}")
return
req.send_response(200)
req.send_header("Content-type", "text/html")
req.end_headers()
req.wfile.write(resp)
def render_template(self, template: str, **kwargs: object) -> bytes:
return self._jinja_env.get_template(template).render(**kwargs).encode()
def main(
port: int,
dump_dir: str | None,
dump_interval: float,
handlers: list[DebugHandler],
enabled_dumps: set[str],
fetch_timeout: float = 60.0,
) -> None:
for handler in handlers:
handler.fetch_timeout = fetch_timeout
logger.setLevel(logging.INFO)
server = FrontendServer(port=port, handlers=handlers)
logger.info("Frontend server started on port %d", server._server.server_port)
dumper: PeriodicDumper | None = None
if dump_dir is not None:
dumper = PeriodicDumper(
[
handler
for handler in handlers
if handler.dump_filename() in enabled_dumps
],
dump_dir,
dump_interval,
)
dumper.start()
logger.info(
"Periodic dumper started, writing to %s every %.0fs",
dump_dir,
dump_interval,
)
try:
server.join()
finally:
if dumper is not None:
dumper.stop()
@@ -0,0 +1,23 @@
import pathlib
import tempfile
import time
from torch._C._distributed_c10d import _register_handler, _Request, _Response
from torch.profiler import _ExperimentalConfig, profile
def _torch_profile(req: _Request, resp: _Response) -> None:
experimental_config = _ExperimentalConfig(
profile_all_threads=True,
)
duration = float(req.get_param("duration"))
with profile(record_shapes=True, experimental_config=experimental_config) as prof:
time.sleep(duration)
with tempfile.NamedTemporaryFile(prefix="torch_debug", suffix=".json") as f:
prof.export_chrome_trace(f.name)
resp.set_content(pathlib.Path(f.name).read_bytes(), "application/json")
resp.set_status(200)
_register_handler("torch_profile", _torch_profile)
@@ -0,0 +1,25 @@
import os
import torch.distributed as dist
def get_rank() -> int:
return int(os.environ["RANK"])
def get_world_size() -> int:
return int(os.environ["WORLD_SIZE"])
def tcpstore_client(prefix: str = "debug_server") -> dist.Store:
MASTER_ADDR = os.environ["MASTER_ADDR"]
MASTER_PORT = int(os.environ["MASTER_PORT"])
store = dist.TCPStore(
host_name=MASTER_ADDR,
port=MASTER_PORT,
is_master=False,
)
if prefix:
store = dist.PrefixStore(prefix, store)
return store