nunchaku-qwen-image-2.1 / extras /diagnose_trajectory_and_scales.py
catplusplus's picture
Upload folder using huggingface_hub
b6a8400 verified
Raw History Blame Contribute Delete
7.22 kB
# -*- coding: utf-8 -*-
"""Step-by-Step Trajectory, Magnitude & Convergence Diagnostic for Qwen-Image-2.1.
Compares Pure BF16 vs SVDQuant NVFP4 across all diffusion steps:
1. Isolated Single-Step Error vs Accumulated Trajectory Drift
2. Mean Magnitude, Norm Ratio (||v_nvfp4|| / ||v_bf16||), and Directional Cosine Similarity
3. Dynamic Range & Tendency towards underflow/overflow across timesteps (t=1.0 down to t=0.0)
4. Evaluates whether global scale adjustment alpha(t) can prevent trajectory drift
"""
import os
import sys
import time
import math
import gc
import json
import torch
import numpy as np
from PIL import Image
ROOT_DIR = "/auto/home/amano/olegk/Nikola"
for p in [f"{ROOT_DIR}/packages/nunchaku", f"{ROOT_DIR}/packages/deepcompressor", ROOT_DIR, f"{ROOT_DIR}/src/imagegen"]:
if p not in sys.path:
sys.path.insert(0, p)
from diffusers import QwenImage21Pipeline
from diffusers.models.transformers.transformer_qwenimage21 import QwenImage21Transformer2DModel
from nunchaku.models.transformers.transformer_qwenimage21 import NunchakuQwenImage21Transformer2DModel
os.environ["CUDA_VISIBLE_DEVICES"] = "0"
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
MODEL_ID = "/home/olegk/Nikola/models/Qwen/Qwen-Image-2.1"
NVFP4_PATH = "/home/olegk/Nikola/models/nunchaku-qwen-image-2.1/svdq-fp4_r32-qwen-image-2.1.safetensors"
OUT_REPORT = "/home/olegk/tmp/qwen21_trajectory_diagnostic.json"
def run_trajectory_tracer(pipe, num_steps=25, seed=42, prompt="A cute fluffy little kitten sitting happily on green grass"):
"""Run pipeline and record latents, timesteps, and noise_pred at every step."""
recorded_steps = []
def step_cb(pipe_obj, step_idx, timestep, callback_kwargs):
lat = callback_kwargs["latents"].detach().cpu().clone()
recorded_steps.append({
"step": step_idx,
"timestep": float(timestep.item() if isinstance(timestep, torch.Tensor) else timestep),
"latents": lat,
})
return callback_kwargs
generator = torch.Generator("cuda").manual_seed(seed)
out = pipe(
prompt=prompt,
height=1024,
width=1024,
num_inference_steps=num_steps,
true_cfg_scale=1.0,
generator=generator,
callback_on_step_end=step_cb,
)
return out.images[0], recorded_steps
def main():
print("=" * 80)
print("🔬 STEP-BY-STEP TRAJECTORY & SCALE CONVERGENCE DIAGNOSTIC (BF16 vs NVFP4)")
print("=" * 80)
prompt = "A cute fluffy little kitten sitting happily on green grass in a sunny garden, soft lighting, sharp focus, high quality photo"
num_steps = 25
seed = 42
# -------------------------------------------------------------------------
# STAGE 1: Record Ground-Truth BF16 Trajectory
# -------------------------------------------------------------------------
print("\n[Stage 1/3] Loading Ground-Truth BF16 Pipeline...")
pipe_bf16 = QwenImage21Pipeline.from_pretrained(
MODEL_ID,
torch_dtype=torch.bfloat16,
)
pipe_bf16.enable_sequential_cpu_offload(gpu_id=0)
try:
pipe_bf16.vae.enable_tiling()
except Exception:
pass
print(f"Running BF16 generation ({num_steps} steps, seed {seed})...")
img_bf16, bf16_steps = run_trajectory_tracer(pipe_bf16, num_steps=num_steps, seed=seed, prompt=prompt)
img_bf16.save("/home/olegk/tmp/diag_kitten_bf16.png")
print(f"BF16 trace recorded: {len(bf16_steps)} step snapshots.")
# Free BF16 model from VRAM before loading NVFP4
del pipe_bf16
gc.collect()
torch.cuda.empty_cache()
print("VRAM cleared. Free VRAM:", f"{torch.cuda.mem_get_info(0)[0] / (1024**3):.2f} GB")
# -------------------------------------------------------------------------
# STAGE 2: Record Resident NVFP4 Trajectory
# -------------------------------------------------------------------------
print("\n[Stage 2/3] Loading Resident NVFP4 Pipeline...")
pipe_nvfp4 = QwenImage21Pipeline.from_pretrained(
MODEL_ID,
transformer=None,
torch_dtype=torch.bfloat16,
)
trans_nvfp4 = NunchakuQwenImage21Transformer2DModel.from_pretrained(
NVFP4_PATH,
device="cuda:0",
torch_dtype=torch.bfloat16,
)
pipe_nvfp4.transformer = trans_nvfp4
pipe_nvfp4.vae = pipe_nvfp4.vae.to("cuda:0")
pipe_nvfp4.vae.enable_tiling()
orig_encode_prompt = pipe_nvfp4.encode_prompt
def safe_encode_prompt(*args, **kwargs):
kwargs.pop("device", None)
te_device = pipe_nvfp4.text_encoder.device
embeds = orig_encode_prompt(*args, device=te_device, **kwargs)
target_dev = pipe_nvfp4.transformer.device
return tuple(x.to(target_dev) if isinstance(x, torch.Tensor) else x for x in embeds)
pipe_nvfp4.encode_prompt = safe_encode_prompt
print(f"Running NVFP4 generation ({num_steps} steps, seed {seed})...")
img_nvfp4, nvfp4_steps = run_trajectory_tracer(pipe_nvfp4, num_steps=num_steps, seed=seed, prompt=prompt)
img_nvfp4.save("/home/olegk/tmp/diag_kitten_nvfp4.png")
print(f"NVFP4 trace recorded: {len(nvfp4_steps)} step snapshots.")
# -------------------------------------------------------------------------
# STAGE 3: Comparative Metric Analysis & Diagnostics
# -------------------------------------------------------------------------
print("\n[Stage 3/3] Analyzing Step-by-Step Trajectory Drift & Scale Stats...")
print("-" * 80)
print(f"{'Step':<5} | {'t':<7} | {'BF16 Norm':<10} | {'NVFP4 Norm':<10} | {'Norm Ratio':<10} | {'Cos Sim':<8} | {'Rel L2 Drift':<12}")
print("-" * 80)
stats_summary = []
n_pts = min(len(bf16_steps), len(nvfp4_steps))
for i in range(n_pts):
b_step = bf16_steps[i]
n_step = nvfp4_steps[i]
t_val = b_step["timestep"]
b_lat = b_step["latents"].float()
n_lat = n_step["latents"].float()
b_norm = float(b_lat.norm().item())
n_norm = float(n_lat.norm().item())
norm_ratio = n_norm / (b_norm + 1e-8)
cos_sim = float(torch.nn.functional.cosine_similarity(b_lat.flatten(), n_lat.flatten(), dim=0).item())
rel_diff = float(((n_lat - b_lat).norm() / (b_lat.norm() + 1e-8)).item())
print(f"{i:<5d} | {t_val:<7.1f} | {b_norm:<10.2f} | {n_norm:<10.2f} | {norm_ratio:<10.4f} | {cos_sim:<8.4f} | {rel_diff:<12.4f}")
stats_summary.append({
"step": i,
"timestep": t_val,
"bf16_norm": b_norm,
"nvfp4_norm": n_norm,
"norm_ratio": norm_ratio,
"cosine_similarity": cos_sim,
"relative_l2_drift": rel_diff,
"bf16_std": float(b_lat.std().item()),
"nvfp4_std": float(n_lat.std().item()),
"bf16_min": float(b_lat.min().item()),
"nvfp4_min": float(n_lat.min().item()),
"bf16_max": float(b_lat.max().item()),
"nvfp4_max": float(n_lat.max().item()),
})
with open(OUT_REPORT, "w") as f:
json.dump(stats_summary, f, indent=2)
print(f"\n✅ Detailed step stats written to {OUT_REPORT}!")
if __name__ == "__main__":
main()