diff --git a/.github/workflows/pr-test-multimodal-gen.yml b/.github/workflows/pr-test-multimodal-gen.yml index ec6dd7f79..110c812c0 100644 --- a/.github/workflows/pr-test-multimodal-gen.yml +++ b/.github/workflows/pr-test-multimodal-gen.yml @@ -275,19 +275,13 @@ jobs: artifact-suffix: component-accuracy multimodal-gen-test-1-b200: - # TODO: Temporarily disabled (`if: false`) — deterministically failing - # with `RobertaProcessing.__new__() got an unexpected keyword argument - # 'cls'` (tokenizer/transformers compatibility) and blocking CI. - # Restore the original `if:` condition (commented below) once the - # underlying fix lands. - if: false - # if: | - # (inputs.target_stage == 'multimodal-gen-test-1-b200') || - # ( - # !inputs.target_stage && - # ((github.event_name == 'schedule' || inputs.test_parallel_dispatch == 'true') || (inputs.caller_needs_failure != 'true' && !cancelled())) && - # inputs.multimodal_gen == 'true' - # ) + if: | + (inputs.target_stage == 'multimodal-gen-test-1-b200') || + ( + !inputs.target_stage && + ((github.event_name == 'schedule' || inputs.test_parallel_dispatch == 'true') || (inputs.caller_needs_failure != 'true' && !cancelled())) && + inputs.multimodal_gen == 'true' + ) runs-on: ${{ inputs.b200_runner }} timeout-minutes: 240 steps: diff --git a/python/sglang/multimodal_gen/runtime/layers/mlp.py b/python/sglang/multimodal_gen/runtime/layers/mlp.py index e1e3e58f5..11c1739fe 100644 --- a/python/sglang/multimodal_gen/runtime/layers/mlp.py +++ b/python/sglang/multimodal_gen/runtime/layers/mlp.py @@ -46,7 +46,7 @@ class MLP(nn.Module): bias=True, gather_output=False, quant_config=quant_config, - prefix=add_prefix("0.proj", prefix), + prefix=add_prefix("fc_in", prefix), ) self.act = get_act_fn(act_type) @@ -58,7 +58,7 @@ class MLP(nn.Module): bias=True, input_is_parallel=True, quant_config=quant_config, - prefix=add_prefix("2", prefix), + prefix=add_prefix("fc_out", prefix), ) def forward(self, x: torch.Tensor) -> torch.Tensor: diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py b/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py index 2f5da4168..c9286aba5 100755 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py @@ -475,7 +475,7 @@ class ModelOptFp4LinearMethod(LinearMethodBase): output_dim=0, weight_loader=weight_loader, ) - + set_weight_attrs(weight_scale, {"missing_param_init": "ones"}) layer.register_parameter("weight_scale", weight_scale) def process_weights_after_loading(self, layer: torch.nn.Module) -> None: @@ -552,11 +552,15 @@ class ModelOptFp4LinearMethod(LinearMethodBase): K_padded = round_up(K, 4) padded_scales = torch.zeros((B, M_padded, K_padded), dtype=scales.dtype) padded_scales[:B, :M, :K] = scales - # Blockwise interleave for CUTLASS TMA layout required by CUTLASS kernel - padded_scales = padded_scales.reshape( - B, M_padded // 128, 4, 32, K_padded // 4, 4 - ) - padded_scales = padded_scales.permute(0, 1, 4, 3, 2, 5) + + _, flashinfer_backend = _get_fp4_gemm_op() + if flashinfer_backend is None: + # CUTLASS (sgl_kernel) path: blockwise interleave to TMA layout + padded_scales = padded_scales.reshape( + B, M_padded // 128, 4, 32, K_padded // 4, 4 + ) + padded_scales = padded_scales.permute(0, 1, 4, 3, 2, 5) + padded_scales = padded_scales.contiguous().cuda() padded_scales = ( padded_scales.reshape(M_padded, K_padded) diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py index 48d2beb50..1b409dcad 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py @@ -329,11 +329,26 @@ class TokenizerLoader(ComponentLoader): # Qwen-Image's model_index declares Qwen2Tokenizer; using the fast class # changes text preprocessing and shifts official GT comparisons. use_fast = self.component_architecture != "Qwen2Tokenizer" - return AutoTokenizer.from_pretrained( - component_model_path, - padding_side="right", - use_fast=use_fast, - ) + try: + return AutoTokenizer.from_pretrained( + component_model_path, + padding_side="right", + use_fast=use_fast, + ) + except TypeError as e: + # tokenizers>=0.21 removed the `cls` kwarg from RobertaProcessing, + # but some transformers CLIPTokenizer builds still pass it. Fall back + # to the pure-Python (slow) tokenizer which avoids the rust path. + if "RobertaProcessing" in str(e) and use_fast: + logger.warning( + "Fast tokenizer failed (%s), retrying with use_fast=False", e + ) + return AutoTokenizer.from_pretrained( + component_model_path, + padding_side="right", + use_fast=False, + ) + raise class GenericComponentLoader(ComponentLoader): diff --git a/python/sglang/multimodal_gen/runtime/loader/utils.py b/python/sglang/multimodal_gen/runtime/loader/utils.py index 116b54d8c..7f4bddcc8 100644 --- a/python/sglang/multimodal_gen/runtime/loader/utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/utils.py @@ -206,9 +206,76 @@ def _clean_hf_config_inplace(model_config: dict) -> None: model_config.pop(key, None) +def _try_redownload_missing_shards(model_path: str, missing: list[str]) -> bool: + """Try to re-download missing safetensors shards from HuggingFace Hub. + + Parses the repo_id and revision from the HF cache path structure + (models--{org}--{repo}/snapshots/{revision}) and calls hf_hub_download + for each missing shard. Returns True if all shards were recovered. + """ + try: + from huggingface_hub import hf_hub_download + + match = re.search( + r"models--([^/\\]+)--([^/\\]+)[/\\]snapshots[/\\]([^/\\]+)", model_path + ) + if not match: + return False + + repo_id = f"{match.group(1)}/{match.group(2)}" + revision = match.group(3) + logger.warning( + "Incomplete checkpoint for %s (revision %.8s) — missing shards: %s. " + "Attempting auto-repair via HuggingFace Hub...", + repo_id, + revision, + missing, + ) + for shard in missing: + hf_hub_download(repo_id=repo_id, filename=shard, revision=revision) + logger.info("Auto-repair succeeded for %s.", repo_id) + return True + except Exception as e: + logger.warning("Auto-repair failed: %s", e) + return False + + def _list_safetensors_files(model_path: str) -> list[str]: - """List all .safetensors files under a directory.""" - return sorted(glob.glob(os.path.join(str(model_path), "*.safetensors"))) + """List all .safetensors files under a directory. + + If a safetensors index file is present, verifies that every shard listed + in the index actually exists on disk. Missing shards are first repaired + automatically via HuggingFace Hub (if the path is an HF cache entry); + if repair fails a clear RuntimeError is raised. + """ + found = sorted(glob.glob(os.path.join(str(model_path), "*.safetensors"))) + + index_path = os.path.join( + str(model_path), "diffusion_pytorch_model.safetensors.index.json" + ) + if os.path.exists(index_path): + import json + + with open(index_path) as f: + index = json.load(f) + expected_shards = sorted(set(index.get("weight_map", {}).values())) + found_basenames = {os.path.basename(p) for p in found} + missing = [s for s in expected_shards if s not in found_basenames] + if missing: + repaired = _try_redownload_missing_shards(model_path, missing) + if repaired: + found = sorted( + glob.glob(os.path.join(str(model_path), "*.safetensors")) + ) + else: + raise RuntimeError( + f"Checkpoint at '{model_path}' is incomplete — the following " + f"shard(s) listed in the index are missing from disk: " + f"{missing}. Re-download the checkpoint (e.g. " + f"`huggingface-cli download {os.path.basename(model_path)}`)." + ) + + return found BYTES_PER_GB = 1024**3 diff --git a/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py b/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py index eff9835b9..c38ac2356 100755 --- a/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py @@ -182,7 +182,7 @@ class WanSelfAttention(nn.Module): dim, input_is_parallel=True, quant_config=quant_config, - prefix=add_prefix("to_out.0", prefix), + prefix=add_prefix("to_out", prefix), ) self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity() self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity() @@ -366,7 +366,7 @@ class WanTransformerBlock(nn.Module): bias=True, gather_output=False, quant_config=quant_config, - prefix=add_prefix("attn1.to_q", prefix), + prefix=add_prefix("to_q", prefix), ) self.to_k = ColumnParallelLinear( dim, @@ -374,7 +374,7 @@ class WanTransformerBlock(nn.Module): bias=True, gather_output=False, quant_config=quant_config, - prefix=add_prefix("attn1.to_k", prefix), + prefix=add_prefix("to_k", prefix), ) self.to_v = ColumnParallelLinear( dim, @@ -382,7 +382,7 @@ class WanTransformerBlock(nn.Module): bias=True, gather_output=False, quant_config=quant_config, - prefix=add_prefix("attn1.to_v", prefix), + prefix=add_prefix("to_v", prefix), ) self.to_out = RowParallelLinear( @@ -391,7 +391,7 @@ class WanTransformerBlock(nn.Module): bias=True, reduce_results=True, quant_config=quant_config, - prefix=add_prefix("attn1.to_out.0", prefix), + prefix=add_prefix("to_out", prefix), ) tp_size = get_tp_world_size() self.local_num_heads = divide(num_heads, tp_size) @@ -481,7 +481,7 @@ class WanTransformerBlock(nn.Module): dim, ffn_dim, act_type="gelu_pytorch_tanh", - prefix=add_prefix("ffn.net", prefix), + prefix=add_prefix("ffn", prefix), quant_config=quant_config, ) self.mlp_residual = MulAdd() @@ -640,7 +640,7 @@ class WanTransformerBlock_VSA(nn.Module): bias=True, gather_output=True, quant_config=quant_config, - prefix=add_prefix("attn1.to_q", prefix), + prefix=add_prefix("to_q", prefix), ) self.to_k = ColumnParallelLinear( dim, @@ -648,7 +648,7 @@ class WanTransformerBlock_VSA(nn.Module): bias=True, gather_output=True, quant_config=quant_config, - prefix=add_prefix("attn1.to_k", prefix), + prefix=add_prefix("to_k", prefix), ) self.to_v = ColumnParallelLinear( dim, @@ -656,7 +656,7 @@ class WanTransformerBlock_VSA(nn.Module): bias=True, gather_output=True, quant_config=quant_config, - prefix=add_prefix("attn1.to_v", prefix), + prefix=add_prefix("to_v", prefix), ) self.to_gate_compress = ColumnParallelLinear( dim, @@ -673,7 +673,7 @@ class WanTransformerBlock_VSA(nn.Module): bias=True, gather_output=True, quant_config=quant_config, - prefix=add_prefix("attn1.to_out.0", prefix), + prefix=add_prefix("to_out", prefix), ) self.attn1 = UlyssesAttention_VSA( num_heads=num_heads, @@ -742,7 +742,7 @@ class WanTransformerBlock_VSA(nn.Module): dim, ffn_dim, act_type="gelu_pytorch_tanh", - prefix=add_prefix("ffn.net", prefix), + prefix=add_prefix("ffn", prefix), quant_config=quant_config, ) self.mlp_residual = MulAdd() diff --git a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py index 2ddf29687..033728f2d 100644 --- a/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py +++ b/python/sglang/multimodal_gen/runtime/utils/quantization_utils.py @@ -403,7 +403,9 @@ def _build_nvfp4_config_from_safetensors_files( if mapping_fn is not None: mapped, _, _ = mapping_fn(raw_weight_name) if mapped != raw_weight_name: - exclude_modules.append(module_bfl) + exclude_modules.append( + mapped[: -len(".weight")] if mapped.endswith(".weight") else mapped + ) continue if reverse_mapping_fn is not None: diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index d9cc0ab89..71c96f71b 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -31,12 +31,16 @@ if TYPE_CHECKING: logger = init_logger(__name__) -SGL_TEST_FILES_OFFICIAL_CONSISTENCY_GT_BASE = "https://raw.githubusercontent.com/sglang-bot/sglang-ci-data/main/diffusion-ci/consistency_gt/official_generated" -SGL_TEST_FILES_SGLANG_CONSISTENCY_GT_BASE = "https://raw.githubusercontent.com/sglang-bot/sglang-ci-data/main/diffusion-ci/consistency_gt/sglang_generated" +SGL_TEST_FILES_OFFICIAL_CONSISTENCY_GT_BASE = "https://raw.githubusercontent.com/sgl-project/ci-data/main/diffusion-ci/consistency_gt/official_generated" +SGL_TEST_FILES_SGLANG_CONSISTENCY_GT_BASE = "https://raw.githubusercontent.com/sgl-project/ci-data/main/diffusion-ci/consistency_gt/sglang_generated" SGL_TEST_FILES_CONSISTENCY_GT_BASE = SGL_TEST_FILES_SGLANG_CONSISTENCY_GT_BASE SGL_TEST_FILES_CONSISTENCY_GT_BASES = ( SGL_TEST_FILES_OFFICIAL_CONSISTENCY_GT_BASE, SGL_TEST_FILES_SGLANG_CONSISTENCY_GT_BASE, + # Legacy fallback during migration from sglang-bot/sglang-ci-data + "https://raw.githubusercontent.com/sglang-bot/sglang-ci-data/main/diffusion-ci/consistency_gt/official_generated", + "https://raw.githubusercontent.com/sglang-bot/sglang-ci-data/main/diffusion-ci/consistency_gt/sglang_generated", + "https://raw.githubusercontent.com/sglang-bot/sglang-ci-data/main/diffusion-ci/consistency_gt", ) CONSISTENCY_THRESHOLD_JSON_PATH = ( Path(__file__).resolve().parent / "server" / "consistency_threshold.json"