Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
43 changes: 33 additions & 10 deletions src/infragraph/__main__.py
Original file line number Diff line number Diff line change
@@ -1,30 +1,42 @@
import os
import stat
import typer
from infragraph.translators.translator_handler import run_translator
from infragraph.visualizer.visualize import run_visualizer

app = typer.Typer()



def _stdin_is_pipe() -> bool:
return stat.S_ISFIFO(os.fstat(0).st_mode)


@app.command()
def translate(
tool = typer.Argument(..., help="Translator to use available lstopo, nccl"),
input_file = typer.Option(None, "--input", "-i", help="Input file Path"),
output_file = typer.Option("device.yaml","--output", "-o", help="Output file path"),
output_file = typer.Option(
None, "--output", "-o",
help="Output file path. Defaults to device.yaml. If omitted and this "
"command's output is piped to another command (e.g. "
"'infragraph translate lstopo | infragraph visualize ...'), the "
"translated data is written to stdout instead.",
),
device_name = typer.Option(None, "--device-name", help="Name of the device or system being described. Required for the 'nccl' translator; inferred from the XML for 'lstopo' if not provided."),
dump = typer.Option("yaml", "--dump", help="Dump format (json or yaml)")
):
"""Translate the tools"""
run_translator(tool, input_file, output_file, dump, device_name)

@app.command()
def visualize(
input_path: str = typer.Option(
...,
None,
"--input", "-i",
help="Path to the InfraGraph infrastructure yaml/json file.",
exists=True,
file_okay=True,
dir_okay=False,
readable=True,
help="Path to the InfraGraph infrastructure yaml/json file. If "
"omitted, reads from stdin when piped in from another command "
"(e.g. 'infragraph translate lstopo | infragraph visualize "
"--output OUT_DIR').",
),
hosts: str = typer.Option(
"",
Expand All @@ -45,6 +57,17 @@ def visualize(
),
):
"""Visualize the graph"""
if input_path is None:
if not _stdin_is_pipe():
raise typer.BadParameter(
"No input provided. Pass --input FILE, or pipe data in, e.g. "
"'infragraph translate lstopo | infragraph visualize --output OUT_DIR'.",
param_hint="--input",
)
input_path = "-"
elif not os.path.isfile(input_path):
raise typer.BadParameter(f"Input file not found: {input_path}", param_hint="--input")

run_visualizer(
input_file=input_path,
hosts=hosts,
Expand Down
35 changes: 21 additions & 14 deletions src/infragraph/translators/lstopo_translator.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import os
import re
import subprocess
import sys
import tempfile
import shutil
import xml.etree.ElementTree as ET
Expand Down Expand Up @@ -627,17 +628,20 @@ def run_lstopo_parser(

input_file = str(tmp_xml)

# If output points to a directory, write a default file (devices.<format>) inside it.
if os.path.isdir(output_file) or output_file.endswith(("/", os.sep)):
output_file = os.path.join(output_file, f"devices.{dump_format.lower()}")
to_stdout = output_file == "-"

_, ext = os.path.splitext(output_file)
ext = ext.lstrip(".").lower()
if not to_stdout:
# If output points to a directory, write a default file (devices.<format>) inside it.
if os.path.isdir(output_file) or output_file.endswith(("/", os.sep)):
output_file = os.path.join(output_file, f"devices.{dump_format.lower()}")

if ext != dump_format.lower():
raise ValueError(
f"Output extension '.{ext}' does not match format '{dump_format}'."
)
_, ext = os.path.splitext(output_file)
ext = ext.lstrip(".").lower()

if ext != dump_format.lower():
raise ValueError(
f"Output extension '.{ext}' does not match format '{dump_format}'."
)

if not os.path.isfile(input_file):
raise FileNotFoundError(f"Input file not found: {input_file}")
Expand All @@ -647,15 +651,18 @@ def run_lstopo_parser(

serialized_data = device_model.serialize(dump_format)

Path(output_file).parent.mkdir(parents=True, exist_ok=True)
with open(output_file, "w", encoding="utf-8") as f:
f.write(serialized_data)
print("translated output file", output_file)
if to_stdout:
sys.stdout.write(serialized_data)
else:
Path(output_file).parent.mkdir(parents=True, exist_ok=True)
with open(output_file, "w", encoding="utf-8") as f:
f.write(serialized_data)
print("translated output file", output_file, file=sys.stderr)

# delete temp file if created
if tmp_xml and tmp_xml.exists():
tmp_xml.unlink()
print("removed /tmp/lstopo_output.xml")
print("removed /tmp/lstopo_output.xml", file=sys.stderr)



Expand Down
49 changes: 30 additions & 19 deletions src/infragraph/translators/nccl_translator.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import os
import sys
import tempfile
import xml.etree.ElementTree as ET
import ctypes
Expand Down Expand Up @@ -570,16 +571,19 @@ def run_nccl_parser(
NcclHelper.generate_nccl_topology()
input_file = str(tmp_xml)

if os.path.isdir(output_file) or output_file.endswith(("/", os.sep)):
output_file = os.path.join(output_file, f"device.{dump_format.lower()}")
to_stdout = output_file == "-"

_, ext = os.path.splitext(output_file)
ext = ext.lstrip(".").lower()
if not to_stdout:
if os.path.isdir(output_file) or output_file.endswith(("/", os.sep)):
output_file = os.path.join(output_file, f"device.{dump_format.lower()}")

if ext != dump_format.lower():
raise ValueError(
f"Output extension '.{ext}' does not match format '{dump_format}'."
)
_, ext = os.path.splitext(output_file)
ext = ext.lstrip(".").lower()

if ext != dump_format.lower():
raise ValueError(
f"Output extension '.{ext}' does not match format '{dump_format}'."
)

if not os.path.isfile(input_file):
raise FileNotFoundError(f"Input file not found: {input_file}")
Expand All @@ -589,16 +593,23 @@ def run_nccl_parser(

serialized_data = device_model.serialize(dump_format)

Path(output_file).parent.mkdir(parents=True, exist_ok=True)
with open(output_file, "w", encoding="utf-8") as f:
f.write(serialized_data)
print(f"Translated output written to: {output_file}")
req = GraphRequest()
req.infragraph.annotations.choice = "full"
annotation_output = parser.get_annotations().get_graph(req)
annotation_file = str(Path(output_file).parent / "annotated_infragraph.json")
with open(annotation_file, "w", encoding="utf-8") as f:
f.write(annotation_output)
print(f"Annotated infragraph (infrastructure + annotations) written to: {annotation_file}")
if to_stdout:
sys.stdout.write(serialized_data)
print(
"Note: annotated_infragraph.json was not written because output is stdout.",
file=sys.stderr,
)
else:
Path(output_file).parent.mkdir(parents=True, exist_ok=True)
with open(output_file, "w", encoding="utf-8") as f:
f.write(serialized_data)
print(f"Translated output written to: {output_file}", file=sys.stderr)
req = GraphRequest()
req.infragraph.annotations.choice = "full"
annotation_output = parser.get_annotations().get_graph(req)
annotation_file = str(Path(output_file).parent / "annotated_infragraph.json")
with open(annotation_file, "w", encoding="utf-8") as f:
f.write(annotation_output)
print(f"Annotated infragraph (infrastructure + annotations) written to: {annotation_file}", file=sys.stderr)
return serialized_data

28 changes: 26 additions & 2 deletions src/infragraph/translators/translator_handler.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,35 @@
import os
import stat

from infragraph.translators.lstopo_translator import run_lstopo_parser
from infragraph.translators.nccl_translator import run_nccl_parser

DEFAULT_OUTPUT_FILE = "device.yaml"


def _resolve_output_path(output_path: str | None) -> str:
"""Resolve the CLI --output value.

An explicit value (including one that happens to equal the old default)
is always used as-is. When --output is omitted, translated data goes to
stdout if this process' stdout is itself a pipe (e.g. `infragraph
translate lstopo | infragraph visualize ...`); otherwise it falls back
to device.yaml exactly as before, so standalone invocations are unchanged.
"""
if output_path is not None:
return output_path
if stat.S_ISFIFO(os.fstat(1).st_mode):
return "-"
return DEFAULT_OUTPUT_FILE

def run_translator(tool: str, input_file: str, output_path: str, dump_format: str, device_name: str) -> str:

def run_translator(tool: str, input_file: str, output_path: str | None, dump_format: str, device_name: str) -> str:
supported_translators = ["lstopo", "nccl"]
if tool not in supported_translators:
raise ValueError(f"Unsupported tool: {tool}")

output_path = _resolve_output_path(output_path)

if tool == "lstopo":
run_lstopo_parser(device_name, input_file, output_path, dump_format)

Expand All @@ -16,5 +40,5 @@ def run_translator(tool: str, input_file: str, output_path: str, dump_format: st
"Please provide it via the --device-name option."
)
run_nccl_parser(device_name, input_file, output_path, dump_format)


18 changes: 13 additions & 5 deletions src/infragraph/visualizer/visualize.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import os
import shutil
import sys
import yaml
import json
from json import JSONDecodeError
Expand Down Expand Up @@ -318,17 +319,24 @@ def _generate_instance_json(self):
def _load_infrastructure(input_file):
"""load the yaml/json file
Params:
input_file: given yaml/json file
input_file: given yaml/json file, or '-' to read from stdin
infrastructure: infrastructure object"""
if not input_file:
raise ValueError("Either input_file or infrastructure must be provided")
try:
with open(input_file, "r", encoding="utf-8") as f:
if input_file == "-":
raw = sys.stdin.read()
try:
data = json.load(f)
data = json.loads(raw)
except (JSONDecodeError, ValueError):
f.seek(0)
data = yaml.safe_load(f)
data = yaml.safe_load(raw)
else:
with open(input_file, "r", encoding="utf-8") as f:
try:
data = json.load(f)
except (JSONDecodeError, ValueError):
f.seek(0)
data = yaml.safe_load(f)
except FileNotFoundError:
raise FileNotFoundError(f"Input file not found: '{input_file}'")
except YAMLError as e:
Expand Down
Loading