[diffusion] CI: fix ModelOpt B200 CI artifact coverage (#22955)
This commit is contained in:
@@ -40,6 +40,7 @@ PostLoadHook = Callable[[nn.Module], None]
|
||||
_PRECISION_VARIANT_SUFFIX_RE = re.compile(
|
||||
r"^(?P<stem>.+?)(?P<precision>\.(?:fp16|bf16|fp32))(?P<shard>-\d+-of-\d+)?(?P<ext>\.safetensors)$"
|
||||
)
|
||||
_MIXED_SAFETENSORS_RE = re.compile(r".*-mixed(?:-\d+-of-\d+)?\.safetensors$")
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -217,6 +218,7 @@ def resolve_transformer_safetensors_to_load(
|
||||
else:
|
||||
safetensors_list = _list_safetensors_files(component_model_path)
|
||||
|
||||
safetensors_list = _prefer_mixed_safetensors_files(safetensors_list)
|
||||
safetensors_list = _filter_duplicate_precision_variant_safetensors(safetensors_list)
|
||||
|
||||
if not safetensors_list:
|
||||
@@ -227,6 +229,31 @@ def resolve_transformer_safetensors_to_load(
|
||||
return safetensors_list
|
||||
|
||||
|
||||
def _prefer_mixed_safetensors_files(safetensors_list: list[str]) -> list[str]:
|
||||
"""Prefer mixed-precision transformer exports over sibling full exports.
|
||||
|
||||
Some raw ModelOpt NVFP4 repos ship both `foo-mixed.safetensors` and
|
||||
`foo.safetensors`. They are alternative full transformer exports, not
|
||||
shards, so loading both trips duplicate tensor-name validation.
|
||||
"""
|
||||
mixed_files = [
|
||||
path
|
||||
for path in safetensors_list
|
||||
if _MIXED_SAFETENSORS_RE.match(os.path.basename(path))
|
||||
]
|
||||
if not mixed_files or len(mixed_files) == len(safetensors_list):
|
||||
return safetensors_list
|
||||
|
||||
logger.info(
|
||||
"Using %d mixed transformer safetensors file(s) and ignoring %d sibling "
|
||||
"non-mixed file(s): %s",
|
||||
len(mixed_files),
|
||||
len(safetensors_list) - len(mixed_files),
|
||||
mixed_files,
|
||||
)
|
||||
return mixed_files
|
||||
|
||||
|
||||
def _filter_duplicate_precision_variant_safetensors(
|
||||
safetensors_list: list[str],
|
||||
) -> list[str]:
|
||||
|
||||
@@ -3,7 +3,7 @@ from sglang.multimodal_gen.test.server.testcase_configs import (
|
||||
MODELOPT_FLUX1_FP8_TRANSFORMER,
|
||||
MODELOPT_FLUX1_NVFP4_TRANSFORMER,
|
||||
MODELOPT_FLUX2_FP8_TRANSFORMER,
|
||||
MODELOPT_FLUX2_NVFP4_MODEL,
|
||||
MODELOPT_FLUX2_NVFP4_WEIGHTS,
|
||||
MODELOPT_NVFP4_B200_ENV_VARS,
|
||||
MODELOPT_WAN22_FP8_TRANSFORMER,
|
||||
MODELOPT_WAN22_NVFP4_TRANSFORMER,
|
||||
@@ -393,10 +393,10 @@ ONE_GPU_CASES_C = [
|
||||
),
|
||||
_make_modelopt_ci_case(
|
||||
"flux2_modelopt_nvfp4_t2i",
|
||||
model_path=MODELOPT_FLUX2_NVFP4_MODEL,
|
||||
model_path=DEFAULT_FLUX_2_DEV_MODEL_NAME_FOR_TEST,
|
||||
modality="image",
|
||||
sampling_params=MODELOPT_T2I_CI_sampling_params,
|
||||
extras=[],
|
||||
extras=["--transformer-weights-path", MODELOPT_FLUX2_NVFP4_WEIGHTS],
|
||||
env_vars=MODELOPT_NVFP4_B200_ENV_VARS,
|
||||
),
|
||||
_make_modelopt_ci_case(
|
||||
|
||||
@@ -707,6 +707,32 @@ Repository: https://github.com/sglang-bot/sglang-ci-data (path: diffusion-ci/con
|
||||
output_path.write_bytes(content)
|
||||
logger.info(f"Saved GT image: {output_path} (format: {detected_format})")
|
||||
|
||||
def _save_diffusion_artifact(
|
||||
self,
|
||||
case: DiffusionTestCase,
|
||||
content: bytes,
|
||||
) -> None:
|
||||
"""Preserve selected generated outputs for CI artifact upload."""
|
||||
artifact_dir = os.environ.get("SGLANG_DIFFUSION_ARTIFACT_DIR")
|
||||
if not artifact_dir or not content or "modelopt" not in case.id.lower():
|
||||
return
|
||||
|
||||
safe_case_id = "".join(c if c.isalnum() or c in "._-" else "_" for c in case.id)
|
||||
is_video = case.server_args.modality == "video"
|
||||
if is_video:
|
||||
filename = f"{safe_case_id}_5s.mp4"
|
||||
else:
|
||||
from sglang.multimodal_gen.test.test_utils import detect_image_format
|
||||
|
||||
suffix = case.sampling_params.output_format or detect_image_format(content)
|
||||
filename = f"{safe_case_id}.{suffix}"
|
||||
|
||||
dst_dir = Path(artifact_dir) / safe_case_id
|
||||
dst_dir.mkdir(parents=True, exist_ok=True)
|
||||
dst = dst_dir / filename
|
||||
dst.write_bytes(content)
|
||||
logger.info("[Artifact] Preserved generated output: %s", dst)
|
||||
|
||||
def _test_lora_api_functionality(
|
||||
self,
|
||||
ctx: ServerContext,
|
||||
@@ -1086,6 +1112,7 @@ Repository: https://github.com/sglang-bot/sglang-ci-data (path: diffusion-ci/con
|
||||
generate_fn,
|
||||
collect_perf=not is_gt_gen_mode,
|
||||
)
|
||||
self._save_diffusion_artifact(case, content)
|
||||
|
||||
if is_gt_gen_mode:
|
||||
# GT generation mode: save output and skip all validations/tests
|
||||
|
||||
@@ -364,7 +364,7 @@ T2I_sampling_params = DiffusionSamplingParams(
|
||||
MODELOPT_T2I_CI_sampling_params = DiffusionSamplingParams(
|
||||
prompt="Doraemon is eating dorayaki",
|
||||
output_size="768x768",
|
||||
extras={"num_inference_steps": 12},
|
||||
extras={"num_inference_steps": 12, "seed": 0},
|
||||
)
|
||||
|
||||
TI2I_sampling_params = DiffusionSamplingParams(
|
||||
@@ -406,8 +406,9 @@ T2V_sampling_params = DiffusionSamplingParams(
|
||||
MODELOPT_T2V_CI_sampling_params = DiffusionSamplingParams(
|
||||
prompt=T2V_PROMPT,
|
||||
output_size="640x384",
|
||||
seconds=5,
|
||||
num_frames=17,
|
||||
extras={"num_inference_steps": 12},
|
||||
extras={"num_inference_steps": 12, "seed": 0},
|
||||
)
|
||||
|
||||
TI2V_sampling_params = DiffusionSamplingParams(
|
||||
@@ -434,7 +435,7 @@ MODELOPT_FLUX1_FP8_TRANSFORMER = "BBuf/flux1-dev-modelopt-fp8-sglang-transformer
|
||||
MODELOPT_FLUX2_FP8_TRANSFORMER = "BBuf/flux2-dev-modelopt-fp8-sglang-transformer"
|
||||
MODELOPT_WAN22_FP8_TRANSFORMER = "BBuf/wan22-t2v-a14b-modelopt-fp8-sglang-transformer"
|
||||
MODELOPT_FLUX1_NVFP4_TRANSFORMER = "BBuf/flux1-dev-modelopt-nvfp4-sglang-transformer"
|
||||
MODELOPT_FLUX2_NVFP4_MODEL = "black-forest-labs/FLUX.2-dev-NVFP4"
|
||||
MODELOPT_FLUX2_NVFP4_WEIGHTS = "black-forest-labs/FLUX.2-dev-NVFP4"
|
||||
MODELOPT_WAN22_NVFP4_TRANSFORMER = (
|
||||
"BBuf/wan22-t2v-a14b-modelopt-nvfp4-sglang-transformer"
|
||||
)
|
||||
|
||||
@@ -100,6 +100,26 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
||||
|
||||
self.assertEqual(resolved, [f.name])
|
||||
|
||||
@patch(
|
||||
"sglang.multimodal_gen.runtime.loader.transformer_load_utils.maybe_download_model",
|
||||
side_effect=lambda path, **kw: path,
|
||||
)
|
||||
def test_resolve_transformer_safetensors_to_load_prefers_mixed_export(
|
||||
self, _mock_download
|
||||
):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
mixed = f"{tmpdir}/flux2-dev-nvfp4-mixed.safetensors"
|
||||
full = f"{tmpdir}/flux2-dev-nvfp4.safetensors"
|
||||
open(mixed, "a").close()
|
||||
open(full, "a").close()
|
||||
|
||||
server_args = self._make_server_args(transformer_weights_path=tmpdir)
|
||||
resolved = resolve_transformer_safetensors_to_load(
|
||||
server_args, "/unused/component/path"
|
||||
)
|
||||
|
||||
self.assertEqual(resolved, [mixed])
|
||||
|
||||
def test_filter_transformer_precision_variants_prefers_canonical_file(self):
|
||||
files = [
|
||||
"/tmp/transformer/diffusion_pytorch_model.fp16.safetensors",
|
||||
|
||||
Reference in New Issue
Block a user