Initial import: grid-bot — grid trading bot for BTC-USDT on Cifra Markets
This commit is contained in:
@@ -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
|
||||
Reference in New Issue
Block a user