Compare commits
15
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
70bd4bb756 | ||
|
|
35e55aa7f6 | ||
|
|
67c023fca9 | ||
|
|
bf839dde18 | ||
|
|
b1d954c894 | ||
|
|
e4eef6ed20 | ||
|
|
7390f528c6 | ||
|
|
92168e5020 | ||
|
|
d4b8f08674 | ||
|
|
13daf39e91 | ||
|
|
2d0e552a7e | ||
|
|
46d1f91c7e | ||
|
|
9307aeb590 | ||
|
|
926f0c9ab4 | ||
|
|
b2ba79a4cd |
@@ -34,6 +34,7 @@ env
|
||||
*.log
|
||||
weights/
|
||||
logs/
|
||||
/Z-Image/
|
||||
official_weights/
|
||||
converted_weights/
|
||||
|
||||
|
||||
@@ -83,6 +83,7 @@ class TextEncoderConfig(EncoderConfig):
|
||||
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
|
||||
is_chat_model: bool = False
|
||||
treat_empty_as_dot: bool = False
|
||||
chat_template_enable_thinking: bool = field(default=False, kw_only=True)
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -327,14 +327,12 @@ class Qwen3ForCausalLM(TextEncoder):
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> nn.Module:
|
||||
from transformers import AutoModelForCausalLM
|
||||
from transformers import AutoModel
|
||||
|
||||
if device.type == "cpu" and torch.cuda.is_available():
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
|
||||
device = get_local_torch_device()
|
||||
|
||||
return AutoModelForCausalLM.from_pretrained(
|
||||
# FastVideo uses Qwen3 only as a text encoder. Loading the body avoids
|
||||
# materializing an unused LM head and full-vocabulary logits, including
|
||||
# for checkpoints whose metadata names Qwen3ForCausalLM.
|
||||
return AutoModel.from_pretrained(
|
||||
model_path,
|
||||
local_files_only=True,
|
||||
torch_dtype=dtype,
|
||||
@@ -368,9 +366,18 @@ class Qwen3ForCausalLM(TextEncoder):
|
||||
residual = None
|
||||
|
||||
if position_ids is None:
|
||||
# Expand to [batch_size, seq_len]: the rotary layer flattens
|
||||
# positions to ``num_tokens`` and reshapes q/k to
|
||||
# ``(num_tokens, -1, head_dim)``. A bare [1, seq_len] only matches
|
||||
# ``num_tokens`` when batch_size == 1; for batched inputs it folds
|
||||
# the batch dim into the head dim and misaligns RoPE. Expanding to
|
||||
# ``batch_size * seq_len`` tokens keeps the layout correct.
|
||||
position_ids = torch.arange(
|
||||
0, hidden_states.shape[1], device=hidden_states.device
|
||||
).unsqueeze(0)
|
||||
0,
|
||||
hidden_states.shape[1],
|
||||
device=hidden_states.device,
|
||||
dtype=torch.long,
|
||||
).unsqueeze(0).expand(hidden_states.shape[0], -1)
|
||||
|
||||
all_hidden_states: tuple[Any, ...] | None = (
|
||||
() if output_hidden_states else None
|
||||
@@ -405,6 +412,20 @@ class Qwen3ForCausalLM(TextEncoder):
|
||||
) -> set[str]:
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
stacked_params_mapping = self.config.arch_config.stacked_params_mapping
|
||||
# A fused destination is initialized either by one already-fused tensor
|
||||
# or after every split source projection has been loaded. Include
|
||||
# auxiliary quantization parameters (for example scale_weight) rather
|
||||
# than limiting completeness checks to weight/bias tensors.
|
||||
expected_stacked_shards = {
|
||||
(name, shard_id)
|
||||
for name in params_dict
|
||||
for param_name, _, shard_id in stacked_params_mapping
|
||||
if param_name in name
|
||||
}
|
||||
fused_param_names = {name for name, _ in expected_stacked_shards}
|
||||
loaded_stacked_shards: set[tuple[str, str | int]] = set()
|
||||
loaded_fused_params: set[str] = set()
|
||||
|
||||
for name, loaded_weight in weights:
|
||||
if name.startswith("model."):
|
||||
@@ -423,37 +444,66 @@ class Qwen3ForCausalLM(TextEncoder):
|
||||
continue
|
||||
name = kv_scale_name
|
||||
|
||||
for (
|
||||
param_name,
|
||||
weight_name,
|
||||
shard_id,
|
||||
) in self.config.arch_config.stacked_params_mapping:
|
||||
matched_stacked_param = False
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
matched_stacked_param = True
|
||||
target_name = name.replace(weight_name, param_name)
|
||||
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
if target_name.endswith(".bias") and target_name not in params_dict:
|
||||
break
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
if target_name not in params_dict:
|
||||
break
|
||||
|
||||
param = params_dict[name]
|
||||
param = params_dict[target_name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
loaded_params.add(target_name)
|
||||
loaded_stacked_shards.add((target_name, shard_id))
|
||||
break
|
||||
|
||||
if matched_stacked_param:
|
||||
continue
|
||||
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
if name in fused_param_names and name.endswith(".scale_weight"):
|
||||
# Merged scale loaders interpret a missing shard id as shard 0.
|
||||
# An exact fused key is already a complete vector, so copy it
|
||||
# atomically and retain the default loader's shape validation.
|
||||
default_weight_loader(param, loaded_weight)
|
||||
else:
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
|
||||
loaded_params.add(name)
|
||||
if name in fused_param_names:
|
||||
loaded_fused_params.add(name)
|
||||
|
||||
required_split_shards = {
|
||||
(name, shard_id)
|
||||
for name, shard_id in expected_stacked_shards
|
||||
if name not in loaded_fused_params
|
||||
}
|
||||
missing_stacked_shards = required_split_shards - loaded_stacked_shards
|
||||
if missing_stacked_shards:
|
||||
formatted_missing = ", ".join(
|
||||
f"{name}[{shard_id}]"
|
||||
for name, shard_id in sorted(
|
||||
missing_stacked_shards,
|
||||
key=lambda item: (item[0], str(item[1])),
|
||||
)
|
||||
)
|
||||
raise ValueError(
|
||||
"Missing required stacked checkpoint shards: "
|
||||
f"{formatted_missing}"
|
||||
)
|
||||
|
||||
return loaded_params
|
||||
|
||||
|
||||
@@ -401,6 +401,10 @@ class TextEncoderLoader(ComponentLoader):
|
||||
dtype=PRECISION_TO_TYPE[dtype],
|
||||
device=target_device,
|
||||
)
|
||||
# HF passthrough encoders return before FastVideo's FSDP
|
||||
# wrapping path, so the text stage needs their placement to
|
||||
# put token tensors on the same device.
|
||||
model._fastvideo_input_device = target_device
|
||||
return model.eval()
|
||||
|
||||
model = model_cls(model_config) # type: ignore
|
||||
|
||||
@@ -78,6 +78,9 @@ _TEXT_ENCODER_MODELS = {
|
||||
"Reason1TextEncoder": ("encoders", "reason1", "Reason1TextEncoder"),
|
||||
"Qwen2_5_VLForConditionalGeneration":
|
||||
("encoders", "reason1", "Reason1TextEncoder"),
|
||||
# Z-Image-Turbo's text_encoder/config.json declares architecture
|
||||
# "Qwen3Model"; route it to the shared Qwen3 encoder (added for Flux2 Klein).
|
||||
"Qwen3Model": ("encoders", "qwen3", "Qwen3ForCausalLM"),
|
||||
"LTX2GemmaTextEncoderModel": ("encoders", "gemma", "LTX2GemmaTextEncoderModel"),
|
||||
"Qwen3ForCausalLM": ("encoders", "qwen3", "Qwen3ForCausalLM"),
|
||||
"Mistral3ForConditionalGeneration":
|
||||
|
||||
@@ -96,6 +96,14 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
The minimum sigma value for the noise schedule.
|
||||
sigma_data (`float`, *optional*):
|
||||
The sigma data value for scaling.
|
||||
use_reference_discrete_timesteps (`bool`, defaults to False):
|
||||
Some reference schedulers (e.g. Z-Image) construct the timestep
|
||||
schedule by linspacing `num_inference_steps + 1` points from
|
||||
`t_max` to `t_min` and dropping the terminal point. Default
|
||||
(`False`) preserves the original `np.linspace(t_max, t_min,
|
||||
num_inference_steps)` (float64) behaviour used by every existing
|
||||
model. Enable this flag only when matching a reference scheduler
|
||||
that expects the +1 + drop-terminal construction.
|
||||
"""
|
||||
|
||||
_compatibles: list[Any] = []
|
||||
@@ -122,6 +130,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
sigma_max: float | None = None,
|
||||
sigma_min: float | None = None,
|
||||
sigma_data: float | None = None,
|
||||
use_reference_discrete_timesteps: bool = False,
|
||||
):
|
||||
if sum([
|
||||
self.config.use_beta_sigmas, self.config.use_exponential_sigmas,
|
||||
@@ -155,7 +164,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
|
||||
self.sigmas = sigmas.to(
|
||||
"cpu") # to avoid too much CPU/GPU communication
|
||||
self.sigma_min = self.sigmas[-1].item()
|
||||
self.sigma_min = sigma_min if sigma_min is not None else self.sigmas[-1].item()
|
||||
self.sigma_max = self.sigmas[0].item()
|
||||
|
||||
BaseScheduler.__init__(self)
|
||||
@@ -350,7 +359,19 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
if timesteps_array is None:
|
||||
t_max = self._sigma_to_t(self.sigma_max)
|
||||
t_min = self._sigma_to_t(self.sigma_min)
|
||||
timesteps_array = np.linspace(t_max, t_min, num_inference_steps)
|
||||
if self.config.use_reference_discrete_timesteps:
|
||||
# Some reference schedulers (for example Z-Image) build a
|
||||
# float64 num_steps+1 linspace and drop the terminal point.
|
||||
timesteps_array = np.linspace(
|
||||
t_max,
|
||||
t_min,
|
||||
num_inference_steps + 1,
|
||||
)[:-1]
|
||||
else:
|
||||
# Preserve the original numpy default (float64) here —
|
||||
# casting to float32 silently shifts rounded timestep
|
||||
# values for every existing model that uses this branch.
|
||||
timesteps_array = np.linspace(t_max, t_min, num_inference_steps)
|
||||
sigmas_array = timesteps_array / self.config.num_train_timesteps
|
||||
else:
|
||||
sigmas_array = np.array(sigmas).astype(np.float32)
|
||||
|
||||
@@ -73,25 +73,36 @@ class DecodingStage(PipelineStage):
|
||||
|
||||
cfg = getattr(self.vae, "config", None)
|
||||
|
||||
# Matrix-Game 2.0-style: z = z * std + mean
|
||||
if (cfg is not None and hasattr(cfg, "latents_mean") and hasattr(cfg, "latents_std")):
|
||||
latents_mean = torch.tensor(cfg.latents_mean, device=latents.device,
|
||||
# Matrix-Game 2.0-style: z = z * std + mean. Some configs declare these
|
||||
# fields as None, so test the values rather than mere attribute presence
|
||||
# -- otherwise torch.tensor(None) raises instead of falling through.
|
||||
latents_mean_value = getattr(cfg, "latents_mean", None) if cfg is not None else None
|
||||
latents_std_value = getattr(cfg, "latents_std", None) if cfg is not None else None
|
||||
if latents_mean_value is not None and latents_std_value is not None:
|
||||
latents_mean = torch.tensor(latents_mean_value, device=latents.device,
|
||||
dtype=latents.dtype).view(1, -1, 1, 1, 1)
|
||||
latents_std = torch.tensor(cfg.latents_std, device=latents.device, dtype=latents.dtype).view(1, -1, 1, 1, 1)
|
||||
latents_std = torch.tensor(latents_std_value, device=latents.device,
|
||||
dtype=latents.dtype).view(1, -1, 1, 1, 1)
|
||||
return latents * latents_std + latents_mean
|
||||
|
||||
# Diffusers-style: scaling_factor (+ optional shift_factor)
|
||||
if hasattr(self.vae, "scaling_factor"):
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(latents.device, latents.dtype)
|
||||
# Diffusers-style: scaling_factor (+ optional shift_factor), read from
|
||||
# the module only. Nearly every VAE *config* also carries a
|
||||
# scaling_factor whose decode() already accounts for it, so sourcing it
|
||||
# from cfg here would newly divide latents for VAEs that have always
|
||||
# passed through untouched (flux2, hunyuanvideo, hunyuanvideo15, ltx2).
|
||||
scaling_factor = getattr(self.vae, "scaling_factor", None)
|
||||
if scaling_factor is not None:
|
||||
if isinstance(scaling_factor, torch.Tensor):
|
||||
latents = latents / scaling_factor.to(latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.scaling_factor
|
||||
latents = latents / scaling_factor
|
||||
|
||||
if hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None:
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latents = latents + self.vae.shift_factor.to(latents.device, latents.dtype)
|
||||
shift_factor = getattr(self.vae, "shift_factor", None)
|
||||
if shift_factor is not None:
|
||||
if isinstance(shift_factor, torch.Tensor):
|
||||
latents = latents + shift_factor.to(latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents + self.vae.shift_factor
|
||||
latents = latents + shift_factor
|
||||
|
||||
return latents
|
||||
|
||||
|
||||
@@ -221,6 +221,20 @@ class TextEncodingStage(PipelineStage):
|
||||
encoder_device = torch.device(target_device)
|
||||
moved_for_forward = True
|
||||
|
||||
# An explicit `device=` wins. Otherwise follow the encoder's real
|
||||
# param device. Once it has been moved for the forward that is the
|
||||
# target device, and an HF-passthrough encoder's
|
||||
# _fastvideo_input_device (stamped at load, e.g. "cpu" under
|
||||
# text_encoder_cpu_offload) is stale -- honouring it would feed cpu
|
||||
# token ids to cuda weights. The marker only speaks when nothing
|
||||
# moved and the module has no parameters to speak for it.
|
||||
if device is not None:
|
||||
input_device = torch.device(target_device)
|
||||
elif moved_for_forward:
|
||||
input_device = encoder_device
|
||||
else:
|
||||
input_device = getattr(text_encoder, "_fastvideo_input_device", encoder_device)
|
||||
|
||||
tok_kwargs = dict(encoder_config.tokenizer_kwargs)
|
||||
if max_length is not None:
|
||||
tok_kwargs["max_length"] = max_length
|
||||
@@ -264,7 +278,7 @@ class TextEncodingStage(PipelineStage):
|
||||
# pre-format prompts into message lists upstream and rely on
|
||||
# the inner tokenizer + full tokenizer_kwargs (which include
|
||||
# add_generation_prompt). Preserve that original path exactly.
|
||||
text_inputs = tok.apply_chat_template(processed_texts, **tok_kwargs).to(encoder_device)
|
||||
text_inputs = tok.apply_chat_template(processed_texts, **tok_kwargs).to(input_device)
|
||||
else:
|
||||
# Two-step approach matching Diffusers: format with chat
|
||||
# template first, then tokenize the resulting strings.
|
||||
@@ -275,12 +289,12 @@ class TextEncodingStage(PipelineStage):
|
||||
messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
enable_thinking=False,
|
||||
enable_thinking=encoder_config.chat_template_enable_thinking,
|
||||
)
|
||||
formatted_texts.append(formatted)
|
||||
text_inputs = tokenizer(formatted_texts, **tok_kwargs).to(encoder_device)
|
||||
text_inputs = tokenizer(formatted_texts, **tok_kwargs).to(input_device)
|
||||
else:
|
||||
text_inputs = tok(processed_texts, **tok_kwargs).to(encoder_device)
|
||||
text_inputs = tok(processed_texts, **tok_kwargs).to(input_device)
|
||||
|
||||
input_ids = text_inputs["input_ids"]
|
||||
attention_mask = text_inputs["attention_mask"]
|
||||
|
||||
@@ -2,6 +2,7 @@ import torch
|
||||
import types
|
||||
import pytest
|
||||
|
||||
from fastvideo.configs.models.encoders.clip import CLIPTextArchConfig, CLIPTextConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
@@ -13,6 +14,9 @@ class TensorDict(dict):
|
||||
return TensorDict({k: v.to(device) for k, v in self.items()})
|
||||
|
||||
class FakeTokenizer:
|
||||
def __init__(self):
|
||||
self.last_chat_template_kwargs = None
|
||||
|
||||
def __call__(self, texts, **kwargs):
|
||||
B = len(texts)
|
||||
seq_len = int(kwargs.get("max_length", 4))
|
||||
@@ -21,6 +25,10 @@ class FakeTokenizer:
|
||||
"attention_mask": torch.ones(B, seq_len, dtype=torch.long),
|
||||
})
|
||||
|
||||
def apply_chat_template(self, messages, **kwargs):
|
||||
self.last_chat_template_kwargs = kwargs
|
||||
return "formatted prompt"
|
||||
|
||||
|
||||
class FakeChatTokenizer:
|
||||
def __init__(self):
|
||||
@@ -40,17 +48,22 @@ class FakeChatTokenizer:
|
||||
"attention_mask": torch.ones(B, seq_len, dtype=torch.long),
|
||||
})
|
||||
|
||||
|
||||
class FakeTextEncoder(torch.nn.Module):
|
||||
def __init__(self, hidden_size=8):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.last_input_device = None
|
||||
self.last_output_hidden_states = None
|
||||
|
||||
def forward(self, input_ids, attention_mask, output_hidden_states=False):
|
||||
self.last_input_device = input_ids.device
|
||||
self.last_output_hidden_states = bool(output_hidden_states)
|
||||
B, T = input_ids.shape
|
||||
last_hidden_state = torch.arange(B * T * self.hidden_size, dtype=torch.float32).view(B, T, self.hidden_size)
|
||||
last_hidden_state = torch.arange(
|
||||
B * T * self.hidden_size,
|
||||
dtype=torch.float32,
|
||||
device=input_ids.device,
|
||||
).view(B, T, self.hidden_size)
|
||||
hidden_states = (last_hidden_state, ) if output_hidden_states else None
|
||||
return types.SimpleNamespace(last_hidden_state=last_hidden_state,
|
||||
hidden_states=hidden_states)
|
||||
@@ -213,3 +226,104 @@ def test_chat_list_preprocess_output_is_not_stripped():
|
||||
{"role": "user", "content": "a robotic arm welding a metal structure"},
|
||||
]]
|
||||
assert tokenizer.last_kwargs["return_tensors"] == "pt"
|
||||
|
||||
|
||||
def test_chat_template_thinking_is_keyword_only_and_preserves_subclass_positions():
|
||||
arch = CLIPTextArchConfig()
|
||||
config = CLIPTextConfig(
|
||||
arch,
|
||||
"legacy-prefix",
|
||||
None,
|
||||
None,
|
||||
False,
|
||||
False,
|
||||
7,
|
||||
False,
|
||||
True,
|
||||
False,
|
||||
)
|
||||
|
||||
assert config.num_hidden_layers_override == 7
|
||||
assert config.require_post_norm is False
|
||||
assert config.enable_scale is True
|
||||
assert config.is_causal is False
|
||||
assert config.chat_template_enable_thinking is False
|
||||
|
||||
configured = CLIPTextConfig(chat_template_enable_thinking=True)
|
||||
assert configured.chat_template_enable_thinking is True
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
CLIPTextConfig(
|
||||
arch,
|
||||
"legacy-prefix",
|
||||
None,
|
||||
None,
|
||||
False,
|
||||
False,
|
||||
7,
|
||||
False,
|
||||
True,
|
||||
False,
|
||||
True,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("enable_thinking", [False, True])
|
||||
def test_encode_text_forwards_chat_template_thinking_config(enable_thinking):
|
||||
fastvideo_args, hidden = make_args(num_encoders=1, text_len=4, hidden_size=8)
|
||||
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[0]
|
||||
encoder_config.is_chat_model = True
|
||||
encoder_config.chat_template_enable_thinking = enable_thinking
|
||||
|
||||
stage = make_stage(num_encoders=1, hidden_size=hidden)
|
||||
stage.encode_text("a", fastvideo_args, encoder_index=[0])
|
||||
|
||||
assert stage.tokenizers[0].last_chat_template_kwargs == {
|
||||
"tokenize": False,
|
||||
"add_generation_prompt": True,
|
||||
"enable_thinking": enable_thinking,
|
||||
}
|
||||
|
||||
|
||||
def test_encode_text_uses_hf_passthrough_input_device(monkeypatch):
|
||||
fastvideo_args, hidden = make_args(num_encoders=1, text_len=4, hidden_size=8)
|
||||
stage = make_stage(num_encoders=1, hidden_size=hidden)
|
||||
stage.text_encoders[0]._fastvideo_input_device = torch.device("cpu")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"fastvideo.pipelines.stages.text_encoding.get_local_torch_device",
|
||||
lambda: torch.device("meta"),
|
||||
)
|
||||
|
||||
batch = ForwardBatch(
|
||||
data_type="video",
|
||||
prompt="a",
|
||||
do_classifier_free_guidance=False,
|
||||
prompt_embeds=[],
|
||||
negative_prompt_embeds=None,
|
||||
prompt_attention_mask=[],
|
||||
negative_attention_mask=None,
|
||||
)
|
||||
output = stage.forward(batch, fastvideo_args)
|
||||
|
||||
# The marker governs where the encoder *receives* its tokens...
|
||||
assert stage.text_encoders[0].last_input_device == torch.device("cpu")
|
||||
# ...while the stage still normalizes the embeds it returns onto the
|
||||
# caller's target device.
|
||||
assert output.prompt_embeds[0].device.type == "meta"
|
||||
|
||||
|
||||
def test_encode_text_explicit_device_overrides_hf_passthrough_marker():
|
||||
fastvideo_args, hidden = make_args(num_encoders=1, text_len=4, hidden_size=8)
|
||||
stage = make_stage(num_encoders=1, hidden_size=hidden)
|
||||
stage.text_encoders[0]._fastvideo_input_device = torch.device("cpu")
|
||||
|
||||
output = stage.encode_text(
|
||||
"a",
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
device=torch.device("meta"),
|
||||
)
|
||||
|
||||
assert stage.text_encoders[0].last_input_device == torch.device("meta")
|
||||
assert output[0].device.type == "meta"
|
||||
|
||||
@@ -0,0 +1,120 @@
|
||||
# Z-Image Port Status
|
||||
|
||||
## Summary
|
||||
|
||||
- model_family: `zimage`
|
||||
- variant: `Z-Image-Turbo`
|
||||
- workload_types: `T2I`
|
||||
- PR scope: component-only (`text_encoder`, `tokenizer`, `vae`, `scheduler`)
|
||||
- official_ref: `https://github.com/Tongyi-MAI/Z-Image`
|
||||
- official_ref_dir: `<repo_root>/Z-Image/src`
|
||||
- official_ref_commit: `26f23eda626ffadda020b04ff79488e1d72004cd`
|
||||
- hf_weights_path: `Tongyi-MAI/Z-Image-Turbo@f332072aa78be7aecdf3ee76d5c247082da564a6`
|
||||
- local_weights_dir: `<repo_root>/official_weights/Z-Image/`
|
||||
- source_layout: `diffusers`
|
||||
- local_tests_readme: `tests/local_tests/zimage/README.md`
|
||||
|
||||
The pinned text encoder has 36 decoder blocks, hidden size 2560,
|
||||
intermediate size 9728, 32 attention heads, 8 KV heads, and head dimension
|
||||
128. FastVideo obtains those values from the pinned `config.json`.
|
||||
|
||||
## Current Phase
|
||||
|
||||
- phase: Phase 6 (component parity revalidation)
|
||||
- status: in_progress
|
||||
- owner: parity
|
||||
- last_updated: 2026-07-04
|
||||
|
||||
## Component Matrix
|
||||
|
||||
| Component | Type | Reuse/Port | Official Definition / Instantiation | FastVideo Target | Prototype | Conversion | Parity | Open Issues |
|
||||
|---|---|---|---|---|---|---|---|---|
|
||||
| Text encoder (Qwen3 body) | text_encoder | reused production body + native implementation | Pinned `Qwen3ForCausalLM` checkpoint; Transformers `AutoModel` body; official pipeline consumes `hidden_states[-2]` | `TextEncoderLoader` -> `Qwen3ForCausalLM.from_pretrained_local` -> independent body-only `AutoModel`, plus native `fastvideo.models.encoders.qwen3.Qwen3ForCausalLM` | DONE | not_needed for current HF subfolder | REVALIDATION REQUIRED (historical unpinned native fp32/bf16 PASS only) | I001, I006, I007 |
|
||||
| Tokenizer | tokenizer | reused | `AutoTokenizer` from `<weights>/tokenizer/`; official thinking template and length 512 | `TokenizerLoader` | DONE | not_needed | REVALIDATION REQUIRED (historical unpinned PASS) | I003, I010 |
|
||||
| VAE | vae | reused with accepted wrapper exception | Pinned `zimage.AutoencoderKL`; official decode uses `(latents / scaling_factor) + shift_factor` | Existing Diffusers-backed `fastvideo.models.vaes.autoencoder_kl.AutoencoderKL` through production `VAELoader`, plus direct implementation coverage | DONE | not_needed | REVALIDATION REQUIRED (historical unpinned direct-decode PASS only) | I008 |
|
||||
| Scheduler | scheduler | reused with Z-Image extension | Pinned `zimage.FlowMatchEulerDiscreteScheduler`; official pipeline sets `sigma_min=0.0` | `FlowMatchEulerDiscreteScheduler(use_reference_discrete_timesteps=True, sigma_min=0.0)` | DONE | not_needed | REVALIDATION REQUIRED; asset-free regressions remain separately runnable | I002, I009 |
|
||||
| Transformer (`ZImageTransformer2DModel`) | dit | planned native port | `<weights>/transformer/` and pinned official source | future PR | NOT_STARTED | unknown_until_prototype | NOT_STARTED | I004 |
|
||||
| Pipeline | pipeline | planned | Pinned `ZImagePipeline` | future PR | NOT_STARTED | depends_on_transformer | NOT_STARTED | I005 |
|
||||
|
||||
## Conversion State
|
||||
|
||||
- conversion_script: none for the reused components in this component-only PR
|
||||
- converted_weights_dir: none
|
||||
- source_layout: diffusers
|
||||
- future_transformer_conversion: unknown until native key/shape prototype exists
|
||||
- strict_load_status: revalidation_required
|
||||
- production_passthrough: body-only Qwen3 `AutoModel`, tokenizer assets
|
||||
- direct_load_components: native Qwen3 parity target, VAE, scheduler
|
||||
- native_qwen_exclusion: `lm_head.weight` only; the native encoder owns the body, not an LM head
|
||||
- native_qwen_fused_contract: every Q/K/V and gate/up shard is required unless the exact fused destination is present
|
||||
- native_qwen_auxiliary_contract: mapped quantization auxiliaries, including scale parameters, must load and may not be silently discarded
|
||||
- retry_history: none
|
||||
|
||||
## Parity Commands
|
||||
|
||||
| Scope | Command | Last Result | Notes |
|
||||
|---|---|---|---|
|
||||
| Scheduler | `pytest tests/local_tests/zimage/test_zimage_scheduler_parity.py -v -s` | RERUN REQUIRED; historical unpinned 2/2 PASS on A40, 2026-05-12 | Now enforces reference SHA/import origin and exact zero endpoint; includes asset-free positional/default regressions. |
|
||||
| Tokenizer | `pytest tests/local_tests/zimage/test_zimage_tokenizer_parity.py -v -s` | RERUN REQUIRED; historical unpinned 2/2 PASS on A40, 2026-05-12 | Now uses `enable_thinking=True`, `max_length=512`, and fails if chat templating is unavailable. |
|
||||
| VAE | `pytest tests/local_tests/zimage/test_zimage_vae_parity.py -v -s` | RERUN REQUIRED; historical unpinned 1/1 direct-decode PASS on A40, 2026-05-12 | Scope is now `both`: direct decode plus `VAELoader` and official latent normalization. |
|
||||
| Text encoder fp32 | `pytest tests/local_tests/zimage/test_zimage_encoder_parity.py::test_zimage_qwen3_encoder_parity_forward[fp32] -v -s` | RERUN REQUIRED; historical unpinned native PASS on L40S, 2026-06-21 | New run must cover independent reference, body-only production path, and native implementation. |
|
||||
| Text encoder bf16 | `pytest tests/local_tests/zimage/test_zimage_encoder_parity.py::test_zimage_qwen3_encoder_parity_forward[bf16] -v -s` | RERUN REQUIRED; historical unpinned native PASS on L40S, 2026-06-21 | Historical thresholds retained pending pinned-snapshot validation. |
|
||||
| Per-layer bf16 diagnostic | `pytest tests/local_tests/zimage/test_zimage_encoder_parity.py::test_zimage_qwen3_encoder_per_layer_bf16_diagnostic -v -s` | RERUN REQUIRED; historical unpinned informational run on L40S, 2026-06-21 | Pinned architecture has 36 blocks and 37 hidden-state entries under the tested contract. |
|
||||
|
||||
## Open Questions
|
||||
|
||||
| ID | Question | Owner | Needed By Phase | Status | Resolution |
|
||||
|---|---|---|---|---|---|
|
||||
| Q001 | Which exact Z-Image source revision is the parity oracle? | prep | Phase 1 | resolved | Pinned `Tongyi-MAI/Z-Image@26f23eda626ffadda020b04ff79488e1d72004cd`; scheduler/VAE tests fail on a wrong HEAD or import origin. |
|
||||
| Q002 | Which immutable HF snapshot supplies component weights? | prep | Phase 1 | resolved | `Tongyi-MAI/Z-Image-Turbo@f332072aa78be7aecdf3ee76d5c247082da564a6`. |
|
||||
| Q003 | What prompt/tokenization/hidden-state contract does the official pipeline use? | pipeline | Phase 6 | resolved | `apply_chat_template(tokenize=False, add_generation_prompt=True, enable_thinking=True)`, tokenization length 512, and valid-token `hidden_states[-2]`. |
|
||||
| Q004 | Does the future native transformer require conversion? | conversion | Phase 5 of future transformer PR | open | Decide from official/native key-shape dumps; Diffusers layout alone is not proof. |
|
||||
|
||||
## Issues And Blockers
|
||||
|
||||
| ID | Phase | Component | Severity | Issue | Evidence | Owner | Status | Resolution |
|
||||
|---|---|---|---|---|---|---|---|---|
|
||||
| I001 | Phase 6 | text_encoder | medium | The full checkpoint includes `lm_head.weight`, while the native FastVideo encoder owns only the body. | Native parity allowlists exactly `lm_head.weight`; production uses body-only Transformers `AutoModel`. | parity | in_progress | Implementation is complete; resolve only after the pinned fp32/bf16 rerun. |
|
||||
| I002 | Phase 7 | scheduler/pipeline | high | Published `scheduler_config.json` omits the reference-schedule flag and official zero endpoint. | Pinned pipeline mutates `sigma_min=0.0`; parity requires both scheduler values. | pipeline | open | Future pipeline config must serialize `use_reference_discrete_timesteps=True` and `sigma_min=0.0`. |
|
||||
| I003 | Phase 6 | tokenizer/text_encoder | medium | Prompt formatting and target length previously differed from the official call. | Pinned source uses thinking chat template, 512 tokens, and `hidden_states[-2]`. | parity | in_progress | Test/config behavior corrected; pinned asset rerun required before resolving. |
|
||||
| I004 | Phase 4 | transformer | high | `ZImageTransformer2DModel` is not ported. | No FastVideo target or parity test exists. | port | open | Separate future PR. |
|
||||
| I005 | Phase 7 | pipeline | high | No FastVideo pipeline/config/preset/registry/example or quality regression exists. | Component-only PR scope. | pipeline | open | Separate future PR after component gates pass. |
|
||||
| I006 | Phase 6 | text_encoder | high | Production loading previously used a causal-LM object and could override CPU placement, while parity exercised only direct native construction. | Production path is now body-only `AutoModel`; loader records `_fastvideo_input_device`; stage moves tokens to that device; parity uses a distinct reference instance. | encoder | in_progress | Implementation complete; pinned production-loader rerun required. |
|
||||
| I007 | Phase 6 | text_encoder | high | Lenient native loading could miss one fused source shard or an auxiliary quantization scale. | Loader completeness now covers every destination, required split shard, exact fused destination, and mapped scale parameter. | encoder | in_progress | Asset-free strictness regressions and pinned checkpoint load must pass before resolving. |
|
||||
| I008 | Phase 6 | vae | medium | Direct decode alone did not cover production registry/config/device/strict-load behavior; the shared target is a runtime Diffusers wrapper. | VAE parity scope is now `both` and includes `VAELoader` plus official normalization. | parity | in_progress | Existing Diffusers `AutoencoderKL` wrapper exception is accepted only for this component-only PR; pinned rerun required. |
|
||||
| I009 | Phase 2 | reference assets | high | Adding a nonexistent clone path to `sys.path` could import another installed `zimage` package and yield false parity. | Tests now verify exact clone HEAD and module paths under `Z-Image/src`; `/Z-Image/` is ignored. | prep | resolved | Wrong SHA/path/import is a failure; absent local assets may skip. |
|
||||
| I010 | Phase 6 | tokenizer | medium | Missing `apply_chat_template` previously skipped and could hide an incompatible tokenizer. | Tokenizer parity now asserts the API and exact official kwargs. | parity | resolved | Missing chat-template behavior fails. Asset-backed numerical rerun remains I003. |
|
||||
|
||||
## Escape Hatches
|
||||
|
||||
| ID | Phase | Decision Type | Question | Recommended Option | Status | Resolution |
|
||||
|---|---|---|---|---|---|---|
|
||||
|
||||
## Decisions
|
||||
|
||||
| Date | Decision | Rationale | Impact |
|
||||
|---|---|---|---|
|
||||
| 2026-07-04 | Pin both external sources and treat all earlier GPU results as historical until rerun. | Earlier evidence did not record an immutable HF revision, and source imports were not origin-checked. | README has executable clone/download commands; tests enforce the source SHA; all asset-backed parity rows are revalidation-required. |
|
||||
| 2026-07-04 | Use a body-only Transformers `AutoModel` for production and a separate `AutoModel` object as the parity oracle. | Every FastVideo consumer needs hidden states, not full-vocabulary logits or an LM head; comparing an object with itself is not independent evidence. | Production avoids the LM head while parity covers production loading and native implementation separately. |
|
||||
| 2026-07-04 | Make fused-source and quantization-scale completeness part of native Qwen strict loading. | A loaded destination name alone did not prove that every required Q/K/V or gate/up shard and auxiliary scale arrived. | Missing split shards or mapped scale parameters fail; an exact fused destination remains valid. |
|
||||
| 2026-07-04 | Accept the existing Diffusers-backed AutoencoderKL wrapper only for this component-only PR. | The target is established shared code and this PR adds production-loader evidence rather than a new VAE implementation. | Exception is limited to the VAE here and is not precedent for the transformer or future full pipeline. |
|
||||
| 2026-07-04 | Treat absent chat-template behavior as incompatibility, not an optional test environment. | Thinking-format prompt construction is part of the official numerical contract. | Tokenizer parity fails when `apply_chat_template` is unavailable. |
|
||||
| 2026-07-04 | Record the pinned Qwen architecture as 36 blocks, hidden size 2560, intermediate size 9728, 32 attention heads, and 8 KV heads. | Earlier notes described a different architecture and could mis-size the native target. | Config/review evidence now matches `f332072aa78be7aecdf3ee76d5c247082da564a6`; expected hidden-state length is 37. |
|
||||
| 2026-06-21 | Expand generated Qwen `position_ids` to `[batch_size, seq_len]`. | The rotary layer flattens positions; `[1, seq_len]` misaligned RoPE for batch sizes greater than one. | Batch-one behavior is unchanged; pinned batch-two parity still requires rerun. |
|
||||
| 2026-06-21 | Reuse the shared native `Qwen3ForCausalLM` implementation rather than maintain a second Z-Image encoder. | The shared implementation supports the same config-driven GQA, fused QKV/gate-up, and hidden-state output contract. | Registry/native parity uses the shared class; architecture values come from the pinned config. |
|
||||
| 2026-05-12 | Use mean and median drift checks for deep bf16 encoder parity. | Cross-kernel fused/unfused bf16 arithmetic produces a long max-error tail; fp32 and per-layer diagnostics distinguish this from structural divergence. | Existing thresholds remain provisional until the pinned 36-block snapshot is rerun. |
|
||||
| 2026-05-12 | Forward the complete scheduler config except loader metadata. | Hand-picking three fields could silently omit future scheduler behavior. | Scheduler parity consumes all on-disk constructor fields and explicitly supplies Z-Image runtime overrides. |
|
||||
|
||||
## Handoff Notes
|
||||
|
||||
- No asset-backed component currently has pinned-snapshot `non_skip_pass`
|
||||
evidence. Do not promote a historical row to PASS without recording the exact
|
||||
command and non-skip result against both immutable pins.
|
||||
- Text-encoder revalidation must cover the independent `AutoModel` oracle,
|
||||
body-only production loader, native implementation, requested input device,
|
||||
fused/split source completeness, and auxiliary quantization scales.
|
||||
- VAE revalidation must cover both direct implementation decode and production
|
||||
`VAELoader` behavior with official scaling/shift normalization.
|
||||
- Next separate-PR work: transformer prototype/parity, conversion decision,
|
||||
pipeline/config/preset/registry/example, pipeline parity, and image-quality
|
||||
regression.
|
||||
@@ -0,0 +1,142 @@
|
||||
# Z-Image Local Tests
|
||||
|
||||
Local-only component coverage for the `zimage` FastVideo port (`T2I`,
|
||||
Z-Image-Turbo only). This component-only PR covers the Qwen3 text encoder,
|
||||
tokenizer, AutoencoderKL VAE, and FlowMatchEulerDiscreteScheduler. The native
|
||||
`ZImageTransformer2DModel`, pipeline, pipeline parity, example, and image-quality
|
||||
regression remain out of scope.
|
||||
|
||||
> **Status:** DRAFT port, in progress. Every test that consumes the reference
|
||||
> clone or HF assets requires a fresh non-skip rerun against the immutable pins
|
||||
> below. Earlier A40/L40S results used an unrecorded HF snapshot and are retained
|
||||
> only as historical diagnostics, not current verification evidence. See
|
||||
> [`PORT_STATUS.md`](./PORT_STATUS.md) for the live state.
|
||||
|
||||
## Reference assets and scope
|
||||
|
||||
| Field | Value |
|
||||
|---|---|
|
||||
| Model family / variant | `zimage` / `Z-Image-Turbo` |
|
||||
| Workload | `T2I` |
|
||||
| Component-only PR scope | text encoder, tokenizer, VAE, scheduler |
|
||||
| Official reference | `https://github.com/Tongyi-MAI/Z-Image` |
|
||||
| Local reference dir | `<repo_root>/Z-Image/src` |
|
||||
| Official commit | `26f23eda626ffadda020b04ff79488e1d72004cd` |
|
||||
| HF weights | `Tongyi-MAI/Z-Image-Turbo@f332072aa78be7aecdf3ee76d5c247082da564a6` |
|
||||
| Local weights dir | `<repo_root>/official_weights/Z-Image/` (`text_encoder/`, `tokenizer/`, `vae/`, `scheduler/`) |
|
||||
| Source layout | Diffusers-style per-component subfolders |
|
||||
| Conversion | not needed for the reused components in this PR; the future native transformer/full-pipeline decision will come from native and official key/shape prototypes |
|
||||
| HF token env | `HF_TOKEN` (the pinned repository is public; never record a token value) |
|
||||
|
||||
The pinned text-encoder config declares `Qwen3ForCausalLM` with 36 decoder
|
||||
blocks, hidden size 2560, intermediate size 9728, 32 attention heads, 8 KV
|
||||
heads, and head dimension 128. FastVideo loads those values from
|
||||
`text_encoder/config.json`; they are not hard-coded in the Z-Image integration.
|
||||
|
||||
## Shared environment setup
|
||||
|
||||
Run from the FastVideo repository root in the same environment used for
|
||||
FastVideo. Keep the reference clone and weights inside the ignored paths shown
|
||||
below.
|
||||
|
||||
Clone and pin the official reference:
|
||||
|
||||
```bash
|
||||
python .agents/skills/add-model-01-prep/scripts/clone_reference_repo.py \
|
||||
https://github.com/Tongyi-MAI/Z-Image.git \
|
||||
Z-Image \
|
||||
--commit 26f23eda626ffadda020b04ff79488e1d72004cd \
|
||||
--update-gitignore
|
||||
git -C Z-Image rev-parse HEAD
|
||||
```
|
||||
|
||||
The final command must print
|
||||
`26f23eda626ffadda020b04ff79488e1d72004cd`. The scheduler and VAE tests also
|
||||
enforce that exact HEAD and verify that imported `zimage` modules resolve under
|
||||
`<repo_root>/Z-Image/src`. A missing clone may skip local parity; a wrong SHA or
|
||||
wrong import origin fails rather than silently testing another installation.
|
||||
|
||||
Download the immutable HF snapshot:
|
||||
|
||||
```bash
|
||||
python .agents/skills/add-model-01-prep/scripts/download_hf_weights.py \
|
||||
Tongyi-MAI/Z-Image-Turbo \
|
||||
official_weights/Z-Image \
|
||||
--revision f332072aa78be7aecdf3ee76d5c247082da564a6 \
|
||||
--allow-pattern 'text_encoder/*' \
|
||||
--allow-pattern 'tokenizer/*' \
|
||||
--allow-pattern 'vae/*' \
|
||||
--allow-pattern 'scheduler/*'
|
||||
```
|
||||
|
||||
Do not change core dependency versions (`torch`, `diffusers`, `transformers`,
|
||||
`flash-attn`, `triton`, or CUDA packages) without explicit approval.
|
||||
|
||||
```text
|
||||
dependency_changes: none
|
||||
official_env_status: not_verified
|
||||
private_dep_stubs: none
|
||||
blocked_on: pinned-snapshot non-skip component reruns
|
||||
```
|
||||
|
||||
## Run the tests
|
||||
|
||||
```bash
|
||||
pytest tests/local_tests/zimage/ -v -s
|
||||
```
|
||||
|
||||
| Component | Coverage scope and contract | Status |
|
||||
|---|---|---|
|
||||
| Scheduler ([`test_zimage_scheduler_parity.py`](./test_zimage_scheduler_parity.py)) | `implementation_subcomponent`; exact official `sigma_min=0.0` plus `use_reference_discrete_timesteps=True`; positional/default-path regressions remain asset-free | REVALIDATION REQUIRED |
|
||||
| Tokenizer ([`test_zimage_tokenizer_parity.py`](./test_zimage_tokenizer_parity.py)) | `production_loader`; exact `apply_chat_template(tokenize=False, add_generation_prompt=True, enable_thinking=True)` and `max_length=512`; absence of `apply_chat_template` is a failure, not a skip | REVALIDATION REQUIRED (historical unpinned PASS) |
|
||||
| VAE ([`test_zimage_vae_parity.py`](./test_zimage_vae_parity.py)) | `both`; direct implementation decode plus production `VAELoader`; production check applies `(latents / scaling_factor) + shift_factor` before decode | REVALIDATION REQUIRED (historical unpinned direct-decode PASS only) |
|
||||
| Text encoder ([`test_zimage_encoder_parity.py`](./test_zimage_encoder_parity.py)) | `both`; independent Transformers `AutoModel` reference, body-only production loader, and FastVideo-native implementation; fp32/bf16 output checks plus fused/split/quant-scale strictness | REVALIDATION REQUIRED (historical unpinned native PASS only) |
|
||||
|
||||
## Component contracts
|
||||
|
||||
### Text encoder and tokenizer
|
||||
|
||||
- The production `TextEncoderLoader` path deliberately loads Transformers
|
||||
`AutoModel`, even when the checkpoint advertises `Qwen3ForCausalLM`. FastVideo
|
||||
consumes hidden states only, so the production path is body-only and does not
|
||||
materialize or execute an LM head.
|
||||
- The parity oracle is a separate `AutoModel` instance; it is not the object
|
||||
returned by the production loader. The native `Qwen3ForCausalLM` implementation
|
||||
is compared separately.
|
||||
- CPU/MPS text-encoder offload remains on the requested device. The loader records
|
||||
`_fastvideo_input_device`, and `TextEncodingStage` sends token tensors there.
|
||||
- Native loading must account for every destination parameter and every required
|
||||
fused source shard. Q/K/V and gate/up shards are all required unless the exact
|
||||
destination is supplied already fused; quantized auxiliary parameters such as
|
||||
`scale_weight` are not silently ignored.
|
||||
- The pinned official prompt path enables thinking, tokenizes to 512, and consumes
|
||||
`hidden_states[-2]` at valid-token positions. A 36-block Qwen model exposes 37
|
||||
hidden-state entries in this contract (embedding/intermediate entries plus the
|
||||
final normalized state).
|
||||
- `TextEncoderConfig.chat_template_enable_thinking` is keyword-only and defaults
|
||||
to `False` for existing model families. A future Z-Image pipeline config must
|
||||
opt in with `True`.
|
||||
|
||||
### VAE
|
||||
|
||||
The test covers both direct implementation parity and production `VAELoader`
|
||||
resolution/strict loading. This component-only PR explicitly accepts the
|
||||
existing shared `fastvideo.models.vaes.autoencoder_kl.AutoencoderKL` wrapper,
|
||||
which subclasses Diffusers `AutoencoderKL`, as a narrowly scoped exception to
|
||||
the native-component boundary. This is not precedent for the future transformer
|
||||
or full Z-Image pipeline, and the exception must be revisited if that scope grows.
|
||||
|
||||
### Scheduler
|
||||
|
||||
The pinned official pipeline mutates `scheduler.sigma_min = 0.0` before building
|
||||
its `num_steps + 1` schedule. A future Z-Image pipeline config must set both
|
||||
`use_reference_discrete_timesteps=True` and `sigma_min=0.0`; the published
|
||||
`scheduler_config.json` alone does not encode the complete runtime contract.
|
||||
|
||||
## Remaining work
|
||||
|
||||
- Run every asset-backed component test non-skip against the pinned clone and HF
|
||||
snapshot before claiming component parity.
|
||||
- Port and validate `ZImageTransformer2DModel` in a separate PR.
|
||||
- Add the Z-Image pipeline/config/preset/registry/example, pipeline parity, and
|
||||
image-quality regression after all component gates pass.
|
||||
@@ -0,0 +1 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -0,0 +1,743 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Parity test for Z-Image Qwen3 text encoder support in FastVideo.
|
||||
|
||||
This compares:
|
||||
1) direct transformers AutoModel output, and
|
||||
2) FastVideo's production ``TextEncoderLoader`` passthrough, and
|
||||
3) FastVideo's shared native Qwen3 encoder (``Qwen3ForCausalLM``, added for Flux2
|
||||
Klein and reused here — Z-Image-Turbo's ``Qwen3Model`` checkpoint routes
|
||||
to it via the model registry),
|
||||
|
||||
using identical local checkpoint, tokenization, and inputs.
|
||||
|
||||
Usage:
|
||||
pytest tests/local_tests/zimage/test_zimage_encoder_parity.py -v -s
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
import gc
|
||||
import json
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
|
||||
def _reclaim_vram() -> None:
|
||||
"""Actually reclaim VRAM after the caller has ``del``'d its model refs.
|
||||
|
||||
HF transformer models hold reference cycles, so ``del`` alone does not free
|
||||
them — a ``gc.collect()`` is required before the CUDA caching allocator will
|
||||
release the blocks. Without this the parity test keeps two full encoder
|
||||
copies resident and OOMs on a 44 GB GPU.
|
||||
"""
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def _gpu_mem_mb() -> float:
|
||||
if torch.cuda.is_available():
|
||||
return torch.cuda.memory_allocated() / (1024 ** 2)
|
||||
return 0.0
|
||||
from torch.testing import assert_close
|
||||
from transformers import AutoModel, AutoModelForCausalLM, AutoTokenizer
|
||||
from safetensors.torch import safe_open
|
||||
|
||||
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
|
||||
from fastvideo.distributed.parallel_state import cleanup_dist_env_and_memory, maybe_init_distributed_environment_and_model_parallel
|
||||
from fastvideo.layers.quantization.absmax_fp8 import AbsMaxFP8MergedParameter
|
||||
from fastvideo.models.encoders.qwen3 import Qwen3ForCausalLM
|
||||
from fastvideo.models.loader.component_loader import TextEncoderLoader
|
||||
from fastvideo.models.registry import ModelRegistry
|
||||
|
||||
PARITY_SCOPE = "both"
|
||||
|
||||
# Strict-load contract: the shared Qwen3 encoder is body-only (embed_tokens +
|
||||
# layers + norm, no lm_head). Z-Image-Turbo ships a full Qwen3 checkpoint, so
|
||||
# `lm_head.weight` is the only key that may go unmatched; anything else means a
|
||||
# real silent drop. Enforced here in the test since the shared encoder's loader
|
||||
# is intentionally lenient (it's used by multiple models).
|
||||
_ALLOWED_UNEXPECTED_KEYS = {"lm_head.weight"}
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
ZIMAGE_TEXT_ENCODER_DIR = REPO_ROOT / "official_weights" / "Z-Image" / "text_encoder"
|
||||
ZIMAGE_TOKENIZER_DIR = REPO_ROOT / "official_weights" / "Z-Image" / "tokenizer"
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", autouse=True)
|
||||
def _init_dist_and_tp_groups():
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29531")
|
||||
os.environ.setdefault("RANK", "0")
|
||||
os.environ.setdefault("WORLD_SIZE", "1")
|
||||
os.environ.setdefault("LOCAL_RANK", "0")
|
||||
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
yield
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
def _load_json(path: Path) -> dict:
|
||||
with path.open("r", encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def _iter_safetensors(path: Path):
|
||||
with safe_open(str(path), framework="pt", device="cpu") as f:
|
||||
for k in f.keys():
|
||||
yield k, f.get_tensor(k)
|
||||
|
||||
|
||||
def _iter_pretrained_safetensors(model_dir: Path):
|
||||
single = model_dir / "model.safetensors"
|
||||
if single.exists():
|
||||
yield from _iter_safetensors(single)
|
||||
return
|
||||
|
||||
index = model_dir / "model.safetensors.index.json"
|
||||
if index.exists():
|
||||
idx = _load_json(index)
|
||||
shard_names = sorted(set(idx["weight_map"].values()))
|
||||
for shard in shard_names:
|
||||
yield from _iter_safetensors(model_dir / shard)
|
||||
return
|
||||
|
||||
raise FileNotFoundError(
|
||||
f"Missing safetensors checkpoint in {model_dir} (expected model.safetensors or model.safetensors.index.json)"
|
||||
)
|
||||
|
||||
|
||||
def _pretrained_safetensor_keys(model_dir: Path) -> set[str]:
|
||||
return {name for name, _ in _iter_pretrained_safetensors(model_dir)}
|
||||
|
||||
|
||||
def _load_qwen3_config() -> Qwen3TextConfig:
|
||||
cfg_raw = _load_json(ZIMAGE_TEXT_ENCODER_DIR / "config.json")
|
||||
for key in ("_name_or_path", "transformers_version", "model_type", "torch_dtype"):
|
||||
cfg_raw.pop(key, None)
|
||||
config = Qwen3TextConfig()
|
||||
config.update_model_arch(cfg_raw)
|
||||
assert config.num_hidden_layers == 36
|
||||
assert config.hidden_size == 2560
|
||||
assert config.intermediate_size == 9728
|
||||
assert config.num_attention_heads == 32
|
||||
assert config.num_key_value_heads == 8
|
||||
return config
|
||||
|
||||
|
||||
def _loader_args(cpu_offload: bool) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
text_encoder_cpu_offload=cpu_offload,
|
||||
override_text_encoder_quant=None,
|
||||
override_text_encoder_safetensors=None,
|
||||
pin_cpu_memory=False,
|
||||
)
|
||||
|
||||
|
||||
def _precision_name(dtype: torch.dtype) -> str:
|
||||
if dtype == torch.float32:
|
||||
return "fp32"
|
||||
if dtype == torch.bfloat16:
|
||||
return "bf16"
|
||||
raise ValueError(f"Unsupported parity dtype: {dtype}")
|
||||
|
||||
|
||||
def _tokenize_official_prompts(tokenizer, prompts: list[str]):
|
||||
formatted = [
|
||||
tokenizer.apply_chat_template(
|
||||
[{"role": "user", "content": prompt}],
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
enable_thinking=True,
|
||||
)
|
||||
for prompt in prompts
|
||||
]
|
||||
return tokenizer(
|
||||
formatted,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
max_length=512,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
|
||||
def _assert_native_load_surface(
|
||||
model,
|
||||
checkpoint_keys: set[str],
|
||||
loaded_params: set[str],
|
||||
) -> None:
|
||||
param_names = {name for name, _ in model.named_parameters()}
|
||||
missing = param_names - loaded_params
|
||||
assert not missing, f"FastVideo parameters not loaded: {sorted(missing)}"
|
||||
|
||||
normalized_keys = {
|
||||
name[len("model."):] if name.startswith("model.") else name
|
||||
for name in checkpoint_keys
|
||||
}
|
||||
mapped_keys: set[str] = set()
|
||||
for checkpoint_name in normalized_keys:
|
||||
for param_name, weight_name, _ in model.config.arch_config.stacked_params_mapping:
|
||||
if weight_name in checkpoint_name:
|
||||
mapped_keys.add(checkpoint_name.replace(weight_name, param_name))
|
||||
break
|
||||
else:
|
||||
mapped_keys.add(checkpoint_name)
|
||||
|
||||
unexpected = mapped_keys - param_names - {
|
||||
"rotary_emb.inv_freq",
|
||||
"rotary_emb.cos_cached",
|
||||
"rotary_emb.sin_cached",
|
||||
}
|
||||
assert unexpected <= _ALLOWED_UNEXPECTED_KEYS, (
|
||||
"Unexpected checkpoint keys not in allowlist: "
|
||||
f"{sorted(unexpected - _ALLOWED_UNEXPECTED_KEYS)}"
|
||||
)
|
||||
|
||||
|
||||
def _fake_qwen3_weight_target(include_quant_scales: bool = False):
|
||||
stacked = Qwen3TextConfig().arch_config.stacked_params_mapping
|
||||
params = {
|
||||
"layers.0.self_attn.qkv_proj.weight": torch.nn.Parameter(torch.empty(1)),
|
||||
"layers.0.mlp.gate_up_proj.weight": torch.nn.Parameter(torch.empty(1)),
|
||||
}
|
||||
if include_quant_scales:
|
||||
qkv_scale = AbsMaxFP8MergedParameter(torch.zeros(3), requires_grad=False)
|
||||
qkv_scale.output_partition_sizes = [1, 1, 1]
|
||||
gate_up_scale = AbsMaxFP8MergedParameter(torch.zeros(2), requires_grad=False)
|
||||
gate_up_scale.output_partition_sizes = [1, 1]
|
||||
params.update({
|
||||
"layers.0.self_attn.qkv_proj.scale_weight": qkv_scale,
|
||||
"layers.0.mlp.gate_up_proj.scale_weight": gate_up_scale,
|
||||
})
|
||||
for param in params.values():
|
||||
if not hasattr(param, "weight_loader"):
|
||||
param.weight_loader = (
|
||||
lambda target, loaded, *args: target.data.copy_(loaded)
|
||||
)
|
||||
|
||||
model = SimpleNamespace(
|
||||
config=SimpleNamespace(
|
||||
arch_config=SimpleNamespace(stacked_params_mapping=stacked)
|
||||
),
|
||||
named_parameters=lambda: params.items(),
|
||||
)
|
||||
return model, params, stacked
|
||||
|
||||
|
||||
# The provisional bf16 thresholds below come from the historical unpinned
|
||||
# 35-layer run. Per add-model-02-parity's calibration block, bf16 max-based
|
||||
# asserts are meaningless for deep encoders — fused-vs-unfused linear order
|
||||
# alone produces a long tail, so the test uses mean + median instead. The
|
||||
# pinned snapshot has 36 blocks and must be rerun before these thresholds are
|
||||
# treated as current evidence.
|
||||
#
|
||||
# Historical empirical numbers (unrecorded snapshot, 35 layers, CUDA):
|
||||
# last_hidden_state (post-norm): mean=0.0152, median=0.0117
|
||||
# hidden_states[-2] (pre-norm): mean=0.0754, median=0.0625
|
||||
# Thresholds below add ~1.6x headroom over observed.
|
||||
_BF16_MEAN_DRIFT_POST_NORM = 0.025
|
||||
_BF16_MEAN_DRIFT_PRE_NORM = 0.120
|
||||
_BF16_MEDIAN_DRIFT_POST_NORM = 0.020
|
||||
_BF16_MEDIAN_DRIFT_PRE_NORM = 0.100
|
||||
_FP32_ATOL = 1e-4
|
||||
_FP32_RTOL = 1e-4
|
||||
|
||||
|
||||
def _print_diag(label: str, ref: torch.Tensor, fv: torch.Tensor) -> torch.Tensor:
|
||||
diff = (ref.float() - fv.float()).abs()
|
||||
flat = diff.flatten()
|
||||
p99 = flat.kthvalue(max(1, int(0.99 * flat.numel()))).values
|
||||
print(
|
||||
f"[{label}] max_diff={diff.max():.4f} mean_diff={diff.mean():.4f} "
|
||||
f"median_diff={diff.median():.4f} p99_diff={p99:.4f}",
|
||||
flush=True,
|
||||
)
|
||||
return diff
|
||||
|
||||
|
||||
def test_qwen3_production_loader_uses_body_only_model_and_honors_cpu_offload(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
"""The production passthrough must avoid the LM head and preserve CPU placement."""
|
||||
import fastvideo.platforms as platforms
|
||||
|
||||
placements: list[torch.device] = []
|
||||
auto_model_calls: list[dict] = []
|
||||
|
||||
class FakeHFModel:
|
||||
def eval(self):
|
||||
return self
|
||||
|
||||
def to(self, device: torch.device):
|
||||
placements.append(torch.device(device))
|
||||
return self
|
||||
|
||||
fake_model = FakeHFModel()
|
||||
|
||||
def fake_auto_model_from_pretrained(*args, **kwargs):
|
||||
auto_model_calls.append(kwargs)
|
||||
return fake_model
|
||||
|
||||
def fail_causal_lm_from_pretrained(*args, **kwargs):
|
||||
raise AssertionError("Qwen text encoding must not load AutoModelForCausalLM")
|
||||
|
||||
monkeypatch.setattr(
|
||||
AutoModel,
|
||||
"from_pretrained",
|
||||
staticmethod(fake_auto_model_from_pretrained),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
AutoModelForCausalLM,
|
||||
"from_pretrained",
|
||||
staticmethod(fail_causal_lm_from_pretrained),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
platforms,
|
||||
"_current_platform",
|
||||
SimpleNamespace(
|
||||
is_mps=lambda: False,
|
||||
verify_model_arch=lambda _arch: None,
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
|
||||
monkeypatch.setattr(
|
||||
"fastvideo.distributed.get_local_torch_device",
|
||||
lambda: torch.device("cuda"),
|
||||
)
|
||||
|
||||
config = Qwen3TextConfig()
|
||||
# The pinned checkpoint names the causal wrapper even though the official
|
||||
# Z-Image loader intentionally asks Transformers for the body-only model.
|
||||
config.arch_config.architectures = ["Qwen3ForCausalLM"]
|
||||
loaded = TextEncoderLoader().load_model(
|
||||
"unused-by-mock",
|
||||
config,
|
||||
torch.device("cuda"),
|
||||
_loader_args(cpu_offload=True),
|
||||
dtype="fp32",
|
||||
use_text_encoder_override=True,
|
||||
)
|
||||
|
||||
assert loaded is fake_model
|
||||
assert auto_model_calls == [{
|
||||
"local_files_only": True,
|
||||
"torch_dtype": torch.float32,
|
||||
"low_cpu_mem_usage": True,
|
||||
}]
|
||||
assert placements == [torch.device("cpu")]
|
||||
assert loaded._fastvideo_input_device == torch.device("cpu")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpoint_prefix", ["", "model."])
|
||||
def test_qwen3_native_load_accepts_already_fused_weights_and_scales(
|
||||
checkpoint_prefix: str,
|
||||
):
|
||||
"""Direct FastVideo/state-dict keys satisfy each fused target atomically."""
|
||||
fake_model, params, _ = _fake_qwen3_weight_target(include_quant_scales=True)
|
||||
expected = {
|
||||
name: torch.arange(1, param.numel() + 1, dtype=param.dtype)
|
||||
for name, param in params.items()
|
||||
}
|
||||
weights = [
|
||||
(f"{checkpoint_prefix}{name}", expected[name])
|
||||
for name in params
|
||||
]
|
||||
|
||||
loaded = Qwen3ForCausalLM.load_weights(fake_model, weights)
|
||||
|
||||
assert loaded == set(params)
|
||||
for name, param in params.items():
|
||||
assert torch.equal(param.data, expected[name])
|
||||
_assert_native_load_surface(
|
||||
fake_model,
|
||||
{name for name, _ in weights},
|
||||
loaded,
|
||||
)
|
||||
|
||||
|
||||
def test_qwen3_native_load_rejects_incompatible_fused_quant_scale_shape():
|
||||
fake_model, params, _ = _fake_qwen3_weight_target(include_quant_scales=True)
|
||||
weights = [
|
||||
(name, torch.ones_like(param))
|
||||
for name, param in params.items()
|
||||
if name != "layers.0.self_attn.qkv_proj.scale_weight"
|
||||
]
|
||||
weights.append(("layers.0.self_attn.qkv_proj.scale_weight", torch.ones(1)))
|
||||
|
||||
with pytest.raises(AssertionError, match="Attempted to load weight"):
|
||||
Qwen3ForCausalLM.load_weights(fake_model, weights)
|
||||
|
||||
|
||||
def test_qwen3_native_load_uses_real_shard_loader_for_split_quant_scales():
|
||||
fake_model, params, stacked = _fake_qwen3_weight_target(include_quant_scales=True)
|
||||
weights = [
|
||||
(name, torch.ones_like(param))
|
||||
for name, param in params.items()
|
||||
if name.endswith(".weight")
|
||||
]
|
||||
weights.extend(
|
||||
(f"model.layers.0.self_attn{weight_name}.scale_weight", torch.tensor(float(index + 1)))
|
||||
for index, (_, weight_name, _) in enumerate(stacked[:3])
|
||||
)
|
||||
weights.extend(
|
||||
(f"model.layers.0.mlp{weight_name}.scale_weight", torch.tensor(float(index + 4)))
|
||||
for index, (_, weight_name, _) in enumerate(stacked[3:])
|
||||
)
|
||||
|
||||
loaded = Qwen3ForCausalLM.load_weights(fake_model, weights)
|
||||
|
||||
assert loaded == set(params)
|
||||
assert torch.equal(
|
||||
params["layers.0.self_attn.qkv_proj.scale_weight"].data,
|
||||
torch.tensor([1.0, 2.0, 3.0]),
|
||||
)
|
||||
assert torch.equal(
|
||||
params["layers.0.mlp.gate_up_proj.scale_weight"].data,
|
||||
torch.tensor([4.0, 5.0]),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("missing_weight_name", "missing_shard_id"),
|
||||
[
|
||||
(".q_proj", "q"),
|
||||
(".k_proj", "k"),
|
||||
(".v_proj", "v"),
|
||||
(".gate_proj", 0),
|
||||
(".up_proj", 1),
|
||||
],
|
||||
)
|
||||
def test_qwen3_native_load_rejects_missing_fused_weight_shard(
|
||||
missing_weight_name: str,
|
||||
missing_shard_id: str | int,
|
||||
):
|
||||
fake_model, _, stacked = _fake_qwen3_weight_target()
|
||||
weights = [
|
||||
(f"model.layers.0.self_attn{weight_name}.weight", torch.empty(1))
|
||||
for _, weight_name, _ in stacked[:3]
|
||||
if weight_name != missing_weight_name
|
||||
]
|
||||
weights.extend(
|
||||
(f"model.layers.0.mlp{weight_name}.weight", torch.empty(1))
|
||||
for _, weight_name, _ in stacked[3:]
|
||||
if weight_name != missing_weight_name
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Missing required stacked checkpoint shards") as exc_info:
|
||||
Qwen3ForCausalLM.load_weights(fake_model, weights)
|
||||
|
||||
destination = (
|
||||
"layers.0.self_attn.qkv_proj.weight"
|
||||
if isinstance(missing_shard_id, str)
|
||||
else "layers.0.mlp.gate_up_proj.weight"
|
||||
)
|
||||
assert f"{destination}[{missing_shard_id}]" in str(exc_info.value)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("missing_weight_name", "missing_shard_id"),
|
||||
[
|
||||
(".q_proj", "q"),
|
||||
(".k_proj", "k"),
|
||||
(".v_proj", "v"),
|
||||
(".gate_proj", 0),
|
||||
(".up_proj", 1),
|
||||
],
|
||||
)
|
||||
def test_qwen3_native_load_rejects_missing_fused_quant_scale_shard(
|
||||
missing_weight_name: str,
|
||||
missing_shard_id: str | int,
|
||||
):
|
||||
fake_model, params, stacked = _fake_qwen3_weight_target(include_quant_scales=True)
|
||||
weights = [
|
||||
(name, torch.empty(1))
|
||||
for name in params
|
||||
if name.endswith(".weight")
|
||||
]
|
||||
weights.extend(
|
||||
(f"model.layers.0.self_attn{weight_name}.scale_weight", torch.empty(1))
|
||||
for _, weight_name, _ in stacked[:3]
|
||||
if weight_name != missing_weight_name
|
||||
)
|
||||
weights.extend(
|
||||
(f"model.layers.0.mlp{weight_name}.scale_weight", torch.empty(1))
|
||||
for _, weight_name, _ in stacked[3:]
|
||||
if weight_name != missing_weight_name
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Missing required stacked checkpoint shards") as exc_info:
|
||||
Qwen3ForCausalLM.load_weights(fake_model, weights)
|
||||
|
||||
destination = (
|
||||
"layers.0.self_attn.qkv_proj.scale_weight"
|
||||
if isinstance(missing_shard_id, str)
|
||||
else "layers.0.mlp.gate_up_proj.scale_weight"
|
||||
)
|
||||
assert f"{destination}[{missing_shard_id}]" in str(exc_info.value)
|
||||
|
||||
|
||||
@pytest.mark.skipif(not ZIMAGE_TEXT_ENCODER_DIR.exists(), reason="Z-Image text encoder checkpoint required")
|
||||
@pytest.mark.parametrize(
|
||||
"dtype",
|
||||
[
|
||||
pytest.param(torch.float32, id="fp32"),
|
||||
pytest.param(
|
||||
torch.bfloat16,
|
||||
id="bf16",
|
||||
marks=pytest.mark.skipif(
|
||||
not (torch.cuda.is_available() and torch.cuda.is_bf16_supported()),
|
||||
reason="bf16 parity requires a bf16-capable CUDA device",
|
||||
),
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_zimage_qwen3_encoder_parity_forward(dtype: torch.dtype):
|
||||
if not ZIMAGE_TOKENIZER_DIR.exists():
|
||||
pytest.skip(f"Z-Image tokenizer dir not found: {ZIMAGE_TOKENIZER_DIR}")
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
torch.manual_seed(11)
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(str(ZIMAGE_TOKENIZER_DIR), local_files_only=True)
|
||||
|
||||
prompts = [
|
||||
"A cinematic shot of a rainy neon street at night.",
|
||||
"A watercolor illustration of a fox in autumn leaves.",
|
||||
]
|
||||
toks = _tokenize_official_prompts(tokenizer, prompts)
|
||||
input_ids = toks["input_ids"].to(device)
|
||||
attention_mask = toks["attention_mask"].to(device)
|
||||
config = _load_qwen3_config()
|
||||
|
||||
ref = AutoModel.from_pretrained(
|
||||
str(ZIMAGE_TEXT_ENCODER_DIR),
|
||||
local_files_only=True,
|
||||
trust_remote_code=True,
|
||||
dtype=dtype,
|
||||
low_cpu_mem_usage=True,
|
||||
).eval().to(device=device, dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_out = ref(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
return_dict=True,
|
||||
)
|
||||
ref_last = ref_out.last_hidden_state.detach().float().cpu()
|
||||
assert ref_out.hidden_states is not None
|
||||
ref_hs_m2 = ref_out.hidden_states[-2].detach().float().cpu()
|
||||
|
||||
print(f"[mem] after HF ref forward: {_gpu_mem_mb():.0f} MiB", flush=True)
|
||||
# Keep only CPU outputs so each full encoder is resident one at a time.
|
||||
del ref, ref_out
|
||||
_reclaim_vram()
|
||||
print(f"[mem] after freeing HF ref: {_gpu_mem_mb():.0f} MiB", flush=True)
|
||||
|
||||
production = TextEncoderLoader().load_model(
|
||||
str(ZIMAGE_TEXT_ENCODER_DIR),
|
||||
config,
|
||||
device,
|
||||
_loader_args(cpu_offload=False),
|
||||
dtype=_precision_name(dtype),
|
||||
use_text_encoder_override=True,
|
||||
)
|
||||
assert production.__class__.__name__ == "Qwen3Model"
|
||||
assert not hasattr(production, "lm_head")
|
||||
assert production._fastvideo_input_device == device
|
||||
|
||||
with torch.no_grad():
|
||||
production_out = production(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
return_dict=True,
|
||||
)
|
||||
production_last = production_out.last_hidden_state.detach().float().cpu()
|
||||
assert production_out.hidden_states is not None
|
||||
production_hs_m2 = production_out.hidden_states[-2].detach().float().cpu()
|
||||
|
||||
# Production loading must preserve the independent Transformers oracle
|
||||
# exactly; native-kernel drift is assessed separately below.
|
||||
assert_close(ref_last, production_last, atol=0.0, rtol=0.0)
|
||||
assert_close(ref_hs_m2, production_hs_m2, atol=0.0, rtol=0.0)
|
||||
del production, production_out, production_last, production_hs_m2
|
||||
_reclaim_vram()
|
||||
|
||||
fv_cls, _ = ModelRegistry.resolve_model_cls("Qwen3Model")
|
||||
fv = fv_cls(config).eval()
|
||||
loaded = fv.load_weights(_iter_pretrained_safetensors(ZIMAGE_TEXT_ENCODER_DIR))
|
||||
assert loaded, "No Qwen3 weights were loaded into FastVideo model"
|
||||
fv = fv.to(device=device, dtype=dtype)
|
||||
print(f"[mem] after FastVideo load: {_gpu_mem_mb():.0f} MiB", flush=True)
|
||||
|
||||
_assert_native_load_surface(
|
||||
fv,
|
||||
_pretrained_safetensor_keys(ZIMAGE_TEXT_ENCODER_DIR),
|
||||
loaded,
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
fv_out = fv(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
fv_last = fv_out.last_hidden_state.detach().float().cpu()
|
||||
assert fv_out.hidden_states is not None
|
||||
fv_hs_m2 = fv_out.hidden_states[-2].detach().float().cpu()
|
||||
|
||||
# Z-Image uses only valid token embeddings selected by attention mask.
|
||||
# Compare parity on that exact path to avoid padded-token artifacts.
|
||||
mask_cpu = attention_mask.detach().bool().cpu()
|
||||
is_bf16 = dtype == torch.bfloat16
|
||||
|
||||
for i in range(mask_cpu.shape[0]):
|
||||
valid = mask_cpu[i]
|
||||
ref_last_v = ref_last[i][valid]
|
||||
fv_last_v = fv_last[i][valid]
|
||||
ref_hs_m2_v = ref_hs_m2[i][valid]
|
||||
fv_hs_m2_v = fv_hs_m2[i][valid]
|
||||
|
||||
last_diff = _print_diag(f"batch{i} last_hidden_state {dtype}", ref_last_v, fv_last_v)
|
||||
hs_m2_diff = _print_diag(f"batch{i} hidden_states[-2] {dtype}", ref_hs_m2_v, fv_hs_m2_v)
|
||||
|
||||
if is_bf16:
|
||||
# Distribution checks instead of element-wise assert_close: cross-
|
||||
# kernel bf16 produces a long max tail (fused vs unfused linears),
|
||||
# but a real op bug pushes mean *and* median high. Mean catches
|
||||
# systematic drift; median catches "wrong weights / swapped
|
||||
# layers" where most elements diverge. Pre-norm tensors get a
|
||||
# looser bound because the final RMSNorm compresses drift ~5x.
|
||||
assert last_diff.mean().item() < _BF16_MEAN_DRIFT_POST_NORM, (
|
||||
f"batch{i} last_hidden_state bf16 mean drift "
|
||||
f"{last_diff.mean():.4f} >= {_BF16_MEAN_DRIFT_POST_NORM}"
|
||||
)
|
||||
assert last_diff.median().item() < _BF16_MEDIAN_DRIFT_POST_NORM, (
|
||||
f"batch{i} last_hidden_state bf16 median drift "
|
||||
f"{last_diff.median():.4f} >= {_BF16_MEDIAN_DRIFT_POST_NORM}"
|
||||
)
|
||||
assert hs_m2_diff.mean().item() < _BF16_MEAN_DRIFT_PRE_NORM, (
|
||||
f"batch{i} hidden_states[-2] bf16 mean drift "
|
||||
f"{hs_m2_diff.mean():.4f} >= {_BF16_MEAN_DRIFT_PRE_NORM}"
|
||||
)
|
||||
assert hs_m2_diff.median().item() < _BF16_MEDIAN_DRIFT_PRE_NORM, (
|
||||
f"batch{i} hidden_states[-2] bf16 median drift "
|
||||
f"{hs_m2_diff.median():.4f} >= {_BF16_MEDIAN_DRIFT_PRE_NORM}"
|
||||
)
|
||||
else:
|
||||
assert_close(ref_last_v, fv_last_v, atol=_FP32_ATOL, rtol=_FP32_RTOL)
|
||||
assert_close(ref_hs_m2_v, fv_hs_m2_v, atol=_FP32_ATOL, rtol=_FP32_RTOL)
|
||||
|
||||
|
||||
@pytest.mark.skipif(not ZIMAGE_TEXT_ENCODER_DIR.exists(), reason="Z-Image text encoder checkpoint required")
|
||||
@pytest.mark.skipif(
|
||||
not (torch.cuda.is_available() and torch.cuda.is_bf16_supported()),
|
||||
reason="bf16 per-layer diagnostic requires a bf16-capable CUDA device",
|
||||
)
|
||||
def test_zimage_qwen3_encoder_per_layer_bf16_diagnostic():
|
||||
"""Diagnostic-only: prints per-layer FastVideo-vs-HF drift for the bf16
|
||||
forward to distinguish monotonic accumulation (bf16 tail) from a
|
||||
single-layer spike (real op mismatch).
|
||||
|
||||
This test always passes; review the captured stdout to interpret.
|
||||
Expected pattern for healthy bf16 accumulation: max/mean/median grow
|
||||
smoothly with layer depth. Suspect pattern: one layer's diff is
|
||||
>>2x the previous, suggesting a real bug at that block.
|
||||
"""
|
||||
if not ZIMAGE_TOKENIZER_DIR.exists():
|
||||
pytest.skip(f"Z-Image tokenizer dir not found: {ZIMAGE_TOKENIZER_DIR}")
|
||||
|
||||
device = torch.device("cuda")
|
||||
dtype = torch.bfloat16
|
||||
torch.manual_seed(11)
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(str(ZIMAGE_TOKENIZER_DIR), local_files_only=True)
|
||||
toks = _tokenize_official_prompts(
|
||||
tokenizer,
|
||||
["A cinematic shot of a rainy neon street at night."],
|
||||
)
|
||||
input_ids = toks["input_ids"].to(device)
|
||||
attention_mask = toks["attention_mask"].to(device)
|
||||
|
||||
ref = AutoModel.from_pretrained(
|
||||
str(ZIMAGE_TEXT_ENCODER_DIR),
|
||||
local_files_only=True,
|
||||
trust_remote_code=True,
|
||||
dtype=dtype,
|
||||
low_cpu_mem_usage=True,
|
||||
).eval().to(device=device, dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_out = ref(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
return_dict=True,
|
||||
)
|
||||
# Cache the reference hidden states on CPU, then free the HF model so
|
||||
# only one encoder is resident when FastVideo runs (avoids OOM).
|
||||
ref_hidden_cpu = [h.detach().float().cpu() for h in ref_out.hidden_states]
|
||||
del ref, ref_out
|
||||
_reclaim_vram()
|
||||
|
||||
fv_cls, _ = ModelRegistry.resolve_model_cls("Qwen3Model")
|
||||
cfg_raw = _load_json(ZIMAGE_TEXT_ENCODER_DIR / "config.json")
|
||||
for k in ("_name_or_path", "transformers_version", "model_type", "torch_dtype"):
|
||||
cfg_raw.pop(k, None)
|
||||
cfg = Qwen3TextConfig()
|
||||
cfg.update_model_arch(cfg_raw)
|
||||
fv = fv_cls(cfg).eval()
|
||||
fv.load_weights(_iter_pretrained_safetensors(ZIMAGE_TEXT_ENCODER_DIR))
|
||||
fv = fv.to(device=device, dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
fv_out = fv(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
fv_hidden_cpu = [h.detach().float().cpu() for h in fv_out.hidden_states]
|
||||
|
||||
assert fv_out.hidden_states is not None
|
||||
assert len(ref_hidden_cpu) == 37, (
|
||||
f"expected embeddings plus 36 block outputs, got {len(ref_hidden_cpu)} entries"
|
||||
)
|
||||
assert len(ref_hidden_cpu) == len(fv_hidden_cpu), (
|
||||
f"len mismatch: HF={len(ref_hidden_cpu)} FV={len(fv_hidden_cpu)}"
|
||||
)
|
||||
|
||||
mask_cpu = attention_mask.detach().bool().cpu()[0]
|
||||
print(
|
||||
f"\n[per-layer bf16 diagnostic] num_hidden_states={len(ref_hidden_cpu)} "
|
||||
f"valid_tokens={mask_cpu.sum().item()}/{mask_cpu.numel()}",
|
||||
flush=True,
|
||||
)
|
||||
print(
|
||||
f"{'idx':>3} {'kind':<10} {'max':>8} {'mean':>8} {'median':>8} {'p99':>8}",
|
||||
flush=True,
|
||||
)
|
||||
for idx, (ref_h, fv_h) in enumerate(zip(ref_hidden_cpu, fv_hidden_cpu)):
|
||||
ref_v = ref_h[0][mask_cpu]
|
||||
fv_v = fv_h[0][mask_cpu]
|
||||
diff = (ref_v - fv_v).abs()
|
||||
flat = diff.flatten()
|
||||
p99 = torch.quantile(flat, 0.99).item()
|
||||
if idx == 0:
|
||||
kind = "embedding"
|
||||
elif idx == len(ref_hidden_cpu) - 1:
|
||||
kind = "post-norm"
|
||||
else:
|
||||
kind = f"layer{idx - 1}-out"
|
||||
print(
|
||||
f"{idx:>3} {kind:<10} "
|
||||
f"{diff.max().item():>8.4f} {diff.mean().item():>8.4f} "
|
||||
f"{diff.median().item():>8.4f} {p99:>8.4f}",
|
||||
flush=True,
|
||||
)
|
||||
@@ -0,0 +1,227 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Parity tests for Z-Image FlowMatchEulerDiscreteScheduler.
|
||||
|
||||
These tests compare FastVideo's scheduler implementation against the
|
||||
reference scheduler shipped in the local Z-Image repository.
|
||||
|
||||
Usage:
|
||||
pytest tests/local_tests/zimage/test_zimage_scheduler_parity.py -v
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import json
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler as FastVideoScheduler,
|
||||
)
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
ZIMAGE_REPO = REPO_ROOT / "Z-Image"
|
||||
ZIMAGE_SRC = REPO_ROOT / "Z-Image" / "src"
|
||||
ZIMAGE_SCHEDULER_CFG = (
|
||||
REPO_ROOT / "official_weights" / "Z-Image" / "scheduler" / "scheduler_config.json"
|
||||
)
|
||||
ZIMAGE_REFERENCE_REVISION = "26f23eda626ffadda020b04ff79488e1d72004cd"
|
||||
PARITY_SCOPE = "implementation_subcomponent"
|
||||
|
||||
|
||||
def _require_pinned_reference_module(module_name: str, source_file: Path):
|
||||
if not ZIMAGE_REPO.exists():
|
||||
pytest.skip(f"Pinned Z-Image reference clone not found: {ZIMAGE_REPO}")
|
||||
if not source_file.is_file():
|
||||
pytest.fail(f"Z-Image reference clone is incomplete; missing {source_file}")
|
||||
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["git", "-C", str(ZIMAGE_REPO), "rev-parse", "HEAD"],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
except (FileNotFoundError, subprocess.CalledProcessError) as exc:
|
||||
pytest.fail(f"Cannot verify Z-Image reference revision: {exc}")
|
||||
|
||||
actual_revision = result.stdout.strip()
|
||||
assert actual_revision == ZIMAGE_REFERENCE_REVISION, (
|
||||
"Z-Image reference clone is not at the pinned revision: "
|
||||
f"expected {ZIMAGE_REFERENCE_REVISION}, got {actual_revision}"
|
||||
)
|
||||
|
||||
if str(ZIMAGE_SRC) not in sys.path:
|
||||
sys.path.insert(0, str(ZIMAGE_SRC))
|
||||
try:
|
||||
module = importlib.import_module(module_name)
|
||||
except Exception as exc:
|
||||
pytest.fail(f"Cannot import pinned Z-Image module {module_name}: {exc}")
|
||||
|
||||
module_file = Path(module.__file__ or "").resolve()
|
||||
assert module_file.is_relative_to(ZIMAGE_SRC.resolve()), (
|
||||
f"{module_name} resolved outside the pinned clone: {module_file}"
|
||||
)
|
||||
return module
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def reference_scheduler_cls():
|
||||
module = _require_pinned_reference_module(
|
||||
"zimage.scheduler",
|
||||
ZIMAGE_SRC / "zimage" / "scheduler.py",
|
||||
)
|
||||
return module.FlowMatchEulerDiscreteScheduler
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def scheduler_config() -> dict:
|
||||
if not ZIMAGE_SCHEDULER_CFG.exists():
|
||||
pytest.skip(f"Z-Image scheduler config not found: {ZIMAGE_SCHEDULER_CFG}")
|
||||
with ZIMAGE_SCHEDULER_CFG.open("r", encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
# Diffusers-style scheduler-config keys we do NOT want to forward — they are
|
||||
# loader/state metadata, not constructor kwargs. Anything else in
|
||||
# scheduler_config.json (including future keys like `time_shift_type`,
|
||||
# `invert_sigmas`, etc.) must be forwarded so parity actually exercises the
|
||||
# config-on-disk and not a hand-curated subset.
|
||||
_SCHEDULER_CONFIG_LOADER_KEYS = frozenset({"_class_name", "_diffusers_version", "_name_or_path"})
|
||||
|
||||
|
||||
def _scheduler_kwargs_from_config(scheduler_config: dict) -> dict:
|
||||
return {k: v for k, v in scheduler_config.items() if k not in _SCHEDULER_CONFIG_LOADER_KEYS}
|
||||
|
||||
|
||||
def _make_scheduler_pair(reference_scheduler_cls, scheduler_config: dict, **overrides):
|
||||
kwargs = _scheduler_kwargs_from_config(scheduler_config)
|
||||
kwargs.update(overrides)
|
||||
|
||||
ref = reference_scheduler_cls(**kwargs)
|
||||
# The pinned Z-Image pipeline mutates this immediately before scheduling.
|
||||
ref.sigma_min = 0.0
|
||||
|
||||
# These values must be serialized by the future production pipeline config.
|
||||
# setdefault also allows this test to consume that config once it lands.
|
||||
fv_kwargs = dict(kwargs)
|
||||
fv_kwargs.setdefault("use_reference_discrete_timesteps", True)
|
||||
fv_kwargs.setdefault("sigma_min", 0.0)
|
||||
fv = FastVideoScheduler(**fv_kwargs)
|
||||
return ref, fv
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def scheduler_pair(reference_scheduler_cls, scheduler_config: dict):
|
||||
return _make_scheduler_pair(reference_scheduler_cls, scheduler_config)
|
||||
|
||||
|
||||
def _run_step_loop(ref_scheduler, fv_scheduler, sample_shape=(2, 4, 32, 32)):
|
||||
torch.manual_seed(123)
|
||||
|
||||
ref_sample = torch.randn(sample_shape, dtype=torch.float32)
|
||||
fv_sample = ref_sample.clone()
|
||||
|
||||
for t in ref_scheduler.timesteps:
|
||||
model_output = torch.randn_like(ref_sample)
|
||||
|
||||
ref_next = ref_scheduler.step(
|
||||
model_output,
|
||||
t,
|
||||
ref_sample,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
fv_next = fv_scheduler.step(
|
||||
model_output,
|
||||
t,
|
||||
fv_sample,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
assert_close(ref_next, fv_next, atol=1e-6, rtol=1e-6)
|
||||
|
||||
ref_sample = ref_next
|
||||
fv_sample = fv_next
|
||||
|
||||
|
||||
def test_zimage_scheduler_parity_default_schedule_and_step(scheduler_pair):
|
||||
ref, fv = scheduler_pair
|
||||
|
||||
num_inference_steps = 8
|
||||
ref.set_timesteps(num_inference_steps=num_inference_steps, device="cpu")
|
||||
fv.set_timesteps(num_inference_steps=num_inference_steps, device="cpu")
|
||||
|
||||
assert_close(ref.timesteps, fv.timesteps, atol=1e-4, rtol=1e-6)
|
||||
assert_close(ref.sigmas, fv.sigmas, atol=1e-7, rtol=1e-6)
|
||||
|
||||
_run_step_loop(ref, fv)
|
||||
|
||||
|
||||
def test_scheduler_positional_args_keep_existing_bindings_and_default_schedule():
|
||||
scheduler = FastVideoScheduler(1000, 1.0, False, 0.7)
|
||||
|
||||
assert scheduler.config.base_shift == 0.7
|
||||
assert scheduler.config.use_reference_discrete_timesteps is False
|
||||
|
||||
scheduler.set_timesteps(num_inference_steps=8, device="cpu")
|
||||
assert_close(scheduler.sigmas[-2], torch.tensor(scheduler.sigma_min))
|
||||
|
||||
|
||||
def test_scheduler_honors_explicit_zero_sigma_min():
|
||||
scheduler = FastVideoScheduler(
|
||||
sigma_min=0.0,
|
||||
use_reference_discrete_timesteps=True,
|
||||
)
|
||||
|
||||
assert scheduler.sigma_min == 0.0
|
||||
scheduler.set_timesteps(num_inference_steps=8, device="cpu")
|
||||
assert_close(scheduler.timesteps[-1], torch.tensor(125.0))
|
||||
|
||||
|
||||
def test_reference_schedule_preserves_float64_linspace_rounding():
|
||||
scheduler = FastVideoScheduler(
|
||||
sigma_min=0.0,
|
||||
use_reference_discrete_timesteps=True,
|
||||
)
|
||||
scheduler.set_timesteps(num_inference_steps=9, device="cpu")
|
||||
|
||||
expected = torch.from_numpy(
|
||||
np.linspace(1000.0, 0.0, 10)[:-1] / 1000.0,
|
||||
).to(torch.float32) * 1000.0
|
||||
float32_regression = torch.from_numpy(
|
||||
np.linspace(1000.0, 0.0, 10, dtype=np.float32)[:-1] / np.float32(1000.0),
|
||||
) * 1000.0
|
||||
|
||||
# The fifth timestep differs by one float32 ULP depending on where the
|
||||
# linspace is rounded, so this assertion fails if dtype=np.float32 returns.
|
||||
assert expected[4].item() != float32_regression[4].item()
|
||||
assert scheduler.timesteps[4].item() == expected[4].item()
|
||||
|
||||
|
||||
def test_zimage_scheduler_parity_dynamic_shifting_with_mu(
|
||||
reference_scheduler_cls,
|
||||
scheduler_config: dict,
|
||||
):
|
||||
ref, fv = _make_scheduler_pair(
|
||||
reference_scheduler_cls,
|
||||
scheduler_config,
|
||||
use_dynamic_shifting=True,
|
||||
)
|
||||
|
||||
mu = 0.75
|
||||
num_inference_steps = 8
|
||||
ref.set_timesteps(num_inference_steps=num_inference_steps, device="cpu", mu=mu)
|
||||
fv.set_timesteps(num_inference_steps=num_inference_steps, device="cpu", mu=mu)
|
||||
|
||||
assert_close(ref.timesteps, fv.timesteps, atol=1e-4, rtol=1e-6)
|
||||
assert_close(ref.sigmas, fv.sigmas, atol=1e-7, rtol=1e-6)
|
||||
|
||||
_run_step_loop(ref, fv)
|
||||
@@ -0,0 +1,130 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Parity test for Z-Image tokenizer loading and tokenization behavior.
|
||||
|
||||
This compares:
|
||||
1) direct transformers AutoTokenizer loading from local tokenizer dir, and
|
||||
2) FastVideo TokenizerLoader loading path,
|
||||
|
||||
using the exact chat-template and tokenization settings from the pinned
|
||||
Z-Image pipeline.
|
||||
|
||||
Usage:
|
||||
pytest tests/local_tests/zimage/test_zimage_tokenizer_parity.py -v
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.models.loader.component_loader import TokenizerLoader
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
ZIMAGE_TOKENIZER_DIR = REPO_ROOT / "official_weights" / "Z-Image" / "tokenizer"
|
||||
PARITY_SCOPE = "production_loader"
|
||||
OFFICIAL_MAX_SEQUENCE_LENGTH = 512
|
||||
OFFICIAL_CHAT_TEMPLATE_KWARGS = {
|
||||
"tokenize": False,
|
||||
"add_generation_prompt": True,
|
||||
"enable_thinking": True,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class _DummyPipelineConfig:
|
||||
text_encoder_configs: tuple = field(default_factory=tuple)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _DummyFastVideoArgs:
|
||||
pipeline_config: _DummyPipelineConfig = field(default_factory=_DummyPipelineConfig)
|
||||
|
||||
|
||||
def _load_reference_tokenizer():
|
||||
if not ZIMAGE_TOKENIZER_DIR.exists():
|
||||
pytest.skip(f"Z-Image tokenizer dir not found: {ZIMAGE_TOKENIZER_DIR}")
|
||||
return AutoTokenizer.from_pretrained(str(ZIMAGE_TOKENIZER_DIR), local_files_only=True)
|
||||
|
||||
|
||||
def _load_fastvideo_tokenizer():
|
||||
if not ZIMAGE_TOKENIZER_DIR.exists():
|
||||
pytest.skip(f"Z-Image tokenizer dir not found: {ZIMAGE_TOKENIZER_DIR}")
|
||||
loader = TokenizerLoader()
|
||||
return loader.load(str(ZIMAGE_TOKENIZER_DIR), _DummyFastVideoArgs())
|
||||
|
||||
|
||||
def test_zimage_tokenizer_loader_and_tokenization_parity():
|
||||
ref_tok = _load_reference_tokenizer()
|
||||
fv_tok = _load_fastvideo_tokenizer()
|
||||
|
||||
prompts = [
|
||||
"A close-up portrait of a cat in warm studio lighting.",
|
||||
"An oil painting of a lighthouse at sunset.",
|
||||
]
|
||||
|
||||
tok_kwargs = {
|
||||
"padding": "max_length",
|
||||
"max_length": OFFICIAL_MAX_SEQUENCE_LENGTH,
|
||||
"truncation": True,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
|
||||
ref_out = ref_tok(prompts, **tok_kwargs)
|
||||
fv_out = fv_tok(prompts, **tok_kwargs)
|
||||
|
||||
assert_close(ref_out["input_ids"], fv_out["input_ids"], atol=0, rtol=0)
|
||||
assert_close(ref_out["attention_mask"], fv_out["attention_mask"], atol=0, rtol=0)
|
||||
|
||||
# Compare key tokenizer attributes that affect generation-time behavior.
|
||||
assert ref_tok.padding_side == fv_tok.padding_side
|
||||
assert ref_tok.pad_token_id == fv_tok.pad_token_id
|
||||
assert ref_tok.eos_token_id == fv_tok.eos_token_id
|
||||
|
||||
|
||||
@pytest.mark.skipif(not ZIMAGE_TOKENIZER_DIR.exists(), reason="Z-Image tokenizer assets are required")
|
||||
def test_zimage_tokenizer_chat_template_parity():
|
||||
ref_tok = _load_reference_tokenizer()
|
||||
fv_tok = _load_fastvideo_tokenizer()
|
||||
|
||||
assert callable(getattr(ref_tok, "apply_chat_template", None)), (
|
||||
"Pinned Z-Image tokenizer must expose apply_chat_template"
|
||||
)
|
||||
assert callable(getattr(fv_tok, "apply_chat_template", None)), (
|
||||
"FastVideo TokenizerLoader dropped the required apply_chat_template API"
|
||||
)
|
||||
|
||||
messages = [{"role": "user", "content": "Describe a futuristic city skyline."}]
|
||||
|
||||
ref_text = ref_tok.apply_chat_template(messages, **OFFICIAL_CHAT_TEMPLATE_KWARGS)
|
||||
fv_text = fv_tok.apply_chat_template(messages, **OFFICIAL_CHAT_TEMPLATE_KWARGS)
|
||||
|
||||
assert ref_text == fv_text
|
||||
|
||||
# Z-Image deliberately enables Qwen's thinking template. Guard against a
|
||||
# pipeline silently using the generic non-thinking prompt path.
|
||||
non_thinking_text = fv_tok.apply_chat_template(
|
||||
messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
enable_thinking=False,
|
||||
)
|
||||
assert fv_text != non_thinking_text
|
||||
|
||||
tokenization_kwargs = {
|
||||
"padding": "max_length",
|
||||
"max_length": OFFICIAL_MAX_SEQUENCE_LENGTH,
|
||||
"truncation": True,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
ref_tokens = ref_tok(ref_text, **tokenization_kwargs)
|
||||
fv_tokens = fv_tok(fv_text, **tokenization_kwargs)
|
||||
|
||||
assert_close(ref_tokens["input_ids"], fv_tokens["input_ids"], atol=0, rtol=0)
|
||||
assert_close(ref_tokens["attention_mask"], fv_tokens["attention_mask"], atol=0, rtol=0)
|
||||
@@ -0,0 +1,226 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Parity test for Z-Image VAE (AutoencoderKL decode path).
|
||||
|
||||
This compares FastVideo's AutoencoderKL wrapper against the Z-Image
|
||||
reference AutoencoderKL implementation using the same local weights.
|
||||
|
||||
Usage:
|
||||
pytest tests/local_tests/zimage/test_zimage_vae_parity.py -v
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import json
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.configs.models.vaes.autoencoder_kl import AutoencoderKLVAEConfig
|
||||
from fastvideo.models.loader.component_loader import VAELoader
|
||||
from fastvideo.models.vaes.autoencoder_kl import AutoencoderKL as FastVideoAutoencoderKL
|
||||
from fastvideo.pipelines.stages.decoding import DecodingStage
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
ZIMAGE_REPO = REPO_ROOT / "Z-Image"
|
||||
ZIMAGE_SRC = REPO_ROOT / "Z-Image" / "src"
|
||||
ZIMAGE_VAE_DIR = REPO_ROOT / "official_weights" / "Z-Image" / "vae"
|
||||
ZIMAGE_VAE_CFG = ZIMAGE_VAE_DIR / "config.json"
|
||||
ZIMAGE_VAE_WEIGHTS = ZIMAGE_VAE_DIR / "diffusion_pytorch_model.safetensors"
|
||||
ZIMAGE_REFERENCE_REVISION = "26f23eda626ffadda020b04ff79488e1d72004cd"
|
||||
ZIMAGE_VAE_SCALING_FACTOR = 0.3611
|
||||
ZIMAGE_VAE_SHIFT_FACTOR = 0.1159
|
||||
PARITY_SCOPE = "both"
|
||||
|
||||
|
||||
def _require_pinned_reference_module(module_name: str, source_file: Path):
|
||||
if not ZIMAGE_REPO.exists():
|
||||
pytest.skip(f"Pinned Z-Image reference clone not found: {ZIMAGE_REPO}")
|
||||
if not source_file.is_file():
|
||||
pytest.fail(f"Z-Image reference clone is incomplete; missing {source_file}")
|
||||
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["git", "-C", str(ZIMAGE_REPO), "rev-parse", "HEAD"],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
except (FileNotFoundError, subprocess.CalledProcessError) as exc:
|
||||
pytest.fail(f"Cannot verify Z-Image reference revision: {exc}")
|
||||
|
||||
actual_revision = result.stdout.strip()
|
||||
assert actual_revision == ZIMAGE_REFERENCE_REVISION, (
|
||||
"Z-Image reference clone is not at the pinned revision: "
|
||||
f"expected {ZIMAGE_REFERENCE_REVISION}, got {actual_revision}"
|
||||
)
|
||||
|
||||
if str(ZIMAGE_SRC) not in sys.path:
|
||||
sys.path.insert(0, str(ZIMAGE_SRC))
|
||||
try:
|
||||
module = importlib.import_module(module_name)
|
||||
except Exception as exc:
|
||||
pytest.fail(f"Cannot import pinned Z-Image module {module_name}: {exc}")
|
||||
|
||||
module_file = Path(module.__file__ or "").resolve()
|
||||
assert module_file.is_relative_to(ZIMAGE_SRC.resolve()), (
|
||||
f"{module_name} resolved outside the pinned clone: {module_file}"
|
||||
)
|
||||
return module
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def reference_autoencoder_cls():
|
||||
module = _require_pinned_reference_module(
|
||||
"zimage.autoencoder",
|
||||
ZIMAGE_SRC / "zimage" / "autoencoder.py",
|
||||
)
|
||||
return module.AutoencoderKL
|
||||
|
||||
|
||||
def _load_cfg() -> dict:
|
||||
if not ZIMAGE_VAE_CFG.exists():
|
||||
pytest.skip(f"Z-Image VAE config not found: {ZIMAGE_VAE_CFG}")
|
||||
with ZIMAGE_VAE_CFG.open("r", encoding="utf-8") as f:
|
||||
cfg = json.load(f)
|
||||
cfg.pop("_class_name", None)
|
||||
cfg.pop("_diffusers_version", None)
|
||||
cfg.pop("_name_or_path", None)
|
||||
return cfg
|
||||
|
||||
|
||||
def _load_weights() -> dict[str, torch.Tensor]:
|
||||
if not ZIMAGE_VAE_WEIGHTS.exists():
|
||||
pytest.skip(f"Z-Image VAE weights not found: {ZIMAGE_VAE_WEIGHTS}")
|
||||
return safetensors_load_file(str(ZIMAGE_VAE_WEIGHTS), device="cpu")
|
||||
|
||||
|
||||
def _build_reference(reference_autoencoder_cls, cfg: dict, weights: dict[str, torch.Tensor]) -> torch.nn.Module:
|
||||
ref = reference_autoencoder_cls(
|
||||
in_channels=cfg["in_channels"],
|
||||
out_channels=cfg["out_channels"],
|
||||
down_block_types=tuple(cfg["down_block_types"]),
|
||||
up_block_types=tuple(cfg["up_block_types"]),
|
||||
block_out_channels=tuple(cfg["block_out_channels"]),
|
||||
layers_per_block=cfg["layers_per_block"],
|
||||
latent_channels=cfg["latent_channels"],
|
||||
norm_num_groups=cfg["norm_num_groups"],
|
||||
scaling_factor=cfg["scaling_factor"],
|
||||
shift_factor=cfg["shift_factor"],
|
||||
use_quant_conv=cfg["use_quant_conv"],
|
||||
use_post_quant_conv=cfg["use_post_quant_conv"],
|
||||
mid_block_add_attention=cfg["mid_block_add_attention"],
|
||||
).eval()
|
||||
ref.load_state_dict(weights, strict=True)
|
||||
return ref
|
||||
|
||||
|
||||
def _build_fastvideo(cfg: dict, weights: dict[str, torch.Tensor]) -> torch.nn.Module:
|
||||
fv_cfg = AutoencoderKLVAEConfig()
|
||||
fv_cfg.update_model_arch(cfg)
|
||||
fv = FastVideoAutoencoderKL(fv_cfg).eval()
|
||||
fv.load_state_dict(weights, strict=True)
|
||||
return fv
|
||||
|
||||
|
||||
def _load_fastvideo_production_vae(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(
|
||||
"fastvideo.models.loader.component_loader.get_local_torch_device",
|
||||
lambda: torch.device("cpu"),
|
||||
)
|
||||
pipeline_config = SimpleNamespace(
|
||||
vae_config=AutoencoderKLVAEConfig(),
|
||||
vae_precision="fp32",
|
||||
)
|
||||
fastvideo_args = SimpleNamespace(
|
||||
model_paths={},
|
||||
pipeline_config=pipeline_config,
|
||||
vae_cpu_offload=False,
|
||||
)
|
||||
vae = VAELoader().load(str(ZIMAGE_VAE_DIR), fastvideo_args)
|
||||
assert isinstance(vae, FastVideoAutoencoderKL)
|
||||
assert fastvideo_args.model_paths["vae"] == str(ZIMAGE_VAE_DIR)
|
||||
return vae
|
||||
|
||||
|
||||
def _official_decode_latents(latents: torch.Tensor, vae: torch.nn.Module) -> torch.Tensor:
|
||||
scaling_factor = vae.config.scaling_factor
|
||||
shift_factor = vae.config.shift_factor or 0.0
|
||||
dtype = next(vae.parameters()).dtype
|
||||
return (latents.to(dtype=dtype) / scaling_factor) + shift_factor
|
||||
|
||||
|
||||
def test_zimage_decode_normalization_uses_autoencoder_config_values():
|
||||
vae = SimpleNamespace(
|
||||
config=SimpleNamespace(
|
||||
latents_mean=None,
|
||||
latents_std=None,
|
||||
scaling_factor=ZIMAGE_VAE_SCALING_FACTOR,
|
||||
shift_factor=ZIMAGE_VAE_SHIFT_FACTOR,
|
||||
)
|
||||
)
|
||||
latents = torch.tensor([[[[0.0, 0.5], [-0.5, 1.0]]]], dtype=torch.float32)
|
||||
|
||||
actual = DecodingStage(vae)._denormalize_latents(latents)
|
||||
expected = latents / ZIMAGE_VAE_SCALING_FACTOR + ZIMAGE_VAE_SHIFT_FACTOR
|
||||
|
||||
assert_close(actual, expected, atol=0, rtol=0)
|
||||
|
||||
|
||||
def test_zimage_vae_decode_parity(reference_autoencoder_cls):
|
||||
torch.manual_seed(7)
|
||||
|
||||
cfg = _load_cfg()
|
||||
weights = _load_weights()
|
||||
|
||||
ref = _build_reference(reference_autoencoder_cls, cfg, weights)
|
||||
fv = _build_fastvideo(cfg, weights)
|
||||
|
||||
# Decode-only parity is the critical path for Z-Image pipeline integration.
|
||||
latents = torch.randn(1, cfg["latent_channels"], 8, 8, dtype=torch.float32)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_out = ref.decode(latents, return_dict=False)[0].detach().float()
|
||||
fv_out = fv.decode(latents, return_dict=False)[0].detach().float()
|
||||
|
||||
assert ref_out.shape == fv_out.shape
|
||||
assert_close(ref_out, fv_out, atol=1e-4, rtol=1e-4)
|
||||
|
||||
|
||||
def test_zimage_vae_production_loader_and_decode_normalization(
|
||||
reference_autoencoder_cls,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
torch.manual_seed(17)
|
||||
|
||||
cfg = _load_cfg()
|
||||
weights = _load_weights()
|
||||
assert cfg["scaling_factor"] == ZIMAGE_VAE_SCALING_FACTOR
|
||||
assert cfg["shift_factor"] == ZIMAGE_VAE_SHIFT_FACTOR
|
||||
ref = _build_reference(reference_autoencoder_cls, cfg, weights)
|
||||
fv = _load_fastvideo_production_vae(monkeypatch)
|
||||
|
||||
assert ref.config.scaling_factor == cfg["scaling_factor"]
|
||||
assert fv.config.scaling_factor == cfg["scaling_factor"]
|
||||
assert ref.config.shift_factor == cfg["shift_factor"]
|
||||
assert fv.config.shift_factor == cfg["shift_factor"]
|
||||
|
||||
pipeline_latents = torch.randn(1, cfg["latent_channels"], 8, 8, dtype=torch.float32)
|
||||
ref_decode_latents = _official_decode_latents(pipeline_latents, ref)
|
||||
fv_decode_latents = DecodingStage(fv)._denormalize_latents(pipeline_latents)
|
||||
assert_close(ref_decode_latents, fv_decode_latents, atol=0, rtol=0)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_out = ref.decode(ref_decode_latents, return_dict=False)[0].detach().float()
|
||||
fv_out = fv.decode(fv_decode_latents, return_dict=False)[0].detach().float()
|
||||
|
||||
assert ref_out.shape == fv_out.shape
|
||||
assert_close(ref_out, fv_out, atol=1e-4, rtol=1e-4)
|
||||
Reference in New Issue
Block a user