Flux2 nvfp4 quantization correctness on Blackwell (B200) (#23625)
This commit is contained in:
@@ -275,19 +275,13 @@ jobs:
|
|||||||
artifact-suffix: component-accuracy
|
artifact-suffix: component-accuracy
|
||||||
|
|
||||||
multimodal-gen-test-1-b200:
|
multimodal-gen-test-1-b200:
|
||||||
# TODO: Temporarily disabled (`if: false`) — deterministically failing
|
if: |
|
||||||
# with `RobertaProcessing.__new__() got an unexpected keyword argument
|
(inputs.target_stage == 'multimodal-gen-test-1-b200') ||
|
||||||
# 'cls'` (tokenizer/transformers compatibility) and blocking CI.
|
(
|
||||||
# Restore the original `if:` condition (commented below) once the
|
!inputs.target_stage &&
|
||||||
# underlying fix lands.
|
((github.event_name == 'schedule' || inputs.test_parallel_dispatch == 'true') || (inputs.caller_needs_failure != 'true' && !cancelled())) &&
|
||||||
if: false
|
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 }}
|
runs-on: ${{ inputs.b200_runner }}
|
||||||
timeout-minutes: 240
|
timeout-minutes: 240
|
||||||
steps:
|
steps:
|
||||||
|
|||||||
@@ -46,7 +46,7 @@ class MLP(nn.Module):
|
|||||||
bias=True,
|
bias=True,
|
||||||
gather_output=False,
|
gather_output=False,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("0.proj", prefix),
|
prefix=add_prefix("fc_in", prefix),
|
||||||
)
|
)
|
||||||
|
|
||||||
self.act = get_act_fn(act_type)
|
self.act = get_act_fn(act_type)
|
||||||
@@ -58,7 +58,7 @@ class MLP(nn.Module):
|
|||||||
bias=True,
|
bias=True,
|
||||||
input_is_parallel=True,
|
input_is_parallel=True,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("2", prefix),
|
prefix=add_prefix("fc_out", prefix),
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
|||||||
@@ -475,7 +475,7 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
|
|||||||
output_dim=0,
|
output_dim=0,
|
||||||
weight_loader=weight_loader,
|
weight_loader=weight_loader,
|
||||||
)
|
)
|
||||||
|
set_weight_attrs(weight_scale, {"missing_param_init": "ones"})
|
||||||
layer.register_parameter("weight_scale", weight_scale)
|
layer.register_parameter("weight_scale", weight_scale)
|
||||||
|
|
||||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||||
@@ -552,11 +552,15 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
|
|||||||
K_padded = round_up(K, 4)
|
K_padded = round_up(K, 4)
|
||||||
padded_scales = torch.zeros((B, M_padded, K_padded), dtype=scales.dtype)
|
padded_scales = torch.zeros((B, M_padded, K_padded), dtype=scales.dtype)
|
||||||
padded_scales[:B, :M, :K] = scales
|
padded_scales[:B, :M, :K] = scales
|
||||||
# Blockwise interleave for CUTLASS TMA layout required by CUTLASS kernel
|
|
||||||
padded_scales = padded_scales.reshape(
|
_, flashinfer_backend = _get_fp4_gemm_op()
|
||||||
B, M_padded // 128, 4, 32, K_padded // 4, 4
|
if flashinfer_backend is None:
|
||||||
)
|
# CUTLASS (sgl_kernel) path: blockwise interleave to TMA layout
|
||||||
padded_scales = padded_scales.permute(0, 1, 4, 3, 2, 5)
|
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.contiguous().cuda()
|
||||||
padded_scales = (
|
padded_scales = (
|
||||||
padded_scales.reshape(M_padded, K_padded)
|
padded_scales.reshape(M_padded, K_padded)
|
||||||
|
|||||||
@@ -329,11 +329,26 @@ class TokenizerLoader(ComponentLoader):
|
|||||||
# Qwen-Image's model_index declares Qwen2Tokenizer; using the fast class
|
# Qwen-Image's model_index declares Qwen2Tokenizer; using the fast class
|
||||||
# changes text preprocessing and shifts official GT comparisons.
|
# changes text preprocessing and shifts official GT comparisons.
|
||||||
use_fast = self.component_architecture != "Qwen2Tokenizer"
|
use_fast = self.component_architecture != "Qwen2Tokenizer"
|
||||||
return AutoTokenizer.from_pretrained(
|
try:
|
||||||
component_model_path,
|
return AutoTokenizer.from_pretrained(
|
||||||
padding_side="right",
|
component_model_path,
|
||||||
use_fast=use_fast,
|
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):
|
class GenericComponentLoader(ComponentLoader):
|
||||||
|
|||||||
@@ -206,9 +206,76 @@ def _clean_hf_config_inplace(model_config: dict) -> None:
|
|||||||
model_config.pop(key, 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]:
|
def _list_safetensors_files(model_path: str) -> list[str]:
|
||||||
"""List all .safetensors files under a directory."""
|
"""List all .safetensors files under a directory.
|
||||||
return sorted(glob.glob(os.path.join(str(model_path), "*.safetensors")))
|
|
||||||
|
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
|
BYTES_PER_GB = 1024**3
|
||||||
|
|||||||
@@ -182,7 +182,7 @@ class WanSelfAttention(nn.Module):
|
|||||||
dim,
|
dim,
|
||||||
input_is_parallel=True,
|
input_is_parallel=True,
|
||||||
quant_config=quant_config,
|
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_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||||
self.norm_k = 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,
|
bias=True,
|
||||||
gather_output=False,
|
gather_output=False,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("attn1.to_q", prefix),
|
prefix=add_prefix("to_q", prefix),
|
||||||
)
|
)
|
||||||
self.to_k = ColumnParallelLinear(
|
self.to_k = ColumnParallelLinear(
|
||||||
dim,
|
dim,
|
||||||
@@ -374,7 +374,7 @@ class WanTransformerBlock(nn.Module):
|
|||||||
bias=True,
|
bias=True,
|
||||||
gather_output=False,
|
gather_output=False,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("attn1.to_k", prefix),
|
prefix=add_prefix("to_k", prefix),
|
||||||
)
|
)
|
||||||
self.to_v = ColumnParallelLinear(
|
self.to_v = ColumnParallelLinear(
|
||||||
dim,
|
dim,
|
||||||
@@ -382,7 +382,7 @@ class WanTransformerBlock(nn.Module):
|
|||||||
bias=True,
|
bias=True,
|
||||||
gather_output=False,
|
gather_output=False,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("attn1.to_v", prefix),
|
prefix=add_prefix("to_v", prefix),
|
||||||
)
|
)
|
||||||
|
|
||||||
self.to_out = RowParallelLinear(
|
self.to_out = RowParallelLinear(
|
||||||
@@ -391,7 +391,7 @@ class WanTransformerBlock(nn.Module):
|
|||||||
bias=True,
|
bias=True,
|
||||||
reduce_results=True,
|
reduce_results=True,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("attn1.to_out.0", prefix),
|
prefix=add_prefix("to_out", prefix),
|
||||||
)
|
)
|
||||||
tp_size = get_tp_world_size()
|
tp_size = get_tp_world_size()
|
||||||
self.local_num_heads = divide(num_heads, tp_size)
|
self.local_num_heads = divide(num_heads, tp_size)
|
||||||
@@ -481,7 +481,7 @@ class WanTransformerBlock(nn.Module):
|
|||||||
dim,
|
dim,
|
||||||
ffn_dim,
|
ffn_dim,
|
||||||
act_type="gelu_pytorch_tanh",
|
act_type="gelu_pytorch_tanh",
|
||||||
prefix=add_prefix("ffn.net", prefix),
|
prefix=add_prefix("ffn", prefix),
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
)
|
)
|
||||||
self.mlp_residual = MulAdd()
|
self.mlp_residual = MulAdd()
|
||||||
@@ -640,7 +640,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
|||||||
bias=True,
|
bias=True,
|
||||||
gather_output=True,
|
gather_output=True,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("attn1.to_q", prefix),
|
prefix=add_prefix("to_q", prefix),
|
||||||
)
|
)
|
||||||
self.to_k = ColumnParallelLinear(
|
self.to_k = ColumnParallelLinear(
|
||||||
dim,
|
dim,
|
||||||
@@ -648,7 +648,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
|||||||
bias=True,
|
bias=True,
|
||||||
gather_output=True,
|
gather_output=True,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("attn1.to_k", prefix),
|
prefix=add_prefix("to_k", prefix),
|
||||||
)
|
)
|
||||||
self.to_v = ColumnParallelLinear(
|
self.to_v = ColumnParallelLinear(
|
||||||
dim,
|
dim,
|
||||||
@@ -656,7 +656,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
|||||||
bias=True,
|
bias=True,
|
||||||
gather_output=True,
|
gather_output=True,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("attn1.to_v", prefix),
|
prefix=add_prefix("to_v", prefix),
|
||||||
)
|
)
|
||||||
self.to_gate_compress = ColumnParallelLinear(
|
self.to_gate_compress = ColumnParallelLinear(
|
||||||
dim,
|
dim,
|
||||||
@@ -673,7 +673,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
|||||||
bias=True,
|
bias=True,
|
||||||
gather_output=True,
|
gather_output=True,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("attn1.to_out.0", prefix),
|
prefix=add_prefix("to_out", prefix),
|
||||||
)
|
)
|
||||||
self.attn1 = UlyssesAttention_VSA(
|
self.attn1 = UlyssesAttention_VSA(
|
||||||
num_heads=num_heads,
|
num_heads=num_heads,
|
||||||
@@ -742,7 +742,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
|||||||
dim,
|
dim,
|
||||||
ffn_dim,
|
ffn_dim,
|
||||||
act_type="gelu_pytorch_tanh",
|
act_type="gelu_pytorch_tanh",
|
||||||
prefix=add_prefix("ffn.net", prefix),
|
prefix=add_prefix("ffn", prefix),
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
)
|
)
|
||||||
self.mlp_residual = MulAdd()
|
self.mlp_residual = MulAdd()
|
||||||
|
|||||||
@@ -403,7 +403,9 @@ def _build_nvfp4_config_from_safetensors_files(
|
|||||||
if mapping_fn is not None:
|
if mapping_fn is not None:
|
||||||
mapped, _, _ = mapping_fn(raw_weight_name)
|
mapped, _, _ = mapping_fn(raw_weight_name)
|
||||||
if mapped != 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
|
continue
|
||||||
|
|
||||||
if reverse_mapping_fn is not None:
|
if reverse_mapping_fn is not None:
|
||||||
|
|||||||
@@ -31,12 +31,16 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
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_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/sglang-bot/sglang-ci-data/main/diffusion-ci/consistency_gt/sglang_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_BASE = SGL_TEST_FILES_SGLANG_CONSISTENCY_GT_BASE
|
||||||
SGL_TEST_FILES_CONSISTENCY_GT_BASES = (
|
SGL_TEST_FILES_CONSISTENCY_GT_BASES = (
|
||||||
SGL_TEST_FILES_OFFICIAL_CONSISTENCY_GT_BASE,
|
SGL_TEST_FILES_OFFICIAL_CONSISTENCY_GT_BASE,
|
||||||
SGL_TEST_FILES_SGLANG_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 = (
|
CONSISTENCY_THRESHOLD_JSON_PATH = (
|
||||||
Path(__file__).resolve().parent / "server" / "consistency_threshold.json"
|
Path(__file__).resolve().parent / "server" / "consistency_threshold.json"
|
||||||
|
|||||||
Reference in New Issue
Block a user