Compare commits

...
Author SHA1 Message Date
Raghav 70bd4bb756 [bugfix]: don't let a stale HF-passthrough marker override a moved encoder (#1339)
_fastvideo_input_device is stamped at load time with the loader's target
device, which is "cpu" when text_encoder_cpu_offload is on. encode_text then
moves such an encoder to the compute device for the forward, after which the
marker is stale. Honouring it fed cpu token ids to cuda weights:

  RuntimeError: Expected all tensors to be on the same device, but got index
  is on cpu, different from other tensors on cuda:0

Seen on FLUX.2 Klein, whose Qwen3 text encoder takes the HF-passthrough path.
Once the encoder has been moved, its param device is authoritative; the marker
only speaks when nothing moved and the module has no parameters of its own.
2026-07-13 17:26:46 -07:00
RaghavandClaude Opus 4.8 35e55aa7f6 [bugfix]: keep DecodingStage scaling_factor module-only (#1339)
_denormalize_latents previously ran its scaling branch only when the VAE
module exposed scaling_factor -- in practice only gamecraftvae does. Falling
back to vae.config.scaling_factor turned that branch on for every VAE whose
config declares one (flux2, hunyuanvideo, hunyuanvideo15, ltx2), newly
dividing latents that had always passed through untouched. Surfaced as
FLUX.2 4B/9B SSIM failures.

Keep the latents_mean/latents_std guard, which only tests the values instead
of attribute presence so a config declaring them as None falls through rather
than raising in torch.tensor(None).

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-13 17:26:46 -07:00
SolitaryThinker 67c023fca9 [fix]: address Z-Image review findings (#1339) 2026-07-13 17:26:46 -07:00
Raghav bf839dde18 [fix]: Z-Image reuses the shared Qwen3 encoder; fix its batched RoPE positions
main gained a config-driven Qwen3 text encoder (Qwen3ForCausalLM + Qwen3TextConfig)
via the Flux2 Klein port (#1349). Reuse it for Z-Image instead of the parallel
~540-line bespoke Qwen3Model this PR originally added: drop the duplicate encoder
+ config, and map Z-Image-Turbo's "Qwen3Model" architecture string to the shared
encoder in the registry. update_model_arch populates Z-Image's dims from config.json.

Validating the reuse on L40S surfaced a real, batch-only divergence in the shared
encoder: it built position_ids as [1, seq_len], but the rotary layer flattens
positions to num_tokens and reshapes q/k to (num_tokens, -1, head_dim). For
batch>1 that folded the batch dim into the head dim and misaligned RoPE (fp32
mean_diff 0.22 at batch=2; batch=1 was fine). Fix: expand position_ids to
[batch_size, seq_len]. batch=1 is byte-identical, so Flux2 Klein is unaffected;
this also fixes a latent batch bug in the shared encoder.

Encoder parity now PASSES on L40S (Z-Image-Turbo): fp32 bit-exact against the
shared encoder (both batch elements, max=0.0000); bf16 within the existing
thresholds (last_hidden mean ~0.016, pre-norm mean ~0.07-0.08).

Also in this PR (unchanged): the gated scheduler option
use_reference_discrete_timesteps (default False) for Z-Image timestep parity,
the Z-Image component parity tests (encoder/scheduler/tokenizer/VAE), and the
encoder parity test's strict-load allowlist + OOM-safe two-model handling
(free HF ref + gc.collect before loading FastVideo).
2026-07-13 17:26:46 -07:00
Raghav b1d954c894 [docs]: pin GPU SKU + final empirical numbers in PORT_STATUS
Component parity validated on Tongyi-MAI/Z-Image-Turbo on NVIDIA A40
(46068 MiB, driver 565.57.01) as of 2026-05-12. All thresholds pass
with 1.5-1.6x headroom over observed:
  last_hidden_state bf16 mean (worst):   0.0168 < 0.025
  last_hidden_state bf16 median (worst): 0.0127 < 0.020
  hidden_states[-2] bf16 mean (worst):   0.0739 < 0.120
  hidden_states[-2] bf16 median (worst): 0.0625 < 0.100

Handoff notes updated to reflect that component parity is done; next
port-stack steps (DiT, pipeline, conversion, SSIM) are scoped to a
future PR, not #1339.
2026-07-13 17:26:46 -07:00
Raghav e4eef6ed20 [test]: calibrate bf16 encoder parity thresholds against observed Z-Image-Turbo
Per-layer bf16 diagnostic test on Z-Image-Turbo (Qwen3 35 layers, not the
24 originally assumed) confirmed the divergence between FastVideo's
Qwen3 port and HF Qwen3 is a textbook bf16-tail signature, not a code
bug:
- fp32 forward is bit-exact (last_hidden_state max=0.0000)
- bf16 per-layer mean drift grows smoothly from 5e-4 to 7.5e-2 over 35
  layers; no single-layer spike (consecutive layer ratio ≤ 1.5x)
- embedding (idx 0) is bit-exact in bf16 too — weight load is right

Root cause is well-understood: FastVideo's TP-aware fused linears
(QKVParallel, MergedColumnParallel, SiluAndMul) execute matmul ops in a
different order than HF's unfused equivalents. In fp32 the per-element
differences cancel statistically (~zero mean drift); in bf16 the
rounding-order accumulates into a long max tail.

Per add-model-02-parity calibration block: replace element-wise
assert_close (meaningless when the max tail is 4.0) with distribution
checks (mean + median). Thresholds calibrated to empirical observation:
  last_hidden_state (post-norm):  mean < 0.025, median < 0.020
  hidden_states[-2] (pre-norm):   mean < 0.120, median < 0.100
Headroom is ~1.6x over observed (0.0152, 0.0117, 0.0754, 0.0625).
fp32 path keeps strict assert_close at atol=1e-4.

Also revert `torch_dtype=dtype` -> `dtype=dtype`: transformers 4.57.3
emits a deprecation warning recommending `dtype=`. The earlier switch
(Copilot review) was wrong-direction for the pinned transformers version.
The two config-key strips that mention `torch_dtype` stay — they're
stripping the legacy serialized config key, not the constructor kwarg.

PORT_STATUS.md updated: pinned Z-Image SHA, pinned HF id (Z-Image-Turbo),
empirical bf16 signature recorded in the decisions table, 35-layer
correction reflected in the parity table.
2026-07-13 17:26:46 -07:00
Raghav 7390f528c6 [test]: add per-layer bf16 diagnostic for Qwen3 encoder
Diagnostic-only test that prints abs-diff (max/mean/median/p99) for each
of the 25 hidden states (embedding + 24 layer outputs + post-norm) in
bf16. Lets a reviewer distinguish:
- Healthy bf16 tail: monotonic growth per layer, median ≪ atol
- Real op mismatch: single-layer spike where one layer's diff is
  >>2x the previous

Always passes; output is captured via `-s` and inspected manually.
Run after the parametrized parity test (which currently fails on the
strict mean-drift threshold) to decide whether to relax thresholds or
investigate a specific layer.
2026-07-13 17:26:46 -07:00
Raghav 92168e5020 [fix]: address Gemini + Copilot review on PR #1339
scheduler:
- Restore numpy default (float64) on the non-reference timestep linspace
  path. The earlier `dtype=np.float32` silently changed rounded timestep
  values for every model already using this scheduler. Float32 is now
  only applied on the new `use_reference_discrete_timesteps=True` branch.
- Document `use_reference_discrete_timesteps` in the class docstring with
  default + when-to-enable + backward-compat notes.

qwen3 encoder:
- `torch.arange(...)` for position_ids now passes `dtype=torch.long`
  explicitly so embedding/RoPE indices are unambiguous.
- Drop the redundant `isinstance(x.device.type, str)` check; `device.type`
  is always a str. Keep the mps -> cpu fallback for torch.autocast.

encoder parity test:
- `AutoModel.from_pretrained` uses `torch_dtype=` to match repo convention
  (the rest of the repo and existing transformers callers all use this
  keyword). transformers 4.57.3 accepts both, but `torch_dtype` keeps the
  test green across the older transformers versions that some downstream
  forks pin.
2026-07-13 17:26:46 -07:00
Raghav d4b8f08674 [docs]: add tests/local_tests/zimage/README.md + PORT_STATUS.md
Required by `.agents/skills/add-model/contracts/port_state.md` and
`add-model-10-pr-review` "Quality and evidence lane" before any handoff.

README captures: reference assets, weight-dir layout, run commands,
per-component status, and live blockers (lm_head allowlist, scheduler
config-flag pin requirement, tokenizer `text_len` reconciliation).

PORT_STATUS captures: component matrix, parity commands + last result,
open questions (pin Z-Image clone SHA, pin HF id), issues I001-I005
(lm_head allowlist resolved, scheduler-config pin still open, transformer
+ pipeline NOT_STARTED), and the bf16 / scheduler-config-widening
decisions made in this PR.
2026-07-13 17:26:46 -07:00
Raghav 13daf39e91 [test]: forward full scheduler_config.json into parity, not 3 keys
The fixture only forwarded `num_train_timesteps`, `shift`, and
`use_dynamic_shifting` from `scheduler_config.json`. Any future field in
the on-disk config (e.g. `time_shift_type`, `invert_sigmas`, etc.) would
be silently dropped on both sides — parity would pass while the produced
schedulers diverged from the file Z-Image actually ships.

Spread the full dict (minus Diffusers loader keys) into the constructor
on both sides. The dynamic-shifting test follows the same shape.
2026-07-13 17:26:46 -07:00
Raghav 2d0e552a7e [test]: add tests/local_tests/zimage/__init__.py
All sibling family dirs (`encoders/`, `pipelines/`, `transformers/`,
`vaes/`, `sd35/`) ship an `__init__.py`. Z-Image was the only port-in-
progress dir missing one; conftest setups that walk packages will skip
the dir without it.
2026-07-13 17:26:46 -07:00
Raghav 46d1f91c7e [test]: Qwen3 encoder parity in bf16 with calibrated tolerance + diagnostics
The original test forced fp32 because the strict `atol=1e-4` failed in bf16
across 24 transformer blocks. Per add-model-02-parity's calibration block,
deep encoders in bf16 accumulate per-GEMM epsilon into a long tail — the
right fix is to keep both dtypes covered, log per-batch
max/mean/median/p99 diagnostics, and require median≈0 + abs-mean drift
below threshold so a real bug (mean ≫ atol) cannot pass.

Also: assert that the unexpected-key surface from the production loader
equals `Qwen3Model.ALLOWED_UNEXPECTED_KEYS` so future `Qwen3ForCausalLM`
checkpoint changes are caught. Fix stale docstring filename. Fix typo
`torch.bloat16` (removed entirely with the dtype rewrite).
2026-07-13 17:26:46 -07:00
Mrinaal Dogra 9307aeb590 [feat]: Add Z-Image Qwen3 Encoder and parity test 2026-07-13 17:26:46 -07:00
Mrinaal Dogra 926f0c9ab4 [test]: add Z-Image VAE decode parity test 2026-07-13 17:26:46 -07:00
Mrinaal Dogra b2ba79a4cd [test]: add Z-Image scheduler parity and reference timestep option 2026-07-13 17:26:46 -07:00
16 changed files with 1857 additions and 49 deletions
+1
View File
@@ -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
+78 -28
View File
@@ -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
+3
View File
@@ -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)
+24 -13
View File
@@ -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
+18 -4
View File
@@ -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"]
+116 -2
View File
@@ -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"
+120
View File
@@ -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.
+142
View File
@@ -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.
+1
View File
@@ -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)