diff --git a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py index 7559f1fda..5ecffdfd0 100644 --- a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py @@ -366,6 +366,20 @@ def _resolve_quant_config_from_transformer_override( transformer_weights_path: str, ) -> Optional[QuantizationConfig]: """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) if not os.path.isdir(override_quantized_path): return None diff --git a/python/sglang/multimodal_gen/test/cli/test_generate_common.py b/python/sglang/multimodal_gen/test/cli/test_generate_common.py index 94d97ff3f..196b82e03 100644 --- a/python/sglang/multimodal_gen/test/cli/test_generate_common.py +++ b/python/sglang/multimodal_gen/test/cli/test_generate_common.py @@ -36,11 +36,14 @@ def run_command(command) -> Optional[float]: command, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, - text=True, - encoding="utf-8", + bufsize=0, ) as process: - for line in process.stdout: - sys.stdout.write(line) + while True: + chunk = process.stdout.read(4096) + if not chunk: + break + sys.stdout.buffer.write(chunk) + sys.stdout.buffer.flush() process.wait() if process.returncode == 0: return True @@ -52,6 +55,7 @@ class CLIBase(unittest.TestCase): model_path: str = None extra_args = [] data_type: DataType = None + log_level: str = "info" # tested on h100 width: int = 720 @@ -83,7 +87,7 @@ class CLIBase(unittest.TestCase): "--prompt", "A curious raccoon", "--save-output", - "--log-level=debug", + f"--log-level={self.log_level}", f"--width={self.width}", f"--height={self.height}", f"--output-path={self.output_path}", diff --git a/python/sglang/multimodal_gen/test/run_suite.py b/python/sglang/multimodal_gen/test/run_suite.py index 553224def..c32968e65 100644 --- a/python/sglang/multimodal_gen/test/run_suite.py +++ b/python/sglang/multimodal_gen/test/run_suite.py @@ -142,7 +142,8 @@ def collect_test_items(files, filter_expr=None): cmd.extend(["-k", filter_expr]) 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) # Check for collection errors @@ -186,91 +187,206 @@ def collect_test_items(files, filter_expr=None): 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: print("No files to run.") return 0 - base_cmd = [sys.executable, "-m", "pytest", "-s", "-v"] - if exitfirst: - base_cmd.append("-x") + base_cmd = [ + sys.executable, + "-m", + "pytest", + "-s", + "-v", + "--tb=short", + "--no-header", + ] - # Add pytest -k filter if provided if filter_expr: base_cmd.extend(["-k", filter_expr]) max_retries = 6 - # retry if the perf assertion failed, for {max_retries} times + attempt_reports = [] + for i in range(max_retries + 1): + is_retry = i > 0 cmd = list(base_cmd) - if i > 0: + if is_retry: 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) - if i > 0: - print( - f"Performance assertion failed. Retrying ({i}/{max_retries}) with --last-failed..." - ) - - print(f"Running command: {' '.join(cmd)}") - - process = subprocess.Popen( - cmd, - stdout=subprocess.PIPE, - stderr=subprocess.STDOUT, - bufsize=0, + mode = "retry failed items" if is_retry else "initial pass" + print( + f"Starting pytest attempt {i + 1}/{max_retries + 1}: {mode} " + f"for {len(files)} assigned item(s)" ) - 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() - returncode = process.returncode + returncode, full_output = _run_pytest_attempt(cmd) + retryable = returncode not in (0, 5) and _is_retryable_failure(full_output) + attempt_reports.append( + { + "attempt": i + 1, + "mode": mode, + "returncode": returncode, + "retryable": retryable, + "collection_line": _extract_collection_line(full_output), + "failure_tail": ( + _extract_failure_tail(full_output) + if returncode not in (0, 5) + else [] + ), + } + ) 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 - # 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: print( "No tests collected (exit code 5). This is expected when filters " "deselect all tests in a partition. Treating as success." ) + _print_attempt_tail_summary(attempt_reports, len(files)) return 0 - # check if the failure is due to an assertion in test_server_utils.py - full_output = output_bytes.decode("utf-8", errors="replace") - 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): + if not retryable: + _print_attempt_tail_summary(attempt_reports, len(files)) return returncode - print(f"Max retry exceeded") - return returncode + if i == max_retries: + 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): @@ -434,37 +550,19 @@ def main(): headers = ["Suite", "Partition"] rows = [[args.suite, partition_info]] msg = tabulate.tabulate(rows, headers=headers, tablefmt="psql") + "\n" - msg += f"✅ Enabled {len(my_items)} test(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:") + msg += f"✅ Assigned {len(my_items)} test(s) from {len(suite_files_abs)} file(s):\n" 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: print("No items assigned to this partition. Exiting success.") 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 - # Fast-fail: stop on first failure unless --continue-on-error is set - 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) + exit_code = run_pytest(my_items) sys.exit(exit_code) diff --git a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py index d0b38f0ca..4d37b5bc7 100644 --- a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py +++ b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py @@ -120,6 +120,9 @@ class TestTransformerQuantHelpers(unittest.TestCase): "sglang.multimodal_gen.runtime.loader.transformer_load_utils.build_nvfp4_config_from_safetensors_list", return_value=None, ) + @patch( + "sglang.multimodal_gen.runtime.loader.transformer_load_utils.maybe_download_model" + ) @patch( "sglang.multimodal_gen.runtime.loader.transformer_load_utils.get_quant_config_from_safetensors_metadata", return_value=None, @@ -136,8 +139,12 @@ class TestTransformerQuantHelpers(unittest.TestCase): _mock_download, mock_metadata, _mock_quant_metadata, + mock_maybe_download, _mock_nvfp4, ): + mock_maybe_download.side_effect = AssertionError( + "local safetensors path should not trigger maybe_download_model" + ) mock_metadata.return_value = { "config": json.dumps({"_class_name": _FakeFluxTransformer.__name__}) } @@ -163,6 +170,7 @@ class TestTransformerQuantHelpers(unittest.TestCase): self.assertIsNone(spec.param_dtype) self.assertEqual(len(spec.post_load_hooks), 1) self.assertIs(nunchaku_config.model_cls, _FakeFluxTransformer) + mock_maybe_download.assert_not_called() def test_flux2_mixed_nvfp4_fallback_disables_conflicting_offloads(self): server_args = self._make_server_args(