[diffusion] CI: improve readability and fix bug of early-return (#22507)
This commit is contained in:
@@ -366,6 +366,20 @@ def _resolve_quant_config_from_transformer_override(
|
|||||||
transformer_weights_path: str,
|
transformer_weights_path: str,
|
||||||
) -> Optional[QuantizationConfig]:
|
) -> Optional[QuantizationConfig]:
|
||||||
"""Resolve quant config from an override transformer repo or directory."""
|
"""Resolve quant config from an override transformer repo or directory."""
|
||||||
|
expanded_path = os.path.expanduser(transformer_weights_path)
|
||||||
|
if os.path.isfile(expanded_path):
|
||||||
|
return None
|
||||||
|
|
||||||
|
# A single local safetensors file does not carry a directory-level config.json.
|
||||||
|
# Let downstream metadata probing handle it instead of misrouting it through HF.
|
||||||
|
if expanded_path.endswith(".safetensors") and (
|
||||||
|
os.path.isabs(expanded_path)
|
||||||
|
or expanded_path.startswith(".")
|
||||||
|
or os.sep in expanded_path
|
||||||
|
or (os.path.altsep and os.path.altsep in expanded_path)
|
||||||
|
):
|
||||||
|
return None
|
||||||
|
|
||||||
override_quantized_path = maybe_download_model(transformer_weights_path)
|
override_quantized_path = maybe_download_model(transformer_weights_path)
|
||||||
if not os.path.isdir(override_quantized_path):
|
if not os.path.isdir(override_quantized_path):
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -36,11 +36,14 @@ def run_command(command) -> Optional[float]:
|
|||||||
command,
|
command,
|
||||||
stdout=subprocess.PIPE,
|
stdout=subprocess.PIPE,
|
||||||
stderr=subprocess.STDOUT,
|
stderr=subprocess.STDOUT,
|
||||||
text=True,
|
bufsize=0,
|
||||||
encoding="utf-8",
|
|
||||||
) as process:
|
) as process:
|
||||||
for line in process.stdout:
|
while True:
|
||||||
sys.stdout.write(line)
|
chunk = process.stdout.read(4096)
|
||||||
|
if not chunk:
|
||||||
|
break
|
||||||
|
sys.stdout.buffer.write(chunk)
|
||||||
|
sys.stdout.buffer.flush()
|
||||||
process.wait()
|
process.wait()
|
||||||
if process.returncode == 0:
|
if process.returncode == 0:
|
||||||
return True
|
return True
|
||||||
@@ -52,6 +55,7 @@ class CLIBase(unittest.TestCase):
|
|||||||
model_path: str = None
|
model_path: str = None
|
||||||
extra_args = []
|
extra_args = []
|
||||||
data_type: DataType = None
|
data_type: DataType = None
|
||||||
|
log_level: str = "info"
|
||||||
# tested on h100
|
# tested on h100
|
||||||
|
|
||||||
width: int = 720
|
width: int = 720
|
||||||
@@ -83,7 +87,7 @@ class CLIBase(unittest.TestCase):
|
|||||||
"--prompt",
|
"--prompt",
|
||||||
"A curious raccoon",
|
"A curious raccoon",
|
||||||
"--save-output",
|
"--save-output",
|
||||||
"--log-level=debug",
|
f"--log-level={self.log_level}",
|
||||||
f"--width={self.width}",
|
f"--width={self.width}",
|
||||||
f"--height={self.height}",
|
f"--height={self.height}",
|
||||||
f"--output-path={self.output_path}",
|
f"--output-path={self.output_path}",
|
||||||
|
|||||||
@@ -142,7 +142,8 @@ def collect_test_items(files, filter_expr=None):
|
|||||||
cmd.extend(["-k", filter_expr])
|
cmd.extend(["-k", filter_expr])
|
||||||
cmd.extend(files)
|
cmd.extend(files)
|
||||||
|
|
||||||
print(f"Collecting tests with command: {' '.join(cmd)}")
|
filter_note = f" with filter: {filter_expr}" if filter_expr else ""
|
||||||
|
print(f"Collecting tests from {len(files)} file(s){filter_note}")
|
||||||
result = subprocess.run(cmd, capture_output=True, text=True)
|
result = subprocess.run(cmd, capture_output=True, text=True)
|
||||||
|
|
||||||
# Check for collection errors
|
# Check for collection errors
|
||||||
@@ -186,91 +187,206 @@ def collect_test_items(files, filter_expr=None):
|
|||||||
return test_items
|
return test_items
|
||||||
|
|
||||||
|
|
||||||
def run_pytest(files, filter_expr=None, exitfirst=False):
|
def _run_pytest_attempt(cmd: list[str]) -> tuple[int, str]:
|
||||||
|
process = subprocess.Popen(
|
||||||
|
cmd,
|
||||||
|
stdout=subprocess.PIPE,
|
||||||
|
stderr=subprocess.STDOUT,
|
||||||
|
bufsize=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
output_bytes = bytearray()
|
||||||
|
while True:
|
||||||
|
chunk = process.stdout.read(4096)
|
||||||
|
if not chunk:
|
||||||
|
break
|
||||||
|
sys.stdout.buffer.write(chunk)
|
||||||
|
sys.stdout.buffer.flush()
|
||||||
|
output_bytes.extend(chunk)
|
||||||
|
|
||||||
|
process.wait()
|
||||||
|
return process.returncode, output_bytes.decode("utf-8", errors="replace")
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_collection_line(full_output: str) -> str | None:
|
||||||
|
for line in full_output.splitlines():
|
||||||
|
stripped = line.strip()
|
||||||
|
if stripped.startswith("collected "):
|
||||||
|
return stripped
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_short_test_summary(full_output: str) -> list[str]:
|
||||||
|
summary_lines = []
|
||||||
|
in_summary = False
|
||||||
|
for line in full_output.splitlines():
|
||||||
|
stripped = line.strip()
|
||||||
|
if "short test summary info" in stripped:
|
||||||
|
in_summary = True
|
||||||
|
continue
|
||||||
|
if not in_summary:
|
||||||
|
continue
|
||||||
|
if stripped.startswith("="):
|
||||||
|
break
|
||||||
|
if not stripped or stripped.startswith("!"):
|
||||||
|
continue
|
||||||
|
summary_lines.append(stripped)
|
||||||
|
return summary_lines
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_failure_tail(full_output: str, max_lines: int = 20) -> list[str]:
|
||||||
|
summary_lines = _extract_short_test_summary(full_output)
|
||||||
|
if summary_lines:
|
||||||
|
return summary_lines
|
||||||
|
|
||||||
|
lines = [line.rstrip() for line in full_output.splitlines() if line.strip()]
|
||||||
|
return lines[-max_lines:]
|
||||||
|
|
||||||
|
|
||||||
|
def _is_retryable_failure(full_output: str) -> bool:
|
||||||
|
is_perf_assertion = (
|
||||||
|
"multimodal_gen/test/server/test_server_utils.py" in full_output
|
||||||
|
and "AssertionError" in full_output
|
||||||
|
)
|
||||||
|
|
||||||
|
is_flaky_ci_assertion = (
|
||||||
|
"SafetensorError" in full_output
|
||||||
|
or "FileNotFoundError" in full_output
|
||||||
|
or "TimeoutError" in full_output
|
||||||
|
)
|
||||||
|
|
||||||
|
is_oom_error = (
|
||||||
|
"out of memory" in full_output.lower() or "oom killer" in full_output.lower()
|
||||||
|
)
|
||||||
|
|
||||||
|
return is_perf_assertion or is_flaky_ci_assertion or is_oom_error
|
||||||
|
|
||||||
|
|
||||||
|
def _print_attempt_tail_summary(
|
||||||
|
attempt_reports: list[dict], assigned_count: int
|
||||||
|
) -> None:
|
||||||
|
if len(attempt_reports) == 1 and attempt_reports[0]["returncode"] in (0, 5):
|
||||||
|
return
|
||||||
|
|
||||||
|
rows = []
|
||||||
|
for report in attempt_reports:
|
||||||
|
if report["returncode"] in (0, 5):
|
||||||
|
result = "success"
|
||||||
|
elif report["retryable"]:
|
||||||
|
result = "retryable failure"
|
||||||
|
else:
|
||||||
|
result = "failure"
|
||||||
|
rows.append(
|
||||||
|
[
|
||||||
|
report["attempt"],
|
||||||
|
report["mode"],
|
||||||
|
result,
|
||||||
|
report["collection_line"] or "-",
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
print("\n" + "=" * 32 + " Pytest Tail Summary " + "=" * 32, flush=True)
|
||||||
|
print(f"Assigned {assigned_count} test item(s)", flush=True)
|
||||||
|
print(
|
||||||
|
tabulate.tabulate(
|
||||||
|
rows,
|
||||||
|
headers=["Attempt", "Mode", "Result", "Collection"],
|
||||||
|
tablefmt="psql",
|
||||||
|
),
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
for report in attempt_reports:
|
||||||
|
if not report["failure_tail"]:
|
||||||
|
continue
|
||||||
|
print(f"\nAttempt {report['attempt']} failure summary:", flush=True)
|
||||||
|
for line in report["failure_tail"]:
|
||||||
|
print(f" {line}", flush=True)
|
||||||
|
|
||||||
|
print("=" * 84, flush=True)
|
||||||
|
|
||||||
|
|
||||||
|
def run_pytest(files, filter_expr=None):
|
||||||
if not files:
|
if not files:
|
||||||
print("No files to run.")
|
print("No files to run.")
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
base_cmd = [sys.executable, "-m", "pytest", "-s", "-v"]
|
base_cmd = [
|
||||||
if exitfirst:
|
sys.executable,
|
||||||
base_cmd.append("-x")
|
"-m",
|
||||||
|
"pytest",
|
||||||
|
"-s",
|
||||||
|
"-v",
|
||||||
|
"--tb=short",
|
||||||
|
"--no-header",
|
||||||
|
]
|
||||||
|
|
||||||
# Add pytest -k filter if provided
|
|
||||||
if filter_expr:
|
if filter_expr:
|
||||||
base_cmd.extend(["-k", filter_expr])
|
base_cmd.extend(["-k", filter_expr])
|
||||||
|
|
||||||
max_retries = 6
|
max_retries = 6
|
||||||
# retry if the perf assertion failed, for {max_retries} times
|
attempt_reports = []
|
||||||
|
|
||||||
for i in range(max_retries + 1):
|
for i in range(max_retries + 1):
|
||||||
|
is_retry = i > 0
|
||||||
cmd = list(base_cmd)
|
cmd = list(base_cmd)
|
||||||
if i > 0:
|
if is_retry:
|
||||||
cmd.append("--last-failed")
|
cmd.append("--last-failed")
|
||||||
# Always include files to constrain test discovery scope
|
|
||||||
# This prevents pytest from scanning the entire rootdir and
|
|
||||||
# discovering unrelated tests that may have missing dependencies
|
|
||||||
cmd.extend(files)
|
cmd.extend(files)
|
||||||
|
|
||||||
if i > 0:
|
mode = "retry failed items" if is_retry else "initial pass"
|
||||||
print(
|
print(
|
||||||
f"Performance assertion failed. Retrying ({i}/{max_retries}) with --last-failed..."
|
f"Starting pytest attempt {i + 1}/{max_retries + 1}: {mode} "
|
||||||
)
|
f"for {len(files)} assigned item(s)"
|
||||||
|
|
||||||
print(f"Running command: {' '.join(cmd)}")
|
|
||||||
|
|
||||||
process = subprocess.Popen(
|
|
||||||
cmd,
|
|
||||||
stdout=subprocess.PIPE,
|
|
||||||
stderr=subprocess.STDOUT,
|
|
||||||
bufsize=0,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
output_bytes = bytearray()
|
returncode, full_output = _run_pytest_attempt(cmd)
|
||||||
while True:
|
retryable = returncode not in (0, 5) and _is_retryable_failure(full_output)
|
||||||
chunk = process.stdout.read(4096)
|
attempt_reports.append(
|
||||||
if not chunk:
|
{
|
||||||
break
|
"attempt": i + 1,
|
||||||
sys.stdout.buffer.write(chunk)
|
"mode": mode,
|
||||||
sys.stdout.buffer.flush()
|
"returncode": returncode,
|
||||||
output_bytes.extend(chunk)
|
"retryable": retryable,
|
||||||
|
"collection_line": _extract_collection_line(full_output),
|
||||||
process.wait()
|
"failure_tail": (
|
||||||
returncode = process.returncode
|
_extract_failure_tail(full_output)
|
||||||
|
if returncode not in (0, 5)
|
||||||
|
else []
|
||||||
|
),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
if returncode == 0:
|
if returncode == 0:
|
||||||
|
if is_retry:
|
||||||
|
print(f"Recovered retryable failures on attempt {i + 1}.")
|
||||||
|
_print_attempt_tail_summary(attempt_reports, len(files))
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
# Exit code 5 means no tests were collected/selected - treat as success
|
|
||||||
# when using filters, since some partitions may have all tests filtered out
|
|
||||||
if returncode == 5:
|
if returncode == 5:
|
||||||
print(
|
print(
|
||||||
"No tests collected (exit code 5). This is expected when filters "
|
"No tests collected (exit code 5). This is expected when filters "
|
||||||
"deselect all tests in a partition. Treating as success."
|
"deselect all tests in a partition. Treating as success."
|
||||||
)
|
)
|
||||||
|
_print_attempt_tail_summary(attempt_reports, len(files))
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
# check if the failure is due to an assertion in test_server_utils.py
|
if not retryable:
|
||||||
full_output = output_bytes.decode("utf-8", errors="replace")
|
_print_attempt_tail_summary(attempt_reports, len(files))
|
||||||
is_perf_assertion = (
|
|
||||||
"multimodal_gen/test/server/test_server_utils.py" in full_output
|
|
||||||
and "AssertionError" in full_output
|
|
||||||
)
|
|
||||||
|
|
||||||
is_flaky_ci_assertion = (
|
|
||||||
"SafetensorError" in full_output
|
|
||||||
or "FileNotFoundError" in full_output
|
|
||||||
or "TimeoutError" in full_output
|
|
||||||
)
|
|
||||||
|
|
||||||
is_oom_error = (
|
|
||||||
"out of memory" in full_output.lower()
|
|
||||||
or "oom killer" in full_output.lower()
|
|
||||||
)
|
|
||||||
|
|
||||||
if not (is_perf_assertion or is_flaky_ci_assertion or is_oom_error):
|
|
||||||
return returncode
|
return returncode
|
||||||
|
|
||||||
print(f"Max retry exceeded")
|
if i == max_retries:
|
||||||
return returncode
|
print(f"Max retry exceeded ({max_retries})")
|
||||||
|
_print_attempt_tail_summary(attempt_reports, len(files))
|
||||||
|
return returncode
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"Retryable failure detected on attempt {i + 1}. "
|
||||||
|
"Retrying only previously failed items."
|
||||||
|
)
|
||||||
|
|
||||||
|
_print_attempt_tail_summary(attempt_reports, len(files))
|
||||||
|
return attempt_reports[-1]["returncode"]
|
||||||
|
|
||||||
|
|
||||||
def partition_test_files(files, partition_id, total_partitions):
|
def partition_test_files(files, partition_id, total_partitions):
|
||||||
@@ -434,37 +550,19 @@ def main():
|
|||||||
headers = ["Suite", "Partition"]
|
headers = ["Suite", "Partition"]
|
||||||
rows = [[args.suite, partition_info]]
|
rows = [[args.suite, partition_info]]
|
||||||
msg = tabulate.tabulate(rows, headers=headers, tablefmt="psql") + "\n"
|
msg = tabulate.tabulate(rows, headers=headers, tablefmt="psql") + "\n"
|
||||||
msg += f"✅ Enabled {len(my_items)} test(s):\n"
|
msg += f"✅ Assigned {len(my_items)} test(s) from {len(suite_files_abs)} file(s):\n"
|
||||||
for item in my_items:
|
|
||||||
msg += f" - {item}\n"
|
|
||||||
print(msg, flush=True)
|
|
||||||
print(
|
|
||||||
f"Suite: {args.suite} | Partition: {args.partition_id}/{args.total_partitions}"
|
|
||||||
)
|
|
||||||
print(f"Selected {len(suite_files_abs)} files:")
|
|
||||||
for f in suite_files_abs:
|
for f in suite_files_abs:
|
||||||
print(f" - {os.path.basename(f)}")
|
msg += f" - {os.path.basename(f)}\n"
|
||||||
|
print(msg, flush=True)
|
||||||
|
|
||||||
if not my_items:
|
if not my_items:
|
||||||
print("No items assigned to this partition. Exiting success.")
|
print("No items assigned to this partition. Exiting success.")
|
||||||
sys.exit(0)
|
sys.exit(0)
|
||||||
|
|
||||||
print(f"Running {len(my_items)} items in this shard: {', '.join(my_items)}")
|
print(f"Running shard with {len(my_items)} assigned test item(s)")
|
||||||
|
|
||||||
# 4. execute with the specific test items
|
# 4. execute with the specific test items
|
||||||
# Fast-fail: stop on first failure unless --continue-on-error is set
|
exit_code = run_pytest(my_items)
|
||||||
exit_code = run_pytest(
|
|
||||||
my_items,
|
|
||||||
filter_expr=args.filter,
|
|
||||||
exitfirst=not args.continue_on_error,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Print tests again at the end for visibility
|
|
||||||
msg = "\n" + tabulate.tabulate(rows, headers=headers, tablefmt="psql") + "\n"
|
|
||||||
msg += f"✅ Executed {len(my_items)} test(s):\n"
|
|
||||||
for item in my_items:
|
|
||||||
msg += f" - {item}\n"
|
|
||||||
print(msg, flush=True)
|
|
||||||
|
|
||||||
sys.exit(exit_code)
|
sys.exit(exit_code)
|
||||||
|
|
||||||
|
|||||||
@@ -120,6 +120,9 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
"sglang.multimodal_gen.runtime.loader.transformer_load_utils.build_nvfp4_config_from_safetensors_list",
|
"sglang.multimodal_gen.runtime.loader.transformer_load_utils.build_nvfp4_config_from_safetensors_list",
|
||||||
return_value=None,
|
return_value=None,
|
||||||
)
|
)
|
||||||
|
@patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.transformer_load_utils.maybe_download_model"
|
||||||
|
)
|
||||||
@patch(
|
@patch(
|
||||||
"sglang.multimodal_gen.runtime.loader.transformer_load_utils.get_quant_config_from_safetensors_metadata",
|
"sglang.multimodal_gen.runtime.loader.transformer_load_utils.get_quant_config_from_safetensors_metadata",
|
||||||
return_value=None,
|
return_value=None,
|
||||||
@@ -136,8 +139,12 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
_mock_download,
|
_mock_download,
|
||||||
mock_metadata,
|
mock_metadata,
|
||||||
_mock_quant_metadata,
|
_mock_quant_metadata,
|
||||||
|
mock_maybe_download,
|
||||||
_mock_nvfp4,
|
_mock_nvfp4,
|
||||||
):
|
):
|
||||||
|
mock_maybe_download.side_effect = AssertionError(
|
||||||
|
"local safetensors path should not trigger maybe_download_model"
|
||||||
|
)
|
||||||
mock_metadata.return_value = {
|
mock_metadata.return_value = {
|
||||||
"config": json.dumps({"_class_name": _FakeFluxTransformer.__name__})
|
"config": json.dumps({"_class_name": _FakeFluxTransformer.__name__})
|
||||||
}
|
}
|
||||||
@@ -163,6 +170,7 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
self.assertIsNone(spec.param_dtype)
|
self.assertIsNone(spec.param_dtype)
|
||||||
self.assertEqual(len(spec.post_load_hooks), 1)
|
self.assertEqual(len(spec.post_load_hooks), 1)
|
||||||
self.assertIs(nunchaku_config.model_cls, _FakeFluxTransformer)
|
self.assertIs(nunchaku_config.model_cls, _FakeFluxTransformer)
|
||||||
|
mock_maybe_download.assert_not_called()
|
||||||
|
|
||||||
def test_flux2_mixed_nvfp4_fallback_disables_conflicting_offloads(self):
|
def test_flux2_mixed_nvfp4_fallback_disables_conflicting_offloads(self):
|
||||||
server_args = self._make_server_args(
|
server_args = self._make_server_args(
|
||||||
|
|||||||
Reference in New Issue
Block a user