Beautify text output in dump comparator (#19683)
This commit is contained in:
@@ -3,10 +3,13 @@ from __future__ import annotations
|
|||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from io import StringIO
|
from io import StringIO
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Optional
|
from typing import TYPE_CHECKING, Any, Optional
|
||||||
|
|
||||||
import polars as pl
|
import polars as pl
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from rich.table import Table
|
||||||
|
|
||||||
from sglang.srt.debug_utils.comparator.output_types import (
|
from sglang.srt.debug_utils.comparator.output_types import (
|
||||||
InputIdsRecord,
|
InputIdsRecord,
|
||||||
RankInfoRecord,
|
RankInfoRecord,
|
||||||
@@ -39,6 +42,21 @@ def emit_display_records(
|
|||||||
|
|
||||||
def _render_polars_as_text(df: pl.DataFrame, *, title: Optional[str] = None) -> str:
|
def _render_polars_as_text(df: pl.DataFrame, *, title: Optional[str] = None) -> str:
|
||||||
from rich.console import Console
|
from rich.console import Console
|
||||||
|
|
||||||
|
table = _build_rich_table(df, title=title)
|
||||||
|
|
||||||
|
buf = StringIO()
|
||||||
|
Console(file=buf, force_terminal=False, width=200).print(table)
|
||||||
|
return buf.getvalue().rstrip("\n")
|
||||||
|
|
||||||
|
|
||||||
|
def _render_polars_as_rich_table(
|
||||||
|
df: pl.DataFrame, *, title: Optional[str] = None
|
||||||
|
) -> "Table":
|
||||||
|
return _build_rich_table(df, title=title)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_rich_table(df: pl.DataFrame, *, title: Optional[str] = None) -> "Table":
|
||||||
from rich.table import Table
|
from rich.table import Table
|
||||||
|
|
||||||
table = Table(title=title)
|
table = Table(title=title)
|
||||||
@@ -47,9 +65,7 @@ def _render_polars_as_text(df: pl.DataFrame, *, title: Optional[str] = None) ->
|
|||||||
for row in df.iter_rows():
|
for row in df.iter_rows():
|
||||||
table.add_row(*[str(v) for v in row])
|
table.add_row(*[str(v) for v in row])
|
||||||
|
|
||||||
buf = StringIO()
|
return table
|
||||||
Console(file=buf, force_terminal=False, width=200).print(table)
|
|
||||||
return buf.getvalue().rstrip("\n")
|
|
||||||
|
|
||||||
|
|
||||||
def _collect_rank_info(
|
def _collect_rank_info(
|
||||||
|
|||||||
@@ -0,0 +1,287 @@
|
|||||||
|
"""Formatting functions for comparator output records.
|
||||||
|
|
||||||
|
Extracted from output_types.py to separate data-structure definitions
|
||||||
|
from rendering / formatting logic.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from rich.console import Group
|
||||||
|
from rich.markup import escape
|
||||||
|
from rich.panel import Panel
|
||||||
|
|
||||||
|
from sglang.srt.debug_utils.comparator.tensor_comparator.formatter import (
|
||||||
|
format_comparison,
|
||||||
|
format_replicated_checks,
|
||||||
|
)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from rich.console import RenderableType
|
||||||
|
|
||||||
|
from sglang.srt.debug_utils.comparator.aligner.entrypoint.traced_types import (
|
||||||
|
TracedAlignerPlan,
|
||||||
|
TracedSubPlan,
|
||||||
|
)
|
||||||
|
from sglang.srt.debug_utils.comparator.aligner.entrypoint.types import AlignerPlan
|
||||||
|
from sglang.srt.debug_utils.comparator.output_types import (
|
||||||
|
ConfigRecord,
|
||||||
|
ErrorLog,
|
||||||
|
InfoLog,
|
||||||
|
LogRecord,
|
||||||
|
NonTensorComparisonRecord,
|
||||||
|
SkipComparisonRecord,
|
||||||
|
SummaryRecord,
|
||||||
|
TensorComparisonRecord,
|
||||||
|
_OutputRecord,
|
||||||
|
_TableRecord,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ── Record-level rendering (body + logs) ─────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _render_record_rich(record: _OutputRecord) -> RenderableType:
|
||||||
|
body: RenderableType = record._format_rich_body()
|
||||||
|
|
||||||
|
log_lines: list[str] = _format_log_lines_rich(
|
||||||
|
errors=record.errors, infos=record.infos
|
||||||
|
)
|
||||||
|
|
||||||
|
if not log_lines:
|
||||||
|
return body
|
||||||
|
|
||||||
|
log_block: str = "\n".join(log_lines)
|
||||||
|
if isinstance(body, str):
|
||||||
|
return body + "\n" + log_block
|
||||||
|
return Group(body, log_block)
|
||||||
|
|
||||||
|
|
||||||
|
def _render_record_text(record: _OutputRecord) -> str:
|
||||||
|
body: str = record._format_body()
|
||||||
|
|
||||||
|
log_suffix: str = _format_log_lines_text(errors=record.errors, infos=record.infos)
|
||||||
|
|
||||||
|
if log_suffix:
|
||||||
|
body += "\n" + log_suffix
|
||||||
|
|
||||||
|
return body
|
||||||
|
|
||||||
|
|
||||||
|
def _format_log_lines_rich(
|
||||||
|
*, errors: list[ErrorLog], infos: list[InfoLog]
|
||||||
|
) -> list[str]:
|
||||||
|
lines: list[str] = []
|
||||||
|
|
||||||
|
if errors:
|
||||||
|
lines.extend(f" [red]✗ {e.to_text()}[/]" for e in errors)
|
||||||
|
if infos:
|
||||||
|
lines.extend(f" [dim]ℹ {i.to_text()}[/]" for i in infos)
|
||||||
|
|
||||||
|
return lines
|
||||||
|
|
||||||
|
|
||||||
|
def _format_log_lines_text(*, errors: list[ErrorLog], infos: list[InfoLog]) -> str:
|
||||||
|
lines: list[str] = []
|
||||||
|
|
||||||
|
if errors:
|
||||||
|
lines.extend(f" ✗ {e.to_text()}" for e in errors)
|
||||||
|
if infos:
|
||||||
|
lines.extend(f" ℹ {i.to_text()}" for i in infos)
|
||||||
|
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
# ── ConfigRecord ──────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _format_config_body(record: ConfigRecord) -> str:
|
||||||
|
return f"Config: {record.config}"
|
||||||
|
|
||||||
|
|
||||||
|
def _format_config_rich_body(record: ConfigRecord) -> RenderableType:
|
||||||
|
lines: list[str] = [f" [bold]{k}[/] : {v}" for k, v in record.config.items()]
|
||||||
|
return Panel("\n".join(lines), title="Comparator Config", border_style="cyan")
|
||||||
|
|
||||||
|
|
||||||
|
# ── SkipComparisonRecord ─────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _format_skip_body(record: SkipComparisonRecord) -> str:
|
||||||
|
return f"Skip: {record.name}{record._format_location_suffix()} ({record.reason})"
|
||||||
|
|
||||||
|
|
||||||
|
def _format_skip_rich_body(record: SkipComparisonRecord) -> RenderableType:
|
||||||
|
suffix: str = record._format_location_suffix()
|
||||||
|
return (
|
||||||
|
f"[dim]⊘ {escape(record.name)}{suffix} ── skipped ({escape(record.reason)})[/]"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ── _TableRecord ─────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _format_table_body(record: _TableRecord) -> str:
|
||||||
|
import polars as pl
|
||||||
|
|
||||||
|
from sglang.srt.debug_utils.comparator.display import _render_polars_as_text
|
||||||
|
|
||||||
|
return _render_polars_as_text(
|
||||||
|
pl.DataFrame(record.rows), title=record._table_title()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _format_table_rich_body(record: _TableRecord) -> RenderableType:
|
||||||
|
import polars as pl
|
||||||
|
|
||||||
|
from sglang.srt.debug_utils.comparator.display import (
|
||||||
|
_render_polars_as_rich_table,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _render_polars_as_rich_table(
|
||||||
|
pl.DataFrame(record.rows), title=record._table_title()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ── TensorComparisonRecord ───────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _format_tensor_comparison_body(record: TensorComparisonRecord) -> str:
|
||||||
|
body: str = record._format_location_prefix() + format_comparison(record)
|
||||||
|
if record.replicated_checks:
|
||||||
|
body += "\n" + format_replicated_checks(record.replicated_checks)
|
||||||
|
if record.traced_plan is not None:
|
||||||
|
body += "\n" + _format_aligner_plan(record.traced_plan)
|
||||||
|
return body
|
||||||
|
|
||||||
|
|
||||||
|
def _format_tensor_comparison_rich_body(
|
||||||
|
record: TensorComparisonRecord,
|
||||||
|
) -> RenderableType:
|
||||||
|
from sglang.srt.debug_utils.comparator.tensor_comparator.formatter import (
|
||||||
|
format_comparison_rich,
|
||||||
|
)
|
||||||
|
|
||||||
|
return record._format_location_prefix_rich() + format_comparison_rich(record=record)
|
||||||
|
|
||||||
|
|
||||||
|
# ── NonTensorComparisonRecord ────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _format_non_tensor_body(record: NonTensorComparisonRecord) -> str:
|
||||||
|
suffix: str = record._format_location_suffix()
|
||||||
|
if record.values_equal:
|
||||||
|
return f"NonTensor: {record.name}{suffix} = {record.baseline_value} ({record.baseline_type}) [equal]"
|
||||||
|
return (
|
||||||
|
f"NonTensor: {record.name}{suffix}\n"
|
||||||
|
f" baseline = {record.baseline_value} ({record.baseline_type})\n"
|
||||||
|
f" target = {record.target_value} ({record.target_type})"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _format_non_tensor_rich_body(
|
||||||
|
record: NonTensorComparisonRecord,
|
||||||
|
) -> RenderableType:
|
||||||
|
suffix: str = record._format_location_suffix()
|
||||||
|
name: str = escape(record.name)
|
||||||
|
baseline_val: str = escape(record.baseline_value)
|
||||||
|
target_val: str = escape(record.target_value)
|
||||||
|
|
||||||
|
if record.values_equal:
|
||||||
|
return (
|
||||||
|
f"═ {name}{suffix} = {baseline_val} "
|
||||||
|
f"({record.baseline_type}) [green]✓[/]"
|
||||||
|
)
|
||||||
|
return (
|
||||||
|
f"═ [bold red]{name}{suffix}[/]\n"
|
||||||
|
f" baseline = {baseline_val} ({record.baseline_type})\n"
|
||||||
|
f" target = {target_val} ({record.target_type})"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ── SummaryRecord ────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _format_summary_body(record: SummaryRecord) -> str:
|
||||||
|
return (
|
||||||
|
f"Summary: {record.passed} passed, {record.failed} failed, "
|
||||||
|
f"{record.skipped} skipped (total {record.total})"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _format_summary_rich_body(record: SummaryRecord) -> RenderableType:
|
||||||
|
text: str = (
|
||||||
|
f"[bold green]{record.passed} passed[/] │ "
|
||||||
|
f"[bold red]{record.failed} failed[/] │ "
|
||||||
|
f"[yellow]{record.skipped} skipped[/] │ "
|
||||||
|
f"{record.total} total"
|
||||||
|
)
|
||||||
|
return Panel(text, title="SUMMARY", border_style="bold")
|
||||||
|
|
||||||
|
|
||||||
|
# ── LogRecord ────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _format_log_body(record: LogRecord) -> str:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
# ── Standalone helpers ───────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _format_aligner_plan(traced_plan: TracedAlignerPlan) -> str:
|
||||||
|
lines: list[str] = ["Aligner Plan:"]
|
||||||
|
|
||||||
|
for side_label, traced_side in [
|
||||||
|
("baseline", traced_plan.per_side.x),
|
||||||
|
("target", traced_plan.per_side.y),
|
||||||
|
]:
|
||||||
|
if not traced_side.step_plans:
|
||||||
|
lines.append(f" {side_label}: (no steps)")
|
||||||
|
continue
|
||||||
|
|
||||||
|
step_summaries: list[str] = []
|
||||||
|
for traced_step in traced_side.step_plans:
|
||||||
|
sub_strs: list[str] = [
|
||||||
|
_format_sub_plan_text(traced_sub)
|
||||||
|
for traced_sub in traced_step.sub_plans
|
||||||
|
]
|
||||||
|
summary: str = ", ".join(sub_strs) if sub_strs else "passthrough"
|
||||||
|
step_summaries.append(f"step={traced_step.step}: {summary}")
|
||||||
|
lines.append(f" {side_label}: [{'; '.join(step_summaries)}]")
|
||||||
|
|
||||||
|
lines.extend(_format_cross_side_plan_text(traced_plan.plan))
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
def _format_sub_plan_text(traced_sub: TracedSubPlan) -> str:
|
||||||
|
sub_desc: str = f"{traced_sub.plan.type}"
|
||||||
|
|
||||||
|
if traced_sub.snapshot is not None:
|
||||||
|
snap = traced_sub.snapshot
|
||||||
|
in_count: int = len(snap.input_shapes)
|
||||||
|
out_count: int = len(snap.output_shapes)
|
||||||
|
in_shape: str = str(snap.input_shapes[0]) if snap.input_shapes else "?"
|
||||||
|
out_shape: str = str(snap.output_shapes[0]) if snap.output_shapes else "?"
|
||||||
|
sub_desc += f" {in_count}x{in_shape} -> {out_count}x{out_shape}"
|
||||||
|
|
||||||
|
return sub_desc
|
||||||
|
|
||||||
|
|
||||||
|
def _format_cross_side_plan_text(plan: AlignerPlan) -> list[str]:
|
||||||
|
lines: list[str] = []
|
||||||
|
|
||||||
|
if plan.token_aligner_plan is not None:
|
||||||
|
num_tokens: int = len(plan.token_aligner_plan.locators.x.steps)
|
||||||
|
lines.append(f" token_aligner: {num_tokens} tokens aligned")
|
||||||
|
|
||||||
|
if plan.axis_aligner_plan is not None:
|
||||||
|
parts: list[str] = []
|
||||||
|
if plan.axis_aligner_plan.pattern.x:
|
||||||
|
parts.append(f"x: {plan.axis_aligner_plan.pattern.x}")
|
||||||
|
if plan.axis_aligner_plan.pattern.y:
|
||||||
|
parts.append(f"y: {plan.axis_aligner_plan.pattern.y}")
|
||||||
|
lines.append(f" axis_aligner: {', '.join(parts)}")
|
||||||
|
|
||||||
|
return lines
|
||||||
@@ -3,15 +3,10 @@ from __future__ import annotations
|
|||||||
from abc import abstractmethod
|
from abc import abstractmethod
|
||||||
from typing import TYPE_CHECKING, Annotated, Any, Literal, Optional, Union
|
from typing import TYPE_CHECKING, Annotated, Any, Literal, Optional, Union
|
||||||
|
|
||||||
import polars as pl
|
|
||||||
from pydantic import ConfigDict, Discriminator, Field, TypeAdapter, model_validator
|
from pydantic import ConfigDict, Discriminator, Field, TypeAdapter, model_validator
|
||||||
from rich.console import RenderableType
|
from rich.console import RenderableType
|
||||||
from rich.markup import escape
|
from rich.markup import escape
|
||||||
|
|
||||||
from sglang.srt.debug_utils.comparator.tensor_comparator.formatter import (
|
|
||||||
format_comparison,
|
|
||||||
format_replicated_checks,
|
|
||||||
)
|
|
||||||
from sglang.srt.debug_utils.comparator.tensor_comparator.types import (
|
from sglang.srt.debug_utils.comparator.tensor_comparator.types import (
|
||||||
DiffInfo,
|
DiffInfo,
|
||||||
TensorComparisonInfo,
|
TensorComparisonInfo,
|
||||||
@@ -21,7 +16,6 @@ from sglang.srt.debug_utils.comparator.utils import Pair, _StrictBase
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.debug_utils.comparator.aligner.entrypoint.traced_types import (
|
from sglang.srt.debug_utils.comparator.aligner.entrypoint.traced_types import (
|
||||||
TracedAlignerPlan,
|
TracedAlignerPlan,
|
||||||
TracedSubPlan,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.debug_utils.comparator.aligner.entrypoint.types import AlignerPlan
|
from sglang.srt.debug_utils.comparator.aligner.entrypoint.types import AlignerPlan
|
||||||
from sglang.srt.debug_utils.comparator.report_sink import Verbosity
|
from sglang.srt.debug_utils.comparator.report_sink import Verbosity
|
||||||
@@ -89,19 +83,45 @@ class _OutputRecord(_StrictBase):
|
|||||||
@abstractmethod
|
@abstractmethod
|
||||||
def _format_body(self) -> str: ...
|
def _format_body(self) -> str: ...
|
||||||
|
|
||||||
def _format_rich_body(self, verbosity: Verbosity = "normal") -> RenderableType:
|
def _format_rich_body(self) -> RenderableType:
|
||||||
return self._format_body()
|
return self._format_body()
|
||||||
|
|
||||||
def to_rich(self, verbosity: Verbosity = "normal") -> RenderableType:
|
def to_rich(self) -> RenderableType:
|
||||||
return self._format_body()
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
|
_render_record_rich,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _render_record_rich(self)
|
||||||
|
|
||||||
def to_text(self) -> str:
|
def to_text(self) -> str:
|
||||||
body = self._format_body()
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
if self.errors:
|
_render_record_text,
|
||||||
body += "\n" + "\n".join(f" ✗ {e.to_text()}" for e in self.errors)
|
)
|
||||||
if self.infos:
|
|
||||||
body += "\n" + "\n".join(f" ℹ {i.to_text()}" for i in self.infos)
|
return _render_record_text(self)
|
||||||
return body
|
|
||||||
|
|
||||||
|
class RecordLocation(_StrictBase):
|
||||||
|
step: Optional[int] = None
|
||||||
|
|
||||||
|
|
||||||
|
class _BaseComparisonRecord(_OutputRecord):
|
||||||
|
location: RecordLocation = Field(default_factory=RecordLocation)
|
||||||
|
|
||||||
|
def _format_location_prefix(self) -> str:
|
||||||
|
if self.location.step is not None:
|
||||||
|
return f"[step={self.location.step}] "
|
||||||
|
return ""
|
||||||
|
|
||||||
|
def _format_location_prefix_rich(self) -> str:
|
||||||
|
if self.location.step is not None:
|
||||||
|
return escape(f"[step={self.location.step}]") + " "
|
||||||
|
return ""
|
||||||
|
|
||||||
|
def _format_location_suffix(self) -> str:
|
||||||
|
if self.location.step is not None:
|
||||||
|
return f" (step={self.location.step})"
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
class RecordLocation(_StrictBase):
|
class RecordLocation(_StrictBase):
|
||||||
@@ -132,7 +152,18 @@ class ConfigRecord(_OutputRecord):
|
|||||||
config: dict[str, Any]
|
config: dict[str, Any]
|
||||||
|
|
||||||
def _format_body(self) -> str:
|
def _format_body(self) -> str:
|
||||||
return f"Config: {self.config}"
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
|
_format_config_body,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _format_config_body(self)
|
||||||
|
|
||||||
|
def _format_rich_body(self) -> RenderableType:
|
||||||
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
|
_format_config_rich_body,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _format_config_rich_body(self)
|
||||||
|
|
||||||
|
|
||||||
class SkipComparisonRecord(_BaseComparisonRecord):
|
class SkipComparisonRecord(_BaseComparisonRecord):
|
||||||
@@ -147,7 +178,18 @@ class SkipComparisonRecord(_BaseComparisonRecord):
|
|||||||
return "skipped"
|
return "skipped"
|
||||||
|
|
||||||
def _format_body(self) -> str:
|
def _format_body(self) -> str:
|
||||||
return f"Skip: {self.name}{self._format_location_suffix()} ({self.reason})"
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
|
_format_skip_body,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _format_skip_body(self)
|
||||||
|
|
||||||
|
def _format_rich_body(self) -> RenderableType:
|
||||||
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
|
_format_skip_rich_body,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _format_skip_rich_body(self)
|
||||||
|
|
||||||
|
|
||||||
class _TableRecord(_OutputRecord):
|
class _TableRecord(_OutputRecord):
|
||||||
@@ -158,12 +200,19 @@ class _TableRecord(_OutputRecord):
|
|||||||
def _table_title(self) -> str: ...
|
def _table_title(self) -> str: ...
|
||||||
|
|
||||||
def _format_body(self) -> str:
|
def _format_body(self) -> str:
|
||||||
from sglang.srt.debug_utils.comparator.display import _render_polars_as_text
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
|
_format_table_body,
|
||||||
return _render_polars_as_text(
|
|
||||||
pl.DataFrame(self.rows), title=self._table_title()
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
return _format_table_body(self)
|
||||||
|
|
||||||
|
def _format_rich_body(self) -> RenderableType:
|
||||||
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
|
_format_table_rich_body,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _format_table_rich_body(self)
|
||||||
|
|
||||||
|
|
||||||
class RankInfoRecord(_TableRecord):
|
class RankInfoRecord(_TableRecord):
|
||||||
type: Literal["rank_info"] = "rank_info"
|
type: Literal["rank_info"] = "rank_info"
|
||||||
@@ -196,12 +245,18 @@ class TensorComparisonRecord(TensorComparisonInfo, _BaseComparisonRecord):
|
|||||||
return "passed" if self.diff is not None and self.diff.passed else "failed"
|
return "passed" if self.diff is not None and self.diff.passed else "failed"
|
||||||
|
|
||||||
def _format_body(self) -> str:
|
def _format_body(self) -> str:
|
||||||
body: str = self._format_location_prefix() + format_comparison(self)
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
if self.replicated_checks:
|
_format_tensor_comparison_body,
|
||||||
body += "\n" + format_replicated_checks(self.replicated_checks)
|
)
|
||||||
if self.traced_plan is not None:
|
|
||||||
body += "\n" + _format_aligner_plan(self.traced_plan)
|
return _format_tensor_comparison_body(self)
|
||||||
return body
|
|
||||||
|
def _format_rich_body(self) -> RenderableType:
|
||||||
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
|
_format_tensor_comparison_rich_body,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _format_tensor_comparison_rich_body(self)
|
||||||
|
|
||||||
|
|
||||||
class NonTensorComparisonRecord(_BaseComparisonRecord):
|
class NonTensorComparisonRecord(_BaseComparisonRecord):
|
||||||
@@ -220,15 +275,19 @@ class NonTensorComparisonRecord(_BaseComparisonRecord):
|
|||||||
return "passed" if self.values_equal else "failed"
|
return "passed" if self.values_equal else "failed"
|
||||||
|
|
||||||
def _format_body(self) -> str:
|
def _format_body(self) -> str:
|
||||||
suffix: str = self._format_location_suffix()
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
if self.values_equal:
|
_format_non_tensor_body,
|
||||||
return f"NonTensor: {self.name}{suffix} = {self.baseline_value} ({self.baseline_type}) [equal]"
|
|
||||||
return (
|
|
||||||
f"NonTensor: {self.name}{suffix}\n"
|
|
||||||
f" baseline = {self.baseline_value} ({self.baseline_type})\n"
|
|
||||||
f" target = {self.target_value} ({self.target_type})"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
return _format_non_tensor_body(self)
|
||||||
|
|
||||||
|
def _format_rich_body(self) -> RenderableType:
|
||||||
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
|
_format_non_tensor_rich_body,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _format_non_tensor_rich_body(self)
|
||||||
|
|
||||||
|
|
||||||
class SummaryRecord(_OutputRecord):
|
class SummaryRecord(_OutputRecord):
|
||||||
type: Literal["summary"] = "summary"
|
type: Literal["summary"] = "summary"
|
||||||
@@ -247,74 +306,38 @@ class SummaryRecord(_OutputRecord):
|
|||||||
return self
|
return self
|
||||||
|
|
||||||
def _format_body(self) -> str:
|
def _format_body(self) -> str:
|
||||||
return (
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
f"Summary: {self.passed} passed, {self.failed} failed, "
|
_format_summary_body,
|
||||||
f"{self.skipped} skipped (total {self.total})"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
return _format_summary_body(self)
|
||||||
|
|
||||||
|
def _format_rich_body(self) -> RenderableType:
|
||||||
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
|
_format_summary_rich_body,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _format_summary_rich_body(self)
|
||||||
|
|
||||||
|
|
||||||
class LogRecord(_OutputRecord):
|
class LogRecord(_OutputRecord):
|
||||||
type: Literal["log"] = "log"
|
type: Literal["log"] = "log"
|
||||||
|
|
||||||
def _format_body(self) -> str:
|
def _format_body(self) -> str:
|
||||||
return ""
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
|
_format_log_body,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _format_log_body(self)
|
||||||
|
|
||||||
|
|
||||||
|
# Re-export _format_aligner_plan for backward compatibility (used by tests)
|
||||||
def _format_aligner_plan(traced_plan: TracedAlignerPlan) -> str:
|
def _format_aligner_plan(traced_plan: TracedAlignerPlan) -> str:
|
||||||
lines: list[str] = ["Aligner Plan:"]
|
from sglang.srt.debug_utils.comparator.output_formatter import (
|
||||||
|
_format_aligner_plan as _impl,
|
||||||
|
)
|
||||||
|
|
||||||
for side_label, traced_side in [
|
return _impl(traced_plan)
|
||||||
("baseline", traced_plan.per_side.x),
|
|
||||||
("target", traced_plan.per_side.y),
|
|
||||||
]:
|
|
||||||
if not traced_side.step_plans:
|
|
||||||
lines.append(f" {side_label}: (no steps)")
|
|
||||||
continue
|
|
||||||
|
|
||||||
step_summaries: list[str] = []
|
|
||||||
for traced_step in traced_side.step_plans:
|
|
||||||
sub_strs: list[str] = [
|
|
||||||
_format_sub_plan_text(traced_sub)
|
|
||||||
for traced_sub in traced_step.sub_plans
|
|
||||||
]
|
|
||||||
summary: str = ", ".join(sub_strs) if sub_strs else "passthrough"
|
|
||||||
step_summaries.append(f"step={traced_step.step}: {summary}")
|
|
||||||
lines.append(f" {side_label}: [{'; '.join(step_summaries)}]")
|
|
||||||
|
|
||||||
lines.extend(_format_cross_side_plan_text(traced_plan.plan))
|
|
||||||
return "\n".join(lines)
|
|
||||||
|
|
||||||
|
|
||||||
def _format_sub_plan_text(traced_sub: TracedSubPlan) -> str:
|
|
||||||
sub_desc: str = f"{traced_sub.plan.type}"
|
|
||||||
|
|
||||||
if traced_sub.snapshot is not None:
|
|
||||||
snap = traced_sub.snapshot
|
|
||||||
in_count: int = len(snap.input_shapes)
|
|
||||||
out_count: int = len(snap.output_shapes)
|
|
||||||
in_shape: str = str(snap.input_shapes[0]) if snap.input_shapes else "?"
|
|
||||||
out_shape: str = str(snap.output_shapes[0]) if snap.output_shapes else "?"
|
|
||||||
sub_desc += f" {in_count}x{in_shape} -> {out_count}x{out_shape}"
|
|
||||||
|
|
||||||
return sub_desc
|
|
||||||
|
|
||||||
|
|
||||||
def _format_cross_side_plan_text(plan: AlignerPlan) -> list[str]:
|
|
||||||
lines: list[str] = []
|
|
||||||
|
|
||||||
if plan.token_aligner_plan is not None:
|
|
||||||
num_tokens: int = len(plan.token_aligner_plan.locators.x.steps)
|
|
||||||
lines.append(f" token_aligner: {num_tokens} tokens aligned")
|
|
||||||
|
|
||||||
if plan.axis_aligner_plan is not None:
|
|
||||||
parts: list[str] = []
|
|
||||||
if plan.axis_aligner_plan.pattern.x:
|
|
||||||
parts.append(f"x: {plan.axis_aligner_plan.pattern.x}")
|
|
||||||
if plan.axis_aligner_plan.pattern.y:
|
|
||||||
parts.append(f"y: {plan.axis_aligner_plan.pattern.y}")
|
|
||||||
lines.append(f" axis_aligner: {', '.join(parts)}")
|
|
||||||
|
|
||||||
return lines
|
|
||||||
|
|
||||||
|
|
||||||
AnyRecord = Annotated[
|
AnyRecord = Annotated[
|
||||||
|
|||||||
@@ -80,7 +80,7 @@ class ReportSink:
|
|||||||
print(record.model_dump_json())
|
print(record.model_dump_json())
|
||||||
else:
|
else:
|
||||||
console: Console = self._get_console()
|
console: Console = self._get_console()
|
||||||
console.print(record.to_rich(verbosity=self._verbosity))
|
console.print(record.to_rich())
|
||||||
console.print() # blank line between records
|
console.print() # blank line between records
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,15 +1,83 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
|
from rich.markup import escape
|
||||||
|
|
||||||
|
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import UnsharderPlan
|
||||||
from sglang.srt.debug_utils.comparator.tensor_comparator.types import (
|
from sglang.srt.debug_utils.comparator.tensor_comparator.types import (
|
||||||
DiffInfo,
|
DiffInfo,
|
||||||
TensorComparisonInfo,
|
TensorComparisonInfo,
|
||||||
|
TensorInfo,
|
||||||
TensorStats,
|
TensorStats,
|
||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.debug_utils.comparator.output_types import ReplicatedCheckResult
|
from sglang.srt.debug_utils.comparator.aligner.entrypoint.traced_types import (
|
||||||
|
TracedAlignerPlan,
|
||||||
|
TracedSubPlan,
|
||||||
|
)
|
||||||
|
from sglang.srt.debug_utils.comparator.aligner.entrypoint.types import AlignerPlan
|
||||||
|
from sglang.srt.debug_utils.comparator.output_types import (
|
||||||
|
BundleSideInfo,
|
||||||
|
ReplicatedCheckResult,
|
||||||
|
ShapeSnapshot,
|
||||||
|
TensorComparisonRecord,
|
||||||
|
)
|
||||||
|
from sglang.srt.debug_utils.comparator.utils import Pair
|
||||||
|
|
||||||
|
|
||||||
|
def _esc_shape(shape: Optional[list[int]]) -> str:
|
||||||
|
return escape(str(shape))
|
||||||
|
|
||||||
|
|
||||||
|
def _strip_torch_prefix(dtype: str) -> str:
|
||||||
|
return dtype.replace("torch.", "")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Number formatting
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _fmt_val(value: float) -> str:
|
||||||
|
return f"{value:.2e}"
|
||||||
|
|
||||||
|
|
||||||
|
def _fmt_diff_colored(diff: float, *, threshold: float = 1e-2) -> str:
|
||||||
|
formatted: str = f"{diff:+.2e}"
|
||||||
|
if abs(diff) >= threshold:
|
||||||
|
return f"[yellow]{formatted}[/]"
|
||||||
|
return f"[dim]{formatted}[/]"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Passed / color / marker helper
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _category_marker(category: str) -> tuple[bool, str, str]:
|
||||||
|
passed: bool = category == "passed"
|
||||||
|
color: str = "green" if passed else "red"
|
||||||
|
marker: str = f"[{color}]✅[/]" if passed else f"[{color}]❌[/]"
|
||||||
|
return passed, color, marker
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Stats formatting helpers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _format_stat_line(stat_name: str, val_b: float, val_t: float, diff: float) -> str:
|
||||||
|
return (
|
||||||
|
f" [blue]{stat_name:10s}[/] {val_b:>10.4f} vs {val_t:>10.4f}"
|
||||||
|
f" Δ {_fmt_diff_colored(diff)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Old text-only formatters (kept for to_text() backward compatibility)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def format_comparison(info: TensorComparisonInfo) -> str:
|
def format_comparison(info: TensorComparisonInfo) -> str:
|
||||||
@@ -129,3 +197,237 @@ def _format_diff(diff: DiffInfo, prefix_text: str = "") -> list[str]:
|
|||||||
lines.append("[abs_diff] " + " ".join(quantile_parts))
|
lines.append("[abs_diff] " + " ".join(quantile_parts))
|
||||||
|
|
||||||
return lines
|
return lines
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# New Rich markup formatters
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def format_comparison_rich(record: TensorComparisonRecord) -> str:
|
||||||
|
return _format_comparison_normal(record=record)
|
||||||
|
|
||||||
|
|
||||||
|
def _format_comparison_normal(*, record: TensorComparisonRecord) -> str:
|
||||||
|
passed, color, marker = _category_marker(record.category)
|
||||||
|
|
||||||
|
baseline: TensorInfo = record.baseline
|
||||||
|
target: TensorInfo = record.target
|
||||||
|
aligned_shape: str = _esc_shape(record.unified_shape)
|
||||||
|
dtype_str: str = _strip_torch_prefix(baseline.dtype)
|
||||||
|
|
||||||
|
lines: list[str] = []
|
||||||
|
|
||||||
|
# L0: Header
|
||||||
|
lines.append(
|
||||||
|
f"{marker} [bold {color}]{escape(record.name)}[/] "
|
||||||
|
f"[dim cyan]── {dtype_str} {aligned_shape}[/]"
|
||||||
|
)
|
||||||
|
|
||||||
|
# L1: Key metrics
|
||||||
|
if record.diff is not None:
|
||||||
|
diff: DiffInfo = record.diff
|
||||||
|
rel_style: str = f"bold {color}" if not passed else color
|
||||||
|
lines.append(
|
||||||
|
f" [{rel_style}]rel_diff={_fmt_val(diff.rel_diff)}[/]"
|
||||||
|
f" max_abs={_fmt_val(diff.max_abs_diff)}"
|
||||||
|
f" mean_abs={_fmt_val(diff.mean_abs_diff)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if not passed:
|
||||||
|
lines.append(
|
||||||
|
f" max_abs @ {_esc_shape(diff.max_diff_coord)}: "
|
||||||
|
f"baseline={diff.baseline_at_max} target={diff.target_at_max}"
|
||||||
|
)
|
||||||
|
elif record.shape_mismatch:
|
||||||
|
lines.append(" [yellow]⚠ Shape mismatch[/]")
|
||||||
|
|
||||||
|
# Downcast info
|
||||||
|
if record.diff_downcast is not None and record.downcast_dtype is not None:
|
||||||
|
dc: DiffInfo = record.diff_downcast
|
||||||
|
dc_marker: str = "[green]✅[/]" if dc.passed else "[red]❌[/]"
|
||||||
|
lines.append(
|
||||||
|
f" {dc_marker} downcast to {record.downcast_dtype}: "
|
||||||
|
f"rel_diff={_fmt_val(dc.rel_diff)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Bundle section
|
||||||
|
if record.raw_bundle_info is not None:
|
||||||
|
lines.append(" [dim]Bundle[/]")
|
||||||
|
lines.extend(_format_bundle_section(bundle_info=record.raw_bundle_info))
|
||||||
|
|
||||||
|
# Plan section
|
||||||
|
if record.traced_plan is not None:
|
||||||
|
lines.append(" [dim]Plan[/]")
|
||||||
|
lines.extend(_format_plan_section_rich(traced_plan=record.traced_plan))
|
||||||
|
|
||||||
|
# Aligned section
|
||||||
|
lines.append(" [dim]Aligned[/]")
|
||||||
|
lines.append(
|
||||||
|
f" {_esc_shape(record.unified_shape)} vs {_esc_shape(target.shape)}"
|
||||||
|
f" {baseline.dtype} vs {target.dtype}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Stats section
|
||||||
|
lines.append(" [dim]Stats[/]")
|
||||||
|
lines.extend(_format_stats_rich(baseline=baseline.stats, target=target.stats))
|
||||||
|
|
||||||
|
# Abs diff percentiles (show when failed)
|
||||||
|
if not passed and record.diff is not None and record.diff.abs_diff_percentiles:
|
||||||
|
lines.append(" [dim]Abs Diff Percentiles[/]")
|
||||||
|
lines.append(" " + _format_abs_diff_percentiles_rich(record.diff))
|
||||||
|
|
||||||
|
# Samples (show when failed)
|
||||||
|
if not passed and baseline.sample is not None:
|
||||||
|
lines.append(" [dim]Samples[/]")
|
||||||
|
lines.append(f" baseline {escape(baseline.sample)}")
|
||||||
|
if target.sample is not None:
|
||||||
|
lines.append(f" target {escape(target.sample)}")
|
||||||
|
|
||||||
|
# Replicated checks (show when failed)
|
||||||
|
if not passed and record.replicated_checks:
|
||||||
|
lines.append(" [dim]Replicated Checks[/]")
|
||||||
|
for check in record.replicated_checks:
|
||||||
|
chk_marker: str = "[green]✅[/]" if check.passed else "[red]❌[/]"
|
||||||
|
if check.diff is not None:
|
||||||
|
lines.append(
|
||||||
|
f" {chk_marker} axis={check.axis} group={check.group_index}"
|
||||||
|
f" idx={check.compared_index} vs {check.baseline_index}"
|
||||||
|
f" rel_diff={_fmt_val(check.diff.rel_diff)}"
|
||||||
|
f" max_abs={_fmt_val(check.diff.max_abs_diff)}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
lines.append(
|
||||||
|
f" {chk_marker} axis={check.axis} group={check.group_index}"
|
||||||
|
f" idx={check.compared_index} vs {check.baseline_index}: n/a"
|
||||||
|
)
|
||||||
|
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
def _format_bundle_section(bundle_info: Pair[BundleSideInfo]) -> list[str]:
|
||||||
|
lines: list[str] = []
|
||||||
|
|
||||||
|
for label, side in [("baseline", bundle_info.x), ("target", bundle_info.y)]:
|
||||||
|
if not side.files:
|
||||||
|
lines.append(f" {label} [dim](no files)[/]")
|
||||||
|
continue
|
||||||
|
|
||||||
|
dtype_desc: str = _strip_torch_prefix(side.files[0].dtype)
|
||||||
|
|
||||||
|
shapes: list[list[int]] = [f.shape for f in side.files]
|
||||||
|
unique_shapes: set[str] = {str(s) for s in shapes}
|
||||||
|
shape_desc: str
|
||||||
|
if len(unique_shapes) == 1:
|
||||||
|
shape_desc = _esc_shape(shapes[0])
|
||||||
|
else:
|
||||||
|
shape_desc = "mixed shapes"
|
||||||
|
|
||||||
|
dims_part: str = f" [dim]dims: {side.dims}[/]" if side.dims else ""
|
||||||
|
lines.append(
|
||||||
|
f" {label} [cyan]{side.num_files} files[/]"
|
||||||
|
f" × {shape_desc} {dtype_desc}{dims_part}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return lines
|
||||||
|
|
||||||
|
|
||||||
|
def _format_plan_section_rich(
|
||||||
|
*,
|
||||||
|
traced_plan: TracedAlignerPlan,
|
||||||
|
) -> list[str]:
|
||||||
|
lines: list[str] = []
|
||||||
|
|
||||||
|
for side_label, traced_side in [
|
||||||
|
("baseline", traced_plan.per_side.x),
|
||||||
|
("target", traced_plan.per_side.y),
|
||||||
|
]:
|
||||||
|
if not traced_side.step_plans:
|
||||||
|
lines.append(f" {side_label} [dim](passthrough)[/]")
|
||||||
|
continue
|
||||||
|
|
||||||
|
parts: list[str] = [
|
||||||
|
_format_sub_plan_rich(traced_sub)
|
||||||
|
for traced_step in traced_side.step_plans
|
||||||
|
for traced_sub in traced_step.sub_plans
|
||||||
|
]
|
||||||
|
lines.append(f" {side_label} " + " → ".join(parts))
|
||||||
|
|
||||||
|
lines.extend(_format_cross_side_plan_rich(traced_plan.plan))
|
||||||
|
return lines
|
||||||
|
|
||||||
|
|
||||||
|
def _format_sub_plan_rich(traced_sub: TracedSubPlan) -> str:
|
||||||
|
sub = traced_sub.plan
|
||||||
|
snapshot: Optional[ShapeSnapshot] = traced_sub.snapshot
|
||||||
|
|
||||||
|
op_name: str = sub.type
|
||||||
|
axis_str: str = ""
|
||||||
|
if isinstance(sub, UnsharderPlan):
|
||||||
|
axis_str = f"({sub.axis})"
|
||||||
|
|
||||||
|
shape_change: str = ""
|
||||||
|
if snapshot:
|
||||||
|
in_count: int = len(snapshot.input_shapes)
|
||||||
|
out_count: int = len(snapshot.output_shapes)
|
||||||
|
in_shape: str = (
|
||||||
|
_esc_shape(snapshot.input_shapes[0]) if snapshot.input_shapes else "?"
|
||||||
|
)
|
||||||
|
out_shape: str = (
|
||||||
|
_esc_shape(snapshot.output_shapes[0]) if snapshot.output_shapes else "?"
|
||||||
|
)
|
||||||
|
shape_change = f" {in_count}×{in_shape} → {out_count}×{out_shape}"
|
||||||
|
|
||||||
|
return f"[magenta]{op_name}{axis_str}[/]{shape_change}"
|
||||||
|
|
||||||
|
|
||||||
|
def _format_cross_side_plan_rich(plan: AlignerPlan) -> list[str]:
|
||||||
|
lines: list[str] = []
|
||||||
|
|
||||||
|
if plan.token_aligner_plan is not None:
|
||||||
|
num_tokens: int = len(plan.token_aligner_plan.locators.x.steps)
|
||||||
|
lines.append(f" token_aligner [dim]{num_tokens} tokens[/]")
|
||||||
|
|
||||||
|
if plan.axis_aligner_plan is not None:
|
||||||
|
parts: list[str] = []
|
||||||
|
if plan.axis_aligner_plan.pattern.x:
|
||||||
|
parts.append(f"x={plan.axis_aligner_plan.pattern.x}")
|
||||||
|
if plan.axis_aligner_plan.pattern.y:
|
||||||
|
parts.append(f"y={plan.axis_aligner_plan.pattern.y}")
|
||||||
|
if parts:
|
||||||
|
lines.append(f" axis_aligner [dim]{', '.join(parts)}[/]")
|
||||||
|
else:
|
||||||
|
lines.append(" axis_aligner [dim](no-op)[/]")
|
||||||
|
|
||||||
|
return lines
|
||||||
|
|
||||||
|
|
||||||
|
def _format_stats_rich(
|
||||||
|
*,
|
||||||
|
baseline: TensorStats,
|
||||||
|
target: TensorStats,
|
||||||
|
) -> list[str]:
|
||||||
|
lines: list[str] = []
|
||||||
|
|
||||||
|
# Compact: mean, std, range (min/max combined)
|
||||||
|
for stat_name in ("mean", "std"):
|
||||||
|
val_b: float = getattr(baseline, stat_name)
|
||||||
|
val_t: float = getattr(target, stat_name)
|
||||||
|
lines.append(_format_stat_line(stat_name, val_b, val_t, val_t - val_b))
|
||||||
|
|
||||||
|
# Range line: combine min/max (escape brackets to avoid Rich markup)
|
||||||
|
range_baseline: str = escape(f"[{baseline.min:.4f}, {baseline.max:.4f}]")
|
||||||
|
range_target: str = escape(f"[{target.min:.4f}, {target.max:.4f}]")
|
||||||
|
lines.append(f" [blue]{'range':10s}[/] {range_baseline} vs {range_target}")
|
||||||
|
|
||||||
|
return lines
|
||||||
|
|
||||||
|
|
||||||
|
def _format_abs_diff_percentiles_rich(diff: DiffInfo) -> str:
|
||||||
|
parts: list[str] = []
|
||||||
|
for p, value in sorted(diff.abs_diff_percentiles.items()):
|
||||||
|
formatted: str = f"p{p}={_fmt_val(value)}"
|
||||||
|
if p >= 99 and value > 0.1:
|
||||||
|
formatted = f"[yellow]{formatted}[/]"
|
||||||
|
parts.append(formatted)
|
||||||
|
return " ".join(parts)
|
||||||
|
|||||||
@@ -19,23 +19,6 @@ from sglang.srt.debug_utils.comparator.report_sink import report_sink
|
|||||||
collect_ignore_glob: list[str] = []
|
collect_ignore_glob: list[str] = []
|
||||||
|
|
||||||
|
|
||||||
def pytest_configure(config: pytest.Config) -> None:
|
|
||||||
config.addinivalue_line(
|
|
||||||
"filterwarnings",
|
|
||||||
"ignore:Unknown config option. asyncio_mode:pytest.PytestConfigWarning",
|
|
||||||
)
|
|
||||||
config.addinivalue_line(
|
|
||||||
"filterwarnings",
|
|
||||||
"ignore:builtin type Swig.*:DeprecationWarning",
|
|
||||||
)
|
|
||||||
config.addinivalue_line(
|
|
||||||
"filterwarnings",
|
|
||||||
"ignore:Named tensors and all their associated APIs:UserWarning",
|
|
||||||
)
|
|
||||||
|
|
||||||
collect_ignore_glob: list[str] = []
|
|
||||||
|
|
||||||
|
|
||||||
def pytest_configure(config: pytest.Config) -> None:
|
def pytest_configure(config: pytest.Config) -> None:
|
||||||
config.addinivalue_line(
|
config.addinivalue_line(
|
||||||
"filterwarnings",
|
"filterwarnings",
|
||||||
|
|||||||
@@ -1,96 +1,63 @@
|
|||||||
import sys
|
import sys
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
from registered.debug_utils.comparator.testing_helpers import make_diff as _make_diff
|
||||||
|
from registered.debug_utils.comparator.testing_helpers import make_stats as _make_stats
|
||||||
|
from registered.debug_utils.comparator.testing_helpers import (
|
||||||
|
make_tensor_info as _make_tensor_info,
|
||||||
|
)
|
||||||
|
|
||||||
|
from sglang.srt.debug_utils.comparator.aligner.axis_aligner import AxisAlignerPlan
|
||||||
|
from sglang.srt.debug_utils.comparator.aligner.entrypoint.traced_types import (
|
||||||
|
TracedAlignerPlan,
|
||||||
|
TracedSidePlan,
|
||||||
|
TracedStepPlan,
|
||||||
|
TracedSubPlan,
|
||||||
|
)
|
||||||
|
from sglang.srt.debug_utils.comparator.aligner.entrypoint.types import (
|
||||||
|
AlignerPerStepPlan,
|
||||||
|
AlignerPlan,
|
||||||
|
)
|
||||||
|
from sglang.srt.debug_utils.comparator.aligner.reorderer.types import (
|
||||||
|
ReordererPlan,
|
||||||
|
ZigzagToNaturalParams,
|
||||||
|
)
|
||||||
|
from sglang.srt.debug_utils.comparator.aligner.token_aligner.smart.types import (
|
||||||
|
TokenAlignerPlan,
|
||||||
|
TokenLocator,
|
||||||
|
)
|
||||||
|
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import (
|
||||||
|
ConcatParams,
|
||||||
|
UnsharderPlan,
|
||||||
|
)
|
||||||
|
from sglang.srt.debug_utils.comparator.dims_spec import ParallelAxis, TokenLayout
|
||||||
|
from sglang.srt.debug_utils.comparator.output_types import (
|
||||||
|
BundleFileInfo,
|
||||||
|
BundleSideInfo,
|
||||||
|
ReplicatedCheckResult,
|
||||||
|
ShapeSnapshot,
|
||||||
|
TensorComparisonRecord,
|
||||||
|
)
|
||||||
from sglang.srt.debug_utils.comparator.tensor_comparator.formatter import (
|
from sglang.srt.debug_utils.comparator.tensor_comparator.formatter import (
|
||||||
|
_format_abs_diff_percentiles_rich,
|
||||||
|
_format_bundle_section,
|
||||||
|
_format_plan_section_rich,
|
||||||
|
_format_stats_rich,
|
||||||
format_comparison,
|
format_comparison,
|
||||||
|
format_comparison_rich,
|
||||||
|
format_replicated_checks,
|
||||||
)
|
)
|
||||||
from sglang.srt.debug_utils.comparator.tensor_comparator.types import (
|
from sglang.srt.debug_utils.comparator.tensor_comparator.types import (
|
||||||
DiffInfo,
|
DiffInfo,
|
||||||
TensorComparisonInfo,
|
TensorComparisonInfo,
|
||||||
TensorInfo,
|
|
||||||
TensorStats,
|
TensorStats,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.debug_utils.comparator.utils import Pair
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
register_cpu_ci(est_time=10, suite="default", nightly=True)
|
register_cpu_ci(est_time=10, suite="default", nightly=True)
|
||||||
|
|
||||||
|
|
||||||
_DEFAULT_PERCENTILES: dict[int, float] = {
|
|
||||||
1: -1.8,
|
|
||||||
5: -1.5,
|
|
||||||
50: 0.0,
|
|
||||||
95: 1.5,
|
|
||||||
99: 1.8,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def _make_stats(
|
|
||||||
mean: float = 0.0,
|
|
||||||
abs_mean: float = 0.8,
|
|
||||||
std: float = 1.0,
|
|
||||||
min: float = -2.0,
|
|
||||||
max: float = 2.0,
|
|
||||||
percentiles: dict[int, float] | None = None,
|
|
||||||
) -> TensorStats:
|
|
||||||
return TensorStats(
|
|
||||||
mean=mean,
|
|
||||||
abs_mean=abs_mean,
|
|
||||||
std=std,
|
|
||||||
min=min,
|
|
||||||
max=max,
|
|
||||||
percentiles=percentiles if percentiles is not None else _DEFAULT_PERCENTILES,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
_DEFAULT_ABS_DIFF_PERCENTILES: dict[int, float] = {
|
|
||||||
1: 0.0001,
|
|
||||||
5: 0.0001,
|
|
||||||
50: 0.0002,
|
|
||||||
95: 0.0004,
|
|
||||||
99: 0.0005,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def _make_diff(
|
|
||||||
rel_diff: float = 0.0001,
|
|
||||||
max_abs_diff: float = 0.0005,
|
|
||||||
mean_abs_diff: float = 0.0002,
|
|
||||||
abs_diff_percentiles: dict[int, float] | None = None,
|
|
||||||
diff_threshold: float = 1e-3,
|
|
||||||
passed: bool = True,
|
|
||||||
) -> DiffInfo:
|
|
||||||
return DiffInfo(
|
|
||||||
rel_diff=rel_diff,
|
|
||||||
max_abs_diff=max_abs_diff,
|
|
||||||
mean_abs_diff=mean_abs_diff,
|
|
||||||
abs_diff_percentiles=(
|
|
||||||
abs_diff_percentiles
|
|
||||||
if abs_diff_percentiles is not None
|
|
||||||
else _DEFAULT_ABS_DIFF_PERCENTILES
|
|
||||||
),
|
|
||||||
max_diff_coord=[2, 3],
|
|
||||||
baseline_at_max=1.0,
|
|
||||||
target_at_max=1.0005,
|
|
||||||
diff_threshold=diff_threshold,
|
|
||||||
passed=passed,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _make_tensor_info(
|
|
||||||
shape: list[int] | None = None,
|
|
||||||
dtype: str = "torch.float32",
|
|
||||||
stats: TensorStats | None = None,
|
|
||||||
sample: str | None = None,
|
|
||||||
) -> TensorInfo:
|
|
||||||
return TensorInfo(
|
|
||||||
shape=shape if shape is not None else [4, 8],
|
|
||||||
dtype=dtype,
|
|
||||||
stats=stats if stats is not None else _make_stats(),
|
|
||||||
sample=sample,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# Snapshot strings below are intentionally spelled out in full per test.
|
# Snapshot strings below are intentionally spelled out in full per test.
|
||||||
# The shared skeleton (stats block, diff block) looks duplicated, but keeping
|
# The shared skeleton (stats block, diff block) looks duplicated, but keeping
|
||||||
# each test self-contained makes failures immediately readable without chasing
|
# each test self-contained makes failures immediately readable without chasing
|
||||||
@@ -294,5 +261,552 @@ class TestFormatComparison:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_comparison_record(
|
||||||
|
name: str = "hidden_states",
|
||||||
|
shape: list[int] | None = None,
|
||||||
|
dtype: str = "torch.float32",
|
||||||
|
diff: DiffInfo | None = None,
|
||||||
|
shape_mismatch: bool = False,
|
||||||
|
sample: str | None = None,
|
||||||
|
diff_downcast: DiffInfo | None = None,
|
||||||
|
downcast_dtype: str | None = None,
|
||||||
|
replicated_checks: list[ReplicatedCheckResult] | None = None,
|
||||||
|
raw_bundle_info: Pair[BundleSideInfo] | None = None,
|
||||||
|
traced_plan: TracedAlignerPlan | None = None,
|
||||||
|
) -> TensorComparisonRecord:
|
||||||
|
s: list[int] = shape if shape is not None else [4, 8]
|
||||||
|
return TensorComparisonRecord(
|
||||||
|
name=name,
|
||||||
|
baseline=_make_tensor_info(shape=s, dtype=dtype, sample=sample),
|
||||||
|
target=_make_tensor_info(shape=s, dtype=dtype, sample=sample),
|
||||||
|
unified_shape=s,
|
||||||
|
shape_mismatch=shape_mismatch,
|
||||||
|
diff=diff,
|
||||||
|
diff_downcast=diff_downcast,
|
||||||
|
downcast_dtype=downcast_dtype,
|
||||||
|
replicated_checks=replicated_checks or [],
|
||||||
|
raw_bundle_info=raw_bundle_info,
|
||||||
|
traced_plan=traced_plan,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_bundle_side_info(
|
||||||
|
num_files: int = 2,
|
||||||
|
shape: list[int] | None = None,
|
||||||
|
dtype: str = "torch.float32",
|
||||||
|
dims: str | None = None,
|
||||||
|
with_parallel_info: bool = False,
|
||||||
|
) -> BundleSideInfo:
|
||||||
|
s: list[int] = shape if shape is not None else [2, 4096]
|
||||||
|
files: list[BundleFileInfo] = []
|
||||||
|
for i in range(num_files):
|
||||||
|
par: dict[str, str] | None = (
|
||||||
|
{"tp": f"{i}/{num_files}"} if with_parallel_info else None
|
||||||
|
)
|
||||||
|
files.append(BundleFileInfo(shape=s, dtype=dtype, rank=i, parallel_info=par))
|
||||||
|
return BundleSideInfo(num_files=num_files, files=files, dims=dims)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_simple_aligner_plan(
|
||||||
|
*,
|
||||||
|
with_unsharder: bool = False,
|
||||||
|
with_reorderer: bool = False,
|
||||||
|
with_token_aligner: bool = False,
|
||||||
|
with_axis_aligner: bool = False,
|
||||||
|
axis_aligner_noop: bool = False,
|
||||||
|
) -> AlignerPlan:
|
||||||
|
baseline_plans: list[AlignerPerStepPlan] = []
|
||||||
|
target_plans: list[AlignerPerStepPlan] = []
|
||||||
|
|
||||||
|
if with_unsharder:
|
||||||
|
unsharder: UnsharderPlan = UnsharderPlan(
|
||||||
|
axis=ParallelAxis.TP,
|
||||||
|
params=ConcatParams(dim_name="h"),
|
||||||
|
groups=[[0, 1]],
|
||||||
|
)
|
||||||
|
target_plans.append(
|
||||||
|
AlignerPerStepPlan(
|
||||||
|
step=0, input_object_indices=[0, 1], sub_plans=[unsharder]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
if with_reorderer:
|
||||||
|
reorderer: ReordererPlan = ReordererPlan(
|
||||||
|
params=ZigzagToNaturalParams(dim_name="s", cp_size=2),
|
||||||
|
)
|
||||||
|
target_plans.append(
|
||||||
|
AlignerPerStepPlan(step=0, input_object_indices=[0], sub_plans=[reorderer])
|
||||||
|
)
|
||||||
|
|
||||||
|
token_aligner_plan: TokenAlignerPlan | None = None
|
||||||
|
if with_token_aligner:
|
||||||
|
token_aligner_plan = TokenAlignerPlan(
|
||||||
|
locators=Pair(
|
||||||
|
x=TokenLocator(steps=[0, 0, 0], token_index_in_step=[0, 1, 2]),
|
||||||
|
y=TokenLocator(steps=[0, 0, 0], token_index_in_step=[0, 1, 2]),
|
||||||
|
),
|
||||||
|
layouts=Pair(x=TokenLayout.T, y=TokenLayout.T),
|
||||||
|
)
|
||||||
|
|
||||||
|
axis_aligner_plan: AxisAlignerPlan | None = None
|
||||||
|
if with_axis_aligner:
|
||||||
|
if axis_aligner_noop:
|
||||||
|
axis_aligner_plan = AxisAlignerPlan(pattern=Pair(x=None, y=None))
|
||||||
|
else:
|
||||||
|
axis_aligner_plan = AxisAlignerPlan(
|
||||||
|
pattern=Pair(x="b s d -> s b d", y=None)
|
||||||
|
)
|
||||||
|
|
||||||
|
return AlignerPlan(
|
||||||
|
per_step_plans=Pair(x=baseline_plans, y=target_plans),
|
||||||
|
token_aligner_plan=token_aligner_plan,
|
||||||
|
axis_aligner_plan=axis_aligner_plan,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_traced_plan(
|
||||||
|
plan: AlignerPlan,
|
||||||
|
*,
|
||||||
|
target_input_shapes: list[list[int]] | None = None,
|
||||||
|
target_output_shapes: list[list[int]] | None = None,
|
||||||
|
) -> TracedAlignerPlan:
|
||||||
|
"""Build a TracedAlignerPlan by attaching snapshots to target sub_plans."""
|
||||||
|
baseline_traced_steps: list[TracedStepPlan] = [
|
||||||
|
TracedStepPlan(
|
||||||
|
step=sp.step,
|
||||||
|
input_object_indices=sp.input_object_indices,
|
||||||
|
sub_plans=[TracedSubPlan(plan=sub) for sub in sp.sub_plans],
|
||||||
|
)
|
||||||
|
for sp in plan.per_step_plans.x
|
||||||
|
]
|
||||||
|
|
||||||
|
target_traced_steps: list[TracedStepPlan] = []
|
||||||
|
for sp in plan.per_step_plans.y:
|
||||||
|
traced_subs: list[TracedSubPlan] = []
|
||||||
|
for sub in sp.sub_plans:
|
||||||
|
snapshot: ShapeSnapshot | None = None
|
||||||
|
if target_input_shapes is not None or target_output_shapes is not None:
|
||||||
|
snapshot = ShapeSnapshot(
|
||||||
|
input_shapes=target_input_shapes or [[2, 4096], [2, 4096]],
|
||||||
|
output_shapes=target_output_shapes or [[4, 4096]],
|
||||||
|
)
|
||||||
|
traced_subs.append(TracedSubPlan(plan=sub, snapshot=snapshot))
|
||||||
|
target_traced_steps.append(
|
||||||
|
TracedStepPlan(
|
||||||
|
step=sp.step,
|
||||||
|
input_object_indices=sp.input_object_indices,
|
||||||
|
sub_plans=traced_subs,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return TracedAlignerPlan(
|
||||||
|
plan=plan,
|
||||||
|
per_side=Pair(
|
||||||
|
x=TracedSidePlan(step_plans=baseline_traced_steps),
|
||||||
|
y=TracedSidePlan(step_plans=target_traced_steps),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Rich format snapshot tests (normal mode only)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class TestFormatComparisonRichNormal:
|
||||||
|
"""format_comparison_rich() snapshot tests."""
|
||||||
|
|
||||||
|
def test_passed(self) -> None:
|
||||||
|
record: TensorComparisonRecord = _make_comparison_record(
|
||||||
|
diff=_make_diff(rel_diff=1e-4, passed=True),
|
||||||
|
)
|
||||||
|
result: str = format_comparison_rich(record)
|
||||||
|
|
||||||
|
assert result == (
|
||||||
|
"[green]✅[/] [bold green]hidden_states[/] [dim cyan]── float32 [4, 8][/]\n"
|
||||||
|
" [green]rel_diff=1.00e-04[/] max_abs=5.00e-04 mean_abs=2.00e-04\n"
|
||||||
|
" [dim]Aligned[/]\n"
|
||||||
|
" [4, 8] vs [4, 8] torch.float32 vs torch.float32\n"
|
||||||
|
" [dim]Stats[/]\n"
|
||||||
|
" [blue]mean [/] 0.0000 vs 0.0000 Δ [dim]+0.00e+00[/]\n"
|
||||||
|
" [blue]std [/] 1.0000 vs 1.0000 Δ [dim]+0.00e+00[/]\n"
|
||||||
|
" [blue]range [/] [-2.0000, 2.0000] vs [-2.0000, 2.0000]"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_failed(self) -> None:
|
||||||
|
record: TensorComparisonRecord = _make_comparison_record(
|
||||||
|
diff=_make_diff(
|
||||||
|
rel_diff=0.5, max_abs_diff=1.0, mean_abs_diff=0.3, passed=False
|
||||||
|
),
|
||||||
|
)
|
||||||
|
result: str = format_comparison_rich(record)
|
||||||
|
|
||||||
|
assert result == (
|
||||||
|
"[red]❌[/] [bold red]hidden_states[/] [dim cyan]── float32 [4, 8][/]\n"
|
||||||
|
" [bold red]rel_diff=5.00e-01[/] max_abs=1.00e+00 mean_abs=3.00e-01\n"
|
||||||
|
" max_abs @ [2, 3]: baseline=1.0 target=1.0005\n"
|
||||||
|
" [dim]Aligned[/]\n"
|
||||||
|
" [4, 8] vs [4, 8] torch.float32 vs torch.float32\n"
|
||||||
|
" [dim]Stats[/]\n"
|
||||||
|
" [blue]mean [/] 0.0000 vs 0.0000 Δ [dim]+0.00e+00[/]\n"
|
||||||
|
" [blue]std [/] 1.0000 vs 1.0000 Δ [dim]+0.00e+00[/]\n"
|
||||||
|
" [blue]range [/] [-2.0000, 2.0000] vs [-2.0000, 2.0000]\n"
|
||||||
|
" [dim]Abs Diff Percentiles[/]\n"
|
||||||
|
" p1=1.00e-04 p5=1.00e-04 p50=2.00e-04 p95=4.00e-04 p99=5.00e-04"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_shape_mismatch(self) -> None:
|
||||||
|
record: TensorComparisonRecord = _make_comparison_record(
|
||||||
|
shape_mismatch=True,
|
||||||
|
)
|
||||||
|
result: str = format_comparison_rich(record)
|
||||||
|
|
||||||
|
assert result == (
|
||||||
|
"[red]❌[/] [bold red]hidden_states[/] [dim cyan]── float32 [4, 8][/]\n"
|
||||||
|
" [yellow]⚠ Shape mismatch[/]\n"
|
||||||
|
" [dim]Aligned[/]\n"
|
||||||
|
" [4, 8] vs [4, 8] torch.float32 vs torch.float32\n"
|
||||||
|
" [dim]Stats[/]\n"
|
||||||
|
" [blue]mean [/] 0.0000 vs 0.0000 Δ [dim]+0.00e+00[/]\n"
|
||||||
|
" [blue]std [/] 1.0000 vs 1.0000 Δ [dim]+0.00e+00[/]\n"
|
||||||
|
" [blue]range [/] [-2.0000, 2.0000] vs [-2.0000, 2.0000]"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_with_downcast(self) -> None:
|
||||||
|
record: TensorComparisonRecord = _make_comparison_record(
|
||||||
|
diff=_make_diff(rel_diff=0.01, passed=False),
|
||||||
|
diff_downcast=_make_diff(rel_diff=1e-5, passed=True),
|
||||||
|
downcast_dtype="torch.bfloat16",
|
||||||
|
)
|
||||||
|
result: str = format_comparison_rich(record)
|
||||||
|
|
||||||
|
assert result == (
|
||||||
|
"[red]❌[/] [bold red]hidden_states[/] [dim cyan]── float32 [4, 8][/]\n"
|
||||||
|
" [bold red]rel_diff=1.00e-02[/] max_abs=5.00e-04 mean_abs=2.00e-04\n"
|
||||||
|
" max_abs @ [2, 3]: baseline=1.0 target=1.0005\n"
|
||||||
|
" [green]✅[/] downcast to torch.bfloat16: rel_diff=1.00e-05\n"
|
||||||
|
" [dim]Aligned[/]\n"
|
||||||
|
" [4, 8] vs [4, 8] torch.float32 vs torch.float32\n"
|
||||||
|
" [dim]Stats[/]\n"
|
||||||
|
" [blue]mean [/] 0.0000 vs 0.0000 Δ [dim]+0.00e+00[/]\n"
|
||||||
|
" [blue]std [/] 1.0000 vs 1.0000 Δ [dim]+0.00e+00[/]\n"
|
||||||
|
" [blue]range [/] [-2.0000, 2.0000] vs [-2.0000, 2.0000]\n"
|
||||||
|
" [dim]Abs Diff Percentiles[/]\n"
|
||||||
|
" p1=1.00e-04 p5=1.00e-04 p50=2.00e-04 p95=4.00e-04 p99=5.00e-04"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_with_bundle_info(self) -> None:
|
||||||
|
bundle_info: Pair[BundleSideInfo] = Pair(
|
||||||
|
x=_make_bundle_side_info(num_files=2, dims="b s h(tp) d"),
|
||||||
|
y=_make_bundle_side_info(num_files=2, dims="b s h(tp) d"),
|
||||||
|
)
|
||||||
|
record: TensorComparisonRecord = _make_comparison_record(
|
||||||
|
diff=_make_diff(passed=True),
|
||||||
|
raw_bundle_info=bundle_info,
|
||||||
|
)
|
||||||
|
result: str = format_comparison_rich(record)
|
||||||
|
|
||||||
|
assert result == (
|
||||||
|
"[green]✅[/] [bold green]hidden_states[/] [dim cyan]── float32 [4, 8][/]\n"
|
||||||
|
" [green]rel_diff=1.00e-04[/] max_abs=5.00e-04 mean_abs=2.00e-04\n"
|
||||||
|
" [dim]Bundle[/]\n"
|
||||||
|
" baseline [cyan]2 files[/] × [2, 4096] float32 [dim]dims: b s h(tp) d[/]\n"
|
||||||
|
" target [cyan]2 files[/] × [2, 4096] float32 [dim]dims: b s h(tp) d[/]\n"
|
||||||
|
" [dim]Aligned[/]\n"
|
||||||
|
" [4, 8] vs [4, 8] torch.float32 vs torch.float32\n"
|
||||||
|
" [dim]Stats[/]\n"
|
||||||
|
" [blue]mean [/] 0.0000 vs 0.0000 Δ [dim]+0.00e+00[/]\n"
|
||||||
|
" [blue]std [/] 1.0000 vs 1.0000 Δ [dim]+0.00e+00[/]\n"
|
||||||
|
" [blue]range [/] [-2.0000, 2.0000] vs [-2.0000, 2.0000]"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_with_plan(self) -> None:
|
||||||
|
plan: AlignerPlan = _make_simple_aligner_plan(with_unsharder=True)
|
||||||
|
record: TensorComparisonRecord = _make_comparison_record(
|
||||||
|
diff=_make_diff(passed=True),
|
||||||
|
traced_plan=_make_traced_plan(plan),
|
||||||
|
)
|
||||||
|
result: str = format_comparison_rich(record)
|
||||||
|
|
||||||
|
assert result == (
|
||||||
|
"[green]✅[/] [bold green]hidden_states[/] [dim cyan]── float32 [4, 8][/]\n"
|
||||||
|
" [green]rel_diff=1.00e-04[/] max_abs=5.00e-04 mean_abs=2.00e-04\n"
|
||||||
|
" [dim]Plan[/]\n"
|
||||||
|
" baseline [dim](passthrough)[/]\n"
|
||||||
|
" target [magenta]unsharder(ParallelAxis.TP)[/]\n"
|
||||||
|
" [dim]Aligned[/]\n"
|
||||||
|
" [4, 8] vs [4, 8] torch.float32 vs torch.float32\n"
|
||||||
|
" [dim]Stats[/]\n"
|
||||||
|
" [blue]mean [/] 0.0000 vs 0.0000 Δ [dim]+0.00e+00[/]\n"
|
||||||
|
" [blue]std [/] 1.0000 vs 1.0000 Δ [dim]+0.00e+00[/]\n"
|
||||||
|
" [blue]range [/] [-2.0000, 2.0000] vs [-2.0000, 2.0000]"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestFormatBundleSection:
|
||||||
|
"""_format_bundle_section() snapshot tests."""
|
||||||
|
|
||||||
|
def test_single_shape(self) -> None:
|
||||||
|
bundle: Pair[BundleSideInfo] = Pair(
|
||||||
|
x=_make_bundle_side_info(num_files=2, shape=[2, 4096]),
|
||||||
|
y=_make_bundle_side_info(num_files=2, shape=[2, 4096]),
|
||||||
|
)
|
||||||
|
lines: list[str] = _format_bundle_section(bundle)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" baseline [cyan]2 files[/] × [2, 4096] float32",
|
||||||
|
" target [cyan]2 files[/] × [2, 4096] float32",
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_mixed_shapes(self) -> None:
|
||||||
|
side: BundleSideInfo = BundleSideInfo(
|
||||||
|
num_files=2,
|
||||||
|
files=[
|
||||||
|
BundleFileInfo(shape=[2, 4096], dtype="torch.float32", rank=0),
|
||||||
|
BundleFileInfo(shape=[3, 4096], dtype="torch.float32", rank=1),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
bundle: Pair[BundleSideInfo] = Pair(x=side, y=side)
|
||||||
|
lines: list[str] = _format_bundle_section(bundle)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" baseline [cyan]2 files[/] × mixed shapes float32",
|
||||||
|
" target [cyan]2 files[/] × mixed shapes float32",
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_no_files(self) -> None:
|
||||||
|
empty: BundleSideInfo = BundleSideInfo(num_files=0, files=[])
|
||||||
|
bundle: Pair[BundleSideInfo] = Pair(x=empty, y=empty)
|
||||||
|
lines: list[str] = _format_bundle_section(bundle)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" baseline [dim](no files)[/]",
|
||||||
|
" target [dim](no files)[/]",
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_with_dims(self) -> None:
|
||||||
|
bundle: Pair[BundleSideInfo] = Pair(
|
||||||
|
x=_make_bundle_side_info(num_files=1, dims="b s h(tp) d"),
|
||||||
|
y=_make_bundle_side_info(num_files=1, dims="b s h(tp) d"),
|
||||||
|
)
|
||||||
|
lines: list[str] = _format_bundle_section(bundle)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" baseline [cyan]1 files[/] × [2, 4096] float32 [dim]dims: b s h(tp) d[/]",
|
||||||
|
" target [cyan]1 files[/] × [2, 4096] float32 [dim]dims: b s h(tp) d[/]",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class TestFormatPlanSectionRich:
|
||||||
|
"""_format_plan_section_rich() snapshot tests."""
|
||||||
|
|
||||||
|
def test_passthrough(self) -> None:
|
||||||
|
plan: AlignerPlan = _make_simple_aligner_plan()
|
||||||
|
traced: TracedAlignerPlan = _make_traced_plan(plan)
|
||||||
|
lines: list[str] = _format_plan_section_rich(traced_plan=traced)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" baseline [dim](passthrough)[/]",
|
||||||
|
" target [dim](passthrough)[/]",
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_unsharder_op(self) -> None:
|
||||||
|
plan: AlignerPlan = _make_simple_aligner_plan(with_unsharder=True)
|
||||||
|
traced: TracedAlignerPlan = _make_traced_plan(plan)
|
||||||
|
lines: list[str] = _format_plan_section_rich(traced_plan=traced)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" baseline [dim](passthrough)[/]",
|
||||||
|
" target [magenta]unsharder(ParallelAxis.TP)[/]",
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_reorderer_op(self) -> None:
|
||||||
|
plan: AlignerPlan = _make_simple_aligner_plan(with_reorderer=True)
|
||||||
|
traced: TracedAlignerPlan = _make_traced_plan(plan)
|
||||||
|
lines: list[str] = _format_plan_section_rich(traced_plan=traced)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" baseline [dim](passthrough)[/]",
|
||||||
|
" target [magenta]reorderer[/]",
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_with_shape_traces(self) -> None:
|
||||||
|
plan: AlignerPlan = _make_simple_aligner_plan(with_unsharder=True)
|
||||||
|
traced: TracedAlignerPlan = _make_traced_plan(
|
||||||
|
plan,
|
||||||
|
target_input_shapes=[[2, 4096], [2, 4096]],
|
||||||
|
target_output_shapes=[[4, 4096]],
|
||||||
|
)
|
||||||
|
lines: list[str] = _format_plan_section_rich(traced_plan=traced)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" baseline [dim](passthrough)[/]",
|
||||||
|
" target [magenta]unsharder(ParallelAxis.TP)[/] 2×[2, 4096] → 1×[4, 4096]",
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_with_token_aligner(self) -> None:
|
||||||
|
plan: AlignerPlan = _make_simple_aligner_plan(with_token_aligner=True)
|
||||||
|
traced: TracedAlignerPlan = _make_traced_plan(plan)
|
||||||
|
lines: list[str] = _format_plan_section_rich(traced_plan=traced)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" baseline [dim](passthrough)[/]",
|
||||||
|
" target [dim](passthrough)[/]",
|
||||||
|
" token_aligner [dim]3 tokens[/]",
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_with_axis_aligner(self) -> None:
|
||||||
|
plan: AlignerPlan = _make_simple_aligner_plan(with_axis_aligner=True)
|
||||||
|
traced: TracedAlignerPlan = _make_traced_plan(plan)
|
||||||
|
lines: list[str] = _format_plan_section_rich(traced_plan=traced)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" baseline [dim](passthrough)[/]",
|
||||||
|
" target [dim](passthrough)[/]",
|
||||||
|
" axis_aligner [dim]x=b s d -> s b d[/]",
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_axis_aligner_noop(self) -> None:
|
||||||
|
plan: AlignerPlan = _make_simple_aligner_plan(
|
||||||
|
with_axis_aligner=True, axis_aligner_noop=True
|
||||||
|
)
|
||||||
|
traced: TracedAlignerPlan = _make_traced_plan(plan)
|
||||||
|
lines: list[str] = _format_plan_section_rich(traced_plan=traced)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" baseline [dim](passthrough)[/]",
|
||||||
|
" target [dim](passthrough)[/]",
|
||||||
|
" axis_aligner [dim](no-op)[/]",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class TestFormatStatsRich:
|
||||||
|
"""_format_stats_rich() snapshot tests."""
|
||||||
|
|
||||||
|
def test_basic(self) -> None:
|
||||||
|
baseline: TensorStats = _make_stats(mean=0.0, std=1.0, min=-2.0, max=2.0)
|
||||||
|
target: TensorStats = _make_stats(
|
||||||
|
mean=0.0001, std=1.0001, min=-2.0001, max=2.0001
|
||||||
|
)
|
||||||
|
lines: list[str] = _format_stats_rich(baseline=baseline, target=target)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" [blue]mean [/] 0.0000 vs 0.0001 Δ [dim]+1.00e-04[/]",
|
||||||
|
" [blue]std [/] 1.0000 vs 1.0001 Δ [dim]+1.00e-04[/]",
|
||||||
|
" [blue]range [/] [-2.0000, 2.0000] vs [-2.0001, 2.0001]",
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_large_delta(self) -> None:
|
||||||
|
baseline: TensorStats = _make_stats(mean=0.0)
|
||||||
|
target: TensorStats = _make_stats(mean=1.0)
|
||||||
|
lines: list[str] = _format_stats_rich(baseline=baseline, target=target)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" [blue]mean [/] 0.0000 vs 1.0000 Δ [yellow]+1.00e+00[/]",
|
||||||
|
" [blue]std [/] 1.0000 vs 1.0000 Δ [dim]+0.00e+00[/]",
|
||||||
|
" [blue]range [/] [-2.0000, 2.0000] vs [-2.0000, 2.0000]",
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_small_delta(self) -> None:
|
||||||
|
baseline: TensorStats = _make_stats(mean=0.0)
|
||||||
|
target: TensorStats = _make_stats(mean=0.001)
|
||||||
|
lines: list[str] = _format_stats_rich(baseline=baseline, target=target)
|
||||||
|
|
||||||
|
assert lines == [
|
||||||
|
" [blue]mean [/] 0.0000 vs 0.0010 Δ [dim]+1.00e-03[/]",
|
||||||
|
" [blue]std [/] 1.0000 vs 1.0000 Δ [dim]+0.00e+00[/]",
|
||||||
|
" [blue]range [/] [-2.0000, 2.0000] vs [-2.0000, 2.0000]",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class TestFormatAbsDiffPercentilesRich:
|
||||||
|
"""_format_abs_diff_percentiles_rich() snapshot tests."""
|
||||||
|
|
||||||
|
def test_normal_values(self) -> None:
|
||||||
|
diff: DiffInfo = _make_diff()
|
||||||
|
result: str = _format_abs_diff_percentiles_rich(diff)
|
||||||
|
|
||||||
|
assert result == (
|
||||||
|
"p1=1.00e-04 p5=1.00e-04 p50=2.00e-04 " "p95=4.00e-04 p99=5.00e-04"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_high_p99_coloring(self) -> None:
|
||||||
|
diff: DiffInfo = _make_diff(
|
||||||
|
abs_diff_percentiles={99: 0.5},
|
||||||
|
)
|
||||||
|
result: str = _format_abs_diff_percentiles_rich(diff)
|
||||||
|
|
||||||
|
assert result == "[yellow]p99=5.00e-01[/]"
|
||||||
|
|
||||||
|
def test_low_p99_no_coloring(self) -> None:
|
||||||
|
diff: DiffInfo = _make_diff(
|
||||||
|
abs_diff_percentiles={99: 0.01},
|
||||||
|
)
|
||||||
|
result: str = _format_abs_diff_percentiles_rich(diff)
|
||||||
|
|
||||||
|
assert result == "p99=1.00e-02"
|
||||||
|
|
||||||
|
|
||||||
|
class TestFormatReplicatedChecks:
|
||||||
|
"""format_replicated_checks() snapshot tests."""
|
||||||
|
|
||||||
|
def test_all_passed(self) -> None:
|
||||||
|
checks: list[ReplicatedCheckResult] = [
|
||||||
|
ReplicatedCheckResult(
|
||||||
|
axis="tp",
|
||||||
|
group_index=0,
|
||||||
|
compared_index=1,
|
||||||
|
baseline_index=0,
|
||||||
|
passed=True,
|
||||||
|
atol=1e-3,
|
||||||
|
diff=_make_diff(rel_diff=1e-6, max_abs_diff=1e-5, mean_abs_diff=1e-6),
|
||||||
|
),
|
||||||
|
]
|
||||||
|
result: str = format_replicated_checks(checks)
|
||||||
|
|
||||||
|
assert result == (
|
||||||
|
"Replicated checks:\n"
|
||||||
|
" ✅ axis=tp group=0 idx=1 vs 0: "
|
||||||
|
"rel_diff=1.000000e-06 max_abs_diff=1.000000e-05 mean_abs_diff=1.000000e-06"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_one_failed(self) -> None:
|
||||||
|
checks: list[ReplicatedCheckResult] = [
|
||||||
|
ReplicatedCheckResult(
|
||||||
|
axis="tp",
|
||||||
|
group_index=0,
|
||||||
|
compared_index=1,
|
||||||
|
baseline_index=0,
|
||||||
|
passed=False,
|
||||||
|
atol=1e-3,
|
||||||
|
diff=_make_diff(rel_diff=0.5, max_abs_diff=1.0, mean_abs_diff=0.3),
|
||||||
|
),
|
||||||
|
]
|
||||||
|
result: str = format_replicated_checks(checks)
|
||||||
|
|
||||||
|
assert result == (
|
||||||
|
"Replicated checks:\n"
|
||||||
|
" ❌ axis=tp group=0 idx=1 vs 0: "
|
||||||
|
"rel_diff=5.000000e-01 max_abs_diff=1.000000e+00 mean_abs_diff=3.000000e-01"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_no_diff(self) -> None:
|
||||||
|
checks: list[ReplicatedCheckResult] = [
|
||||||
|
ReplicatedCheckResult(
|
||||||
|
axis="tp",
|
||||||
|
group_index=0,
|
||||||
|
compared_index=1,
|
||||||
|
baseline_index=0,
|
||||||
|
passed=True,
|
||||||
|
atol=1e-3,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
result: str = format_replicated_checks(checks)
|
||||||
|
|
||||||
|
assert result == (
|
||||||
|
"Replicated checks:\n" " ✅ axis=tp group=0 idx=1 vs 0: n/a diff"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__]))
|
sys.exit(pytest.main([__file__]))
|
||||||
|
|||||||
@@ -1,10 +1,12 @@
|
|||||||
import sys
|
import sys
|
||||||
|
from io import StringIO
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Optional
|
from typing import Any, Optional
|
||||||
|
|
||||||
import polars as pl
|
import polars as pl
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
from rich.console import Console
|
||||||
|
|
||||||
from sglang.srt.debug_utils.comparator.display import (
|
from sglang.srt.debug_utils.comparator.display import (
|
||||||
_collect_input_ids_and_positions,
|
_collect_input_ids_and_positions,
|
||||||
@@ -21,6 +23,12 @@ from sglang.test.ci.ci_register import register_cpu_ci
|
|||||||
register_cpu_ci(est_time=10, suite="default", nightly=True)
|
register_cpu_ci(est_time=10, suite="default", nightly=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _render_rich(renderable: object) -> str:
|
||||||
|
buf: StringIO = StringIO()
|
||||||
|
Console(file=buf, force_terminal=False, width=120).print(renderable)
|
||||||
|
return buf.getvalue().rstrip("\n")
|
||||||
|
|
||||||
|
|
||||||
def _save_dump_file(
|
def _save_dump_file(
|
||||||
directory: Path,
|
directory: Path,
|
||||||
*,
|
*,
|
||||||
@@ -276,6 +284,24 @@ class TestRankInfoRecordSnapshot:
|
|||||||
assert "1/2" in text
|
assert "1/2" in text
|
||||||
assert "0/1" in text
|
assert "0/1" in text
|
||||||
|
|
||||||
|
def test_to_rich_snapshot(self) -> None:
|
||||||
|
from rich.table import Table
|
||||||
|
|
||||||
|
record = RankInfoRecord(
|
||||||
|
label="baseline",
|
||||||
|
rows=[
|
||||||
|
{"rank": 0, "tp": "0/2", "pp": "0/1"},
|
||||||
|
{"rank": 1, "tp": "1/2", "pp": "0/1"},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
body = record._format_rich_body()
|
||||||
|
|
||||||
|
assert isinstance(body, Table)
|
||||||
|
rendered: str = _render_rich(body)
|
||||||
|
assert "baseline ranks" in rendered
|
||||||
|
assert "0/2" in rendered
|
||||||
|
assert "1/2" in rendered
|
||||||
|
|
||||||
def test_json_roundtrip(self) -> None:
|
def test_json_roundtrip(self) -> None:
|
||||||
record = RankInfoRecord(
|
record = RankInfoRecord(
|
||||||
label="target",
|
label="target",
|
||||||
@@ -310,6 +336,29 @@ class TestInputIdsRecordSnapshot:
|
|||||||
assert "10, 20, 30" in text
|
assert "10, 20, 30" in text
|
||||||
assert "0, 1, 2" in text
|
assert "0, 1, 2" in text
|
||||||
|
|
||||||
|
def test_to_rich_snapshot(self) -> None:
|
||||||
|
from rich.table import Table
|
||||||
|
|
||||||
|
record = InputIdsRecord(
|
||||||
|
label="target",
|
||||||
|
rows=[
|
||||||
|
{
|
||||||
|
"step": 0,
|
||||||
|
"rank": 0,
|
||||||
|
"num_tokens": 3,
|
||||||
|
"input_ids": "[10, 20, 30]",
|
||||||
|
"positions": "[0, 1, 2]",
|
||||||
|
},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
body = record._format_rich_body()
|
||||||
|
|
||||||
|
assert isinstance(body, Table)
|
||||||
|
rendered: str = _render_rich(body)
|
||||||
|
assert "target input_ids & positions" in rendered
|
||||||
|
assert "10, 20, 30" in rendered
|
||||||
|
assert "0, 1, 2" in rendered
|
||||||
|
|
||||||
def test_json_roundtrip(self) -> None:
|
def test_json_roundtrip(self) -> None:
|
||||||
record = InputIdsRecord(
|
record = InputIdsRecord(
|
||||||
label="baseline",
|
label="baseline",
|
||||||
|
|||||||
@@ -6,7 +6,8 @@ from registered.debug_utils.comparator.testing_helpers import make_diff as _make
|
|||||||
from registered.debug_utils.comparator.testing_helpers import (
|
from registered.debug_utils.comparator.testing_helpers import (
|
||||||
make_tensor_info as _make_tensor_info,
|
make_tensor_info as _make_tensor_info,
|
||||||
)
|
)
|
||||||
from rich.console import Console
|
from rich.console import Console, Group
|
||||||
|
from rich.panel import Panel
|
||||||
|
|
||||||
from sglang.srt.debug_utils.comparator.aligner.axis_aligner import AxisAlignerPlan
|
from sglang.srt.debug_utils.comparator.aligner.axis_aligner import AxisAlignerPlan
|
||||||
from sglang.srt.debug_utils.comparator.aligner.entrypoint.traced_types import (
|
from sglang.srt.debug_utils.comparator.aligner.entrypoint.traced_types import (
|
||||||
@@ -121,6 +122,23 @@ class TestConfigRecord:
|
|||||||
record: ConfigRecord = ConfigRecord(config={"a": 1, "b": "two"})
|
record: ConfigRecord = ConfigRecord(config={"a": 1, "b": "two"})
|
||||||
assert record._format_body() == "Config: {'a': 1, 'b': 'two'}"
|
assert record._format_body() == "Config: {'a': 1, 'b': 'two'}"
|
||||||
|
|
||||||
|
def test_format_rich_body(self) -> None:
|
||||||
|
record: ConfigRecord = ConfigRecord(config={"threshold": 0.001, "mode": "fast"})
|
||||||
|
body = record._format_rich_body()
|
||||||
|
|
||||||
|
assert isinstance(body, Panel)
|
||||||
|
rendered: str = _render_rich(body)
|
||||||
|
assert rendered == (
|
||||||
|
"╭───────────────────────────────────────────────── Comparator Config "
|
||||||
|
"──────────────────────────────────────────────────╮\n"
|
||||||
|
"│ threshold : 0.001"
|
||||||
|
" │\n"
|
||||||
|
"│ mode : fast"
|
||||||
|
" │\n"
|
||||||
|
"╰──────────────────────────────────────────────────────────────────────"
|
||||||
|
"────────────────────────────────────────────────╯"
|
||||||
|
)
|
||||||
|
|
||||||
def test_to_text_with_errors(self) -> None:
|
def test_to_text_with_errors(self) -> None:
|
||||||
record: ConfigRecord = ConfigRecord(
|
record: ConfigRecord = ConfigRecord(
|
||||||
config={"x": 1},
|
config={"x": 1},
|
||||||
@@ -152,6 +170,14 @@ class TestSkipComparisonRecord:
|
|||||||
)
|
)
|
||||||
assert record._format_body() == "Skip: layer.weight (step=3) (scalar)"
|
assert record._format_body() == "Skip: layer.weight (step=3) (scalar)"
|
||||||
|
|
||||||
|
def test_format_rich_body(self) -> None:
|
||||||
|
record: SkipComparisonRecord = SkipComparisonRecord(
|
||||||
|
name="attn.qkv",
|
||||||
|
reason="no baseline",
|
||||||
|
)
|
||||||
|
body: str = record._format_rich_body()
|
||||||
|
assert body == "[dim]⊘ attn.qkv ── skipped (no baseline)[/]"
|
||||||
|
|
||||||
def test_category_skipped(self) -> None:
|
def test_category_skipped(self) -> None:
|
||||||
record: SkipComparisonRecord = SkipComparisonRecord(
|
record: SkipComparisonRecord = SkipComparisonRecord(
|
||||||
name="x",
|
name="x",
|
||||||
@@ -200,6 +226,32 @@ class TestNonTensorComparisonRecord:
|
|||||||
" target = 0.01 (float)"
|
" target = 0.01 (float)"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_format_rich_body_equal(self) -> None:
|
||||||
|
record: NonTensorComparisonRecord = NonTensorComparisonRecord(
|
||||||
|
name="config.lr",
|
||||||
|
baseline_value="0.001",
|
||||||
|
target_value="0.001",
|
||||||
|
baseline_type="float",
|
||||||
|
target_type="float",
|
||||||
|
values_equal=True,
|
||||||
|
)
|
||||||
|
assert record._format_rich_body() == ("═ config.lr = 0.001 (float) [green]✓[/]")
|
||||||
|
|
||||||
|
def test_format_rich_body_not_equal(self) -> None:
|
||||||
|
record: NonTensorComparisonRecord = NonTensorComparisonRecord(
|
||||||
|
name="config.lr",
|
||||||
|
baseline_value="0.001",
|
||||||
|
target_value="0.01",
|
||||||
|
baseline_type="float",
|
||||||
|
target_type="float",
|
||||||
|
values_equal=False,
|
||||||
|
)
|
||||||
|
assert record._format_rich_body() == (
|
||||||
|
"═ [bold red]config.lr[/]\n"
|
||||||
|
" baseline = 0.001 (float)\n"
|
||||||
|
" target = 0.01 (float)"
|
||||||
|
)
|
||||||
|
|
||||||
def test_with_step(self) -> None:
|
def test_with_step(self) -> None:
|
||||||
record: NonTensorComparisonRecord = NonTensorComparisonRecord(
|
record: NonTensorComparisonRecord = NonTensorComparisonRecord(
|
||||||
name="bias",
|
name="bias",
|
||||||
@@ -250,6 +302,26 @@ class TestSummaryRecord:
|
|||||||
"Summary: 7 passed, 2 failed, 1 skipped (total 10)"
|
"Summary: 7 passed, 2 failed, 1 skipped (total 10)"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_format_rich_body(self) -> None:
|
||||||
|
record: SummaryRecord = SummaryRecord(
|
||||||
|
total=10,
|
||||||
|
passed=7,
|
||||||
|
failed=2,
|
||||||
|
skipped=1,
|
||||||
|
)
|
||||||
|
body = record._format_rich_body()
|
||||||
|
assert isinstance(body, Panel)
|
||||||
|
|
||||||
|
rendered: str = _render_rich(body)
|
||||||
|
assert rendered == (
|
||||||
|
"╭────────────────────────────────────────────────────── SUMMARY "
|
||||||
|
"───────────────────────────────────────────────────────╮\n"
|
||||||
|
"│ 7 passed │ 2 failed │ 1 skipped │ 10 total"
|
||||||
|
" │\n"
|
||||||
|
"╰──────────────────────────────────────────────────────────────────────"
|
||||||
|
"────────────────────────────────────────────────╯"
|
||||||
|
)
|
||||||
|
|
||||||
def test_validation_error(self) -> None:
|
def test_validation_error(self) -> None:
|
||||||
with pytest.raises(ValueError, match="total=5 !="):
|
with pytest.raises(ValueError, match="total=5 !="):
|
||||||
SummaryRecord(total=5, passed=1, failed=1, skipped=1)
|
SummaryRecord(total=5, passed=1, failed=1, skipped=1)
|
||||||
@@ -567,7 +639,7 @@ class TestFormatAlignerPlan:
|
|||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# _OutputRecord log attachment (to_text)
|
# _OutputRecord log attachment (to_text / to_rich)
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
@@ -606,6 +678,27 @@ class TestOutputRecordLogAttachment:
|
|||||||
|
|
||||||
assert text == "Config: {'a': 1}\n ✗ err1\n ℹ note1"
|
assert text == "Config: {'a': 1}\n ✗ err1\n ℹ note1"
|
||||||
|
|
||||||
|
def test_to_rich_string_body(self) -> None:
|
||||||
|
record: SkipComparisonRecord = SkipComparisonRecord(
|
||||||
|
name="x",
|
||||||
|
reason="r",
|
||||||
|
errors=[ErrorLog(category="e", message="oops")],
|
||||||
|
)
|
||||||
|
body = record.to_rich()
|
||||||
|
|
||||||
|
assert isinstance(body, str)
|
||||||
|
assert body == "[dim]⊘ x ── skipped (r)[/]\n [red]✗ oops[/]"
|
||||||
|
|
||||||
|
def test_to_rich_group_body(self) -> None:
|
||||||
|
record: ConfigRecord = ConfigRecord(
|
||||||
|
config={"a": 1},
|
||||||
|
errors=[ErrorLog(category="e", message="oops")],
|
||||||
|
)
|
||||||
|
body = record.to_rich()
|
||||||
|
|
||||||
|
# Panel body + log block → Group
|
||||||
|
assert isinstance(body, Group)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__]))
|
sys.exit(pytest.main([__file__]))
|
||||||
|
|||||||
Reference in New Issue
Block a user