Skip to content
Open
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
Original file line number Diff line number Diff line change
Expand Up @@ -25,11 +25,17 @@
BenchmarkJob,
RecommendationJob,
)
from sagemaker.serve.ai_inference_recommender.listing import (
list_benchmarks,
list_recommendations,
)
from sagemaker.serve.ai_inference_recommender.result import (
BenchmarkComparison,
BenchmarkMetric,
BenchmarkMetrics,
BenchmarkResult,
BenchmarkSearchResult,
compare_benchmarks,
)
from sagemaker.serve.ai_inference_recommender.secrets import Secret
from sagemaker.serve.ai_inference_recommender.workload import Workload
Expand All @@ -39,6 +45,7 @@


__all__ = [
"BenchmarkComparison",
"BenchmarkJob",
"BenchmarkMetric",
"BenchmarkMetrics",
Expand All @@ -51,5 +58,8 @@
"Secret",
"Workload",
"WorkloadValidationError",
"compare_benchmarks",
"list_benchmarks",
"list_recommendations",
"start_benchmark",
]
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
_fmt_number,
_format_table,
_indent,
_require_pandas,
)


Expand Down Expand Up @@ -63,9 +64,7 @@ def stats(self) -> Dict[str, float]:
return dict(self._stats)

def __repr__(self) -> str:
parts = ", ".join(
f"{stat}={_fmt_number(v)}" for stat, v in self._stats.items()
)
parts = ", ".join(f"{stat}={_fmt_number(v)}" for stat, v in self._stats.items())
unit = f" {self.unit}" if self.unit else ""
return f"<{parts}{unit}>"

Expand All @@ -82,9 +81,7 @@ class _ExpectedPerformanceView:
__slots__ = ("_by_metric",)

def __init__(self, raw_rows: Optional[List[Any]]):
by_metric: Dict[str, Dict[str, Any]] = defaultdict(
lambda: {"unit": None, "stats": {}}
)
by_metric: Dict[str, Dict[str, Any]] = defaultdict(lambda: {"unit": None, "stats": {}})
for row in raw_rows or []:
metric = getattr(row, "metric", None)
if not metric:
Expand Down Expand Up @@ -139,9 +136,9 @@ def __len__(self) -> int:
return len(self._by_metric)

def __repr__(self) -> str:
return "{" + ", ".join(
f"{name}: {metric!r}" for name, metric in self._by_metric.items()
) + "}"
return (
"{" + ", ".join(f"{name}: {metric!r}" for name, metric in self._by_metric.items()) + "}"
)


def _to_float(value):
Expand Down Expand Up @@ -195,7 +192,6 @@ def __getattr__(self, name):
def __str__(self) -> str:
md = getattr(self._raw, "model_details", None)
dc = getattr(self._raw, "deployment_configuration", None)
ep = getattr(self._raw, "expected_performance", None) or []

config_lines = [
f"instance_type: {_safe_str(dc, 'instance_type')}",
Expand All @@ -217,14 +213,16 @@ def __str__(self) -> str:
env_lines = [f" {k} = {v}" for k, v in items]
env_block = "\nenv vars ({0}):\n{1}".format(len(items), "\n".join(env_lines))

perf_rows = []
for m in ep:
perf_rows.append([
_safe_str(m, "metric"),
_safe_str(m, "stat"),
_fmt_number(_safe_float(m, "value")),
_safe_str(m, "unit"),
])
# Same records to_dataframe() uses, formatted for display here.
perf_rows = [
[
rec["metric"] if rec["metric"] not in (None, "") else "-",
rec["stat"] if rec["stat"] not in (None, "") else "-",
_fmt_number(rec["value"]),
rec["unit"] if rec["unit"] not in (None, "") else "-",
]
for rec in self._perf_records()
]
perf_table = _format_table(
headers=["metric", "stat", "value", "unit"],
rows=perf_rows,
Expand All @@ -249,6 +247,30 @@ def _repr_pretty_(self, p, cycle):
# Render the full table in notebooks (Jupyter uses this hook).
p.text("..." if cycle else str(self))

def _perf_records(self) -> List[Dict[str, Any]]:
"""(metric, stat, value, unit) records for this row's expected
performance — the rows of the printed table, one per (metric, stat)."""
ep = getattr(self._raw, "expected_performance", None) or []
return [
{
"metric": getattr(m, "metric", None),
"stat": getattr(m, "stat", None),
"value": _safe_float(m, "value"),
"unit": getattr(m, "unit", None),
}
for m in ep
]

def to_dataframe(self):
"""Return this recommendation's expected performance as a pandas
``DataFrame`` — the same ``metric``/``stat``/``value``/``unit`` rows the
printed ``expected performance`` table shows, one row per (metric, stat).

Requires pandas.
"""
pd = _require_pandas()
return pd.DataFrame(self._perf_records(), columns=["metric", "stat", "value", "unit"])


def _safe_str(obj, attr) -> str:
if obj is None:
Expand Down Expand Up @@ -291,48 +313,105 @@ def _repr_pretty_(self, p, cycle):
# Render the full table in notebooks (Jupyter uses this hook).
p.text("..." if cycle else str(self))

# Column labels for the comparative table / DataFrame, in display order.
_TABLE_COLUMNS = (
"idx",
"spec_name",
"instance_type",
"instances",
"copies/inst",
"container",
"req/s",
"tok/s",
"lat_p50",
"lat_p90",
"lat_p99",
"ttft_p50",
"itl_p50",
)

def _row_records(self) -> List[Dict[str, Any]]:
"""One record per row, keyed by ``_TABLE_COLUMNS``, with native values.

Shared by ``__str__`` and ``to_dataframe()`` so they cannot drift.
"""
records = []
for view in self:
dc = getattr(view.raw, "deployment_configuration", None)
ep = view.expected_performance
records.append(
{
"idx": view._index,
"spec_name": view.recommendation_spec_name,
"instance_type": getattr(dc, "instance_type", None) if dc else None,
"instances": getattr(dc, "instance_count", None) if dc else None,
"copies/inst": (getattr(dc, "copy_count_per_instance", None) if dc else None),
# None when absent (not "-"), so the DataFrame keeps it as
# missing; __str__ renders the dash.
"container": (
_short_container_tag(dc.image_uri)
if dc is not None and getattr(dc, "image_uri", None)
else None
),
"req/s": _get_metric_stat(ep, "request_throughput", "avg"),
"tok/s": _get_metric_stat(ep, "output_token_throughput", "avg"),
"lat_p50": _get_metric_stat(ep, "request_latency", "p50"),
"lat_p90": _get_metric_stat(ep, "request_latency", "p90"),
"lat_p99": _get_metric_stat(ep, "request_latency", "p99"),
"ttft_p50": _get_metric_stat(ep, "time_to_first_token", "p50"),
"itl_p50": _get_metric_stat(ep, "inter_token_latency", "p50"),
}
)
return records

def __str__(self) -> str:
if not self:
return "Recommendations[0] (no rows)"

_NUMERIC = {"req/s", "tok/s", "lat_p50", "lat_p90", "lat_p99", "ttft_p50", "itl_p50"}
rows = []
for view in self:
dc = getattr(view.raw, "deployment_configuration", None)
ep = view.expected_performance
rows.append([
f"[{view._index}]",
view.recommendation_spec_name or "-",
_safe_str(dc, "instance_type"),
_safe_str(dc, "instance_count"),
_safe_str(dc, "copy_count_per_instance"),
_short_container_tag(_safe_str(dc, "image_uri")),
_fmt_number(_get_metric_stat(ep, "request_throughput", "avg")),
_fmt_number(_get_metric_stat(ep, "output_token_throughput", "avg")),
_fmt_number(_get_metric_stat(ep, "request_latency", "p50")),
_fmt_number(_get_metric_stat(ep, "request_latency", "p90")),
_fmt_number(_get_metric_stat(ep, "request_latency", "p99")),
_fmt_number(_get_metric_stat(ep, "time_to_first_token", "p50")),
_fmt_number(_get_metric_stat(ep, "inter_token_latency", "p50")),
])

table = _format_table(
headers=[
"idx", "spec_name", "instance_type",
"instances", "copies/inst",
"container",
"req/s", "tok/s",
"lat_p50", "lat_p90", "lat_p99",
"ttft_p50", "itl_p50",
],
rows=rows,
)
for rec in self._row_records():
row = []
for col in self._TABLE_COLUMNS:
value = rec[col]
if col == "idx":
row.append(f"[{value}]")
elif col in _NUMERIC:
row.append(_fmt_number(value))
else:
row.append("-" if value in (None, "") else str(value))
rows.append(row)

table = _format_table(headers=list(self._TABLE_COLUMNS), rows=rows)

return (
f"Recommendations[{len(self)}] (.best = top row; index by [N] for full detail)\n"
f"{table}\n"
f"lat/ttft/itl in ms; req/s = requests/sec; tok/s = output tokens/sec"
)

def to_dataframe(self):
"""Return the recommendations as a pandas ``DataFrame`` — one row per
recommendation, columns matching the printed comparative table
(``instance_type``, ``instances``, ``req/s``, ``lat_p50``, ...), indexed
by the recommendation index ``idx``.

Latency columns are milliseconds; ``req/s`` is requests/sec and ``tok/s``
is output tokens/sec, same as the printed table's footnote. Numeric
columns hold real numbers (``NaN`` where a metric is absent), not

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

F7 (major): when a metric is absent from every row, the column is object dtype holding None, not NaN — and sorting raises.

pandas infers object from an all-None list, so the promise here does not hold for the sparse case. Verified on two rows with no expected_performance:

dtypes:  req/s object, tok/s object, lat_p50 object, lat_p90 object,
         lat_p99 object, ttft_p50 object, itl_p50 object
>>> df.nlargest(1, "req/s")
TypeError: Column 'req/s' has dtype object, cannot use method 'nlargest' with this dtype

BenchmarkComparison.to_dataframe's docstring makes the same claim about deltas (result.py:724-725) and has the same behaviour when every delta is undefined.

This matters because sorting, nlargest, and .mean() are the reason a DataFrame was asked for, and the columns most likely to be entirely absent are exactly the ones a customer would rank on. test_numeric_columns_are_native_and_missing_is_nan asserts pd.isna() on a partially populated column, which pandas does coerce to float64 — so it passes while this case is broken.

Suggested direction: coerce the numeric columns explicitly, e.g. .astype({c: "float64" for c in numeric_cols}) or pd.to_numeric(..., errors="coerce") after construction, and add a case where a metric is missing from every row.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This still needs to be addressed.

preformatted strings.

Requires pandas.
"""
pd = _require_pandas()
columns = [c for c in self._TABLE_COLUMNS if c != "idx"]
records = self._row_records()
return pd.DataFrame(
[{c: rec[c] for c in columns} for rec in records],
columns=columns,
index=pd.Index([rec["idx"] for rec in records], name="idx"),
)


def _get_metric_stat(ep_view, metric_name: str, stat: str):
try:
Expand Down
Loading
Loading